code wiki / (root) / nx_zimage_sdpa_verify.nx

nx_zimage_sdpa_verify.nx source

↩ module page · 159 lines · 5998 B

1// nx_zimage_sdpa_verify.nx -- SOVEREIGN Z-Image attention sub-piece 2b+2c: qk-norm+rope+SDPA, verified. 2// 3// Small organ (no big linears): takes the oracle's dumped qkv [4,11520], applies per-head QK-RMSNorm(qn/kn) 4// + axial RoPE(cosr/sinr) to Q,K, then per-head SDPA (softmax(q.k/sqrt128) v) -> ao [4,3840]; verifies vs 5// the oracle's dumped ao_pre (attention output BEFORE out-proj). Reference dumps oracle-only; organ = Nishi. 6// license_tier: ORIGINAL 7import "nx_syscalls.nx" 8import "nx_le.nx" 9import "nx_f32.nx" 10import "nx_f32_div.nx" 11import "nx_f32_cvt.nx" 12import "nx_f32_exp.nx" 13import "nx_strconv.nx" 14const K_MAGIC_3840: i64 = 3840 15const K_MAGIC_11520: i64 = 11520 16const K_MAGIC_100000: i64 = 100000 17 18func zsd_load(name: *u8, nl: i64, n_floats: i64) -> *u8 { 19 let base: *u8 = "/mnt/c/Users/elder/AppData/Local/Temp/claude/C--Users-elder/7be78b15-304c-449e-afe8-4d5bd7ddaa9c/scratchpad/zblk/" as *u8 20 let path: *u8 = sys_mmap(256) 21 var p: i64 = 0 22 var i: i64 = 0 23 while base[i] != 0 { path[p] = base[i]; p = p + 1; i = i + 1 } 24 i = 0 25 while i < nl { path[p] = name[i]; p = p + 1; i = i + 1 } 26 path[p] = 0x2E; p = p + 1 27 path[p] = 0x66; p = p + 1 28 path[p] = 0x33; p = p + 1 29 path[p] = 0x32; p = p + 1 30 path[p] = 0 31 let fd: i64 = sys_openat_rd(path) 32 if fd < 0 { return 0 as *u8 } 33 let bytes: i64 = n_floats * 4 34 let buf: *u8 = sys_mmap(bytes + 64) 35 var tot: i64 = 0 36 var go: i64 = 1 37 while go == 1 { 38 let r: i64 = sys_read(fd, ((buf as i64) + tot) as *u8, bytes - tot) 39 if r <= 0 { go = 0 } else { tot = tot + r; if tot >= bytes { go = 0 } } 40 } 41 sys_close(fd) 42 return buf 43} 44 45func zsd_qknorm_rope(qkv: *i64, off: i64, w: *u8, cosr: *u8, sinr: *u8, t: i64, eps: i64) -> i64 { 46 var ss: i64 = 0 47 var d: i64 = 0 48 while d < 128 { let v: i64 = qkv[off + d]; ss = nx_f32_add(ss, nx_f32_mul(v, v)); d = d + 1 } 49 let rms: i64 = nx_f32_sqrt(nx_f32_add(nx_f32_div(ss, nx_i32_to_f32(128)), eps)) 50 d = 0 51 while d < 128 { qkv[off + d] = nx_f32_div(nx_f32_mul(qkv[off + d], nx_le_read_u32(w, d * 4)), rms); d = d + 1 } 52 var j: i64 = 0 53 while j < 64 { 54 let c: i64 = nx_le_read_u32(cosr, (t * 64 + j) * 4) 55 let s: i64 = nx_le_read_u32(sinr, (t * 64 + j) * 4) 56 let x0: i64 = qkv[off + 2 * j] 57 let x1: i64 = qkv[off + 2 * j + 1] 58 qkv[off + 2 * j] = nx_f32_sub(nx_f32_mul(x0, c), nx_f32_mul(x1, s)) 59 qkv[off + 2 * j + 1] = nx_f32_add(nx_f32_mul(x0, s), nx_f32_mul(x1, c)) 60 j = j + 1 61 } 62 return 0 63} 64 65func main() -> i64 { 66 let D: i64 = K_MAGIC_3840 67 let H: i64 = 30 68 let HD: i64 = 128 69 let NT: i64 = 4 70 let QKV: i64 = K_MAGIC_11520 71 let qkvd: *u8 = zsd_load("qkv" as *u8, 3, NT * QKV) 72 let qn: *u8 = zsd_load("qn" as *u8, 2, HD) 73 let kn: *u8 = zsd_load("kn" as *u8, 2, HD) 74 let cosr: *u8 = zsd_load("cosr" as *u8, 4, NT * 64) 75 let sinr: *u8 = zsd_load("sinr" as *u8, 4, NT * 64) 76 let aog: *u8 = zsd_load("ao_pre" as *u8, 6, NT * D) 77 if (qkvd as i64) == 0 { return 30 } 78 if (aog as i64) == 0 { return 31 } 79 80 let eps: i64 = nx_f32_div(nx_i32_to_f32(1), nx_i32_to_f32(K_MAGIC_100000)) 81 let scale: i64 = nx_f32_div(nx_i32_to_f32(1), nx_f32_sqrt(nx_i32_to_f32(128))) 82 83 // mutable copy of qkv 84 let qkva: *i64 = sys_mmap(NT * QKV * 8) as *i64 85 var k0: i64 = 0 86 while k0 < NT * QKV { qkva[k0] = nx_le_read_u32(qkvd, k0 * 4); k0 = k0 + 1 } 87 88 var t: i64 = 0 89 while t < NT { 90 var hd: i64 = 0 91 while hd < H { 92 zsd_qknorm_rope(qkva, t * QKV + hd * HD, qn, cosr, sinr, t, eps) 93 zsd_qknorm_rope(qkva, t * QKV + D + hd * HD, kn, cosr, sinr, t, eps) 94 hd = hd + 1 95 } 96 t = t + 1 97 } 98 99 let ao: *i64 = sys_mmap(NT * D * 8) as *i64 100 let sc: *i64 = sys_mmap(NT * 8) as *i64 101 var hd2: i64 = 0 102 while hd2 < H { 103 var t1: i64 = 0 104 while t1 < NT { 105 var mx: i64 = 0 106 var t2: i64 = 0 107 while t2 < NT { 108 var s: i64 = 0 109 var d: i64 = 0 110 while d < HD { s = nx_f32_add(s, nx_f32_mul(qkva[t1 * QKV + hd2 * HD + d], qkva[t2 * QKV + D + hd2 * HD + d])); d = d + 1 } 111 s = nx_f32_mul(s, scale) 112 sc[t2] = s 113 if t2 == 0 { mx = s } else { if nx_f32_lt(mx, s) == 1 { mx = s } } 114 t2 = t2 + 1 115 } 116 var sum: i64 = 0 117 t2 = 0 118 while t2 < NT { let e: i64 = nx_f32_exp(nx_f32_sub(sc[t2], mx)); sc[t2] = e; sum = nx_f32_add(sum, e); t2 = t2 + 1 } 119 var d2: i64 = 0 120 while d2 < HD { 121 var a: i64 = 0 122 t2 = 0 123 while t2 < NT { a = nx_f32_add(a, nx_f32_mul(nx_f32_div(sc[t2], sum), qkva[t2 * QKV + D + D + hd2 * HD + d2])); t2 = t2 + 1 } 124 ao[t1 * D + hd2 * HD + d2] = a 125 d2 = d2 + 1 126 } 127 t1 = t1 + 1 128 } 129 hd2 = hd2 + 1 130 } 131 132 let tolc: i64 = nx_f32_div(nx_i32_to_f32(3), nx_i32_to_f32(100)) // 0.03 133 var fails: i64 = 0 134 var first_bad: i64 = 0 - 1 135 var kk: i64 = 0 136 while kk < NT * D { 137 let g: i64 = nx_le_read_u32(aog, kk * 4) 138 var thr: i64 = tolc 139 let ag: i64 = g & 0x7FFFFFFF 140 if nx_f32_lt(thr, nx_f32_mul(tolc, ag)) == 1 { thr = nx_f32_mul(tolc, ag) } 141 if (nx_f32_sub(ao[kk], g) & 0x7FFFFFFF) >= thr { fails = fails + 1; if first_bad < 0 { first_bad = kk } } 142 kk = kk + 1 143 } 144 145 let ofd: i64 = sys_openat_wr("/tmp/zsd.txt" as *u8, 0x1a4) 146 if ofd >= 0 { 147 let dec: *u8 = sys_mmap(32) 148 sys_write(ofd, "fails=" as *u8, 6) 149 let n1: i64 = nx_strconv_format_i64(fails, dec) 150 sys_write(ofd, dec, n1) 151 sys_write(ofd, " first_bad=" as *u8, 11) 152 let n2: i64 = nx_strconv_format_i64(first_bad, dec) 153 sys_write(ofd, dec, n2) 154 sys_write(ofd, "\n" as *u8, 1) 155 sys_close(ofd) 156 } 157 if fails > 0 { return 20 } 158 return 0 159}