code wiki / (root) / nx_zimage_attn_verify.nx

nx_zimage_attn_verify.nx source

↩ module page · 205 lines · 7888 B

1// nx_zimage_attn_verify.nx -- SOVEREIGN Z-Image block piece #2: JointAttention, verified vs the oracle. 2// 3// z_image.hpp JointAttention: qkv = Linear(h); split Q|K|V each 30 heads x 128; per-head QK-RMSNorm(qn/kn); 4// axial RoPE; per-head SDPA (softmax(q.k^T / sqrt(128)) v); out = Linear. Reads oracle dumps h_attn (input), 5// qkv_w/out_w/qn/kn (weights), cosr/sinr (rope), and golden ao_out; computes sovereignly; verifies to tol. 6// Reference dumps are oracle-only; the ORGAN is pure Nishi (nx_f32_*, raw syscalls, static ELF). 7// license_tier: ORIGINAL 8import "nx_syscalls.nx" 9import "nx_le.nx" 10import "nx_f32.nx" 11import "nx_f32_div.nx" 12import "nx_f32_cvt.nx" 13import "nx_f32_exp.nx" 14import "nx_strconv.nx" 15const K_MAGIC_3840: i64 = 3840 16const K_MAGIC_11520: i64 = 11520 17const K_MAGIC_100000: i64 = 100000 18const K_MAGIC_2000: i64 = 2000 19 20func zav_load(name: *u8, nl: i64, n_floats: i64) -> *u8 { 21 let base: *u8 = "/mnt/c/Users/elder/AppData/Local/Temp/claude/C--Users-elder/7be78b15-304c-449e-afe8-4d5bd7ddaa9c/scratchpad/zblk/" as *u8 22 let path: *u8 = sys_mmap(256) 23 var p: i64 = 0 24 var i: i64 = 0 25 while base[i] != 0 { path[p] = base[i]; p = p + 1; i = i + 1 } 26 i = 0 27 while i < nl { path[p] = name[i]; p = p + 1; i = i + 1 } 28 path[p] = 0x2E; p = p + 1 29 path[p] = 0x66; p = p + 1 30 path[p] = 0x33; p = p + 1 31 path[p] = 0x32; p = p + 1 32 path[p] = 0 33 let fd: i64 = sys_openat_rd(path) 34 if fd < 0 { return 0 as *u8 } 35 let bytes: i64 = n_floats * 4 36 let buf: *u8 = sys_mmap(bytes + 64) 37 var tot: i64 = 0 38 var go: i64 = 1 39 while go == 1 { 40 let r: i64 = sys_read(fd, ((buf as i64) + tot) as *u8, bytes - tot) 41 if r <= 0 { go = 0 } else { tot = tot + r; if tot >= bytes { go = 0 } } 42 } 43 sys_close(fd) 44 return buf 45} 46 47// per-head RMSNorm(128) then axial-RoPE, in place on the i64 f32-bit array at qkv[off .. off+128] 48func zav_qknorm_rope(qkv: *i64, off: i64, w: *u8, cosr: *u8, sinr: *u8, t: i64, eps: i64) -> i64 { 49 var ss: i64 = 0 50 var d: i64 = 0 51 while d < 128 { let v: i64 = qkv[off + d]; ss = nx_f32_add(ss, nx_f32_mul(v, v)); d = d + 1 } 52 let ms: i64 = nx_f32_div(ss, nx_i32_to_f32(128)) 53 let rms: i64 = nx_f32_sqrt(nx_f32_add(ms, eps)) 54 d = 0 55 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 } 56 var j: i64 = 0 57 while j < 64 { 58 let c: i64 = nx_le_read_u32(cosr, (t * 64 + j) * 4) 59 let s: i64 = nx_le_read_u32(sinr, (t * 64 + j) * 4) 60 let x0: i64 = qkv[off + 2 * j] 61 let x1: i64 = qkv[off + 2 * j + 1] 62 qkv[off + 2 * j] = nx_f32_sub(nx_f32_mul(x0, c), nx_f32_mul(x1, s)) 63 qkv[off + 2 * j + 1] = nx_f32_add(nx_f32_mul(x0, s), nx_f32_mul(x1, c)) 64 j = j + 1 65 } 66 return 0 67} 68 69func zstage(n: i64) -> i64 { 70 let fd: i64 = sys_openat_wr("/tmp/zstage.txt" as *u8, 0x1a4) 71 if fd >= 0 { let dec: *u8 = sys_mmap(16); let nd: i64 = nx_strconv_format_i64(n, dec); sys_write(fd, dec, nd); sys_write(fd, "\n" as *u8, 1); sys_close(fd) } 72 return 0 73} 74 75func main() -> i64 { 76 let D: i64 = K_MAGIC_3840 77 let H: i64 = 30 78 let HD: i64 = 128 79 let NT: i64 = 4 80 let QKV: i64 = K_MAGIC_11520 81 let ha: *u8 = zav_load("h_attn" as *u8, 6, NT * D) 82 let qkvw: *u8 = zav_load("qkv_w" as *u8, 5, QKV * D) 83 let outw: *u8 = zav_load("out_w" as *u8, 5, D * D) 84 let qn: *u8 = zav_load("qn" as *u8, 2, HD) 85 let kn: *u8 = zav_load("kn" as *u8, 2, HD) 86 let cosr: *u8 = zav_load("cosr" as *u8, 4, NT * 64) 87 let sinr: *u8 = zav_load("sinr" as *u8, 4, NT * 64) 88 let aog: *u8 = zav_load("ao_out" as *u8, 6, NT * D) 89 if (ha as i64) == 0 { return 30 } 90 if (qkvw as i64) == 0 { return 31 } 91 if (aog as i64) == 0 { return 32 } 92 if (outw as i64) == 0 { return 33 } 93 if (qn as i64) == 0 { return 34 } 94 if (cosr as i64) == 0 { return 35 } 95 zstage(1) 96 97 let eps: i64 = nx_f32_div(nx_i32_to_f32(1), nx_i32_to_f32(K_MAGIC_100000)) 98 let scale: i64 = nx_f32_div(nx_i32_to_f32(1), nx_f32_sqrt(nx_i32_to_f32(128))) 99 100 // qkv = h_attn @ qkv_w^T -> [NT, 11520] f32-bit 101 let qkv: *i64 = sys_mmap(NT * QKV * 8) as *i64 102 zstage((qkvw as i64) >> 28) // marker: high bits of the qkvw pointer (valid mmap => large) 103 var t: i64 = 0 104 while t < NT { 105 var o: i64 = 0 106 while o < QKV { 107 var acc: i64 = 0 108 var i: i64 = 0 109 let hb: i64 = t * D * 4 110 let wb: i64 = o * D * 4 111 while i < D { acc = nx_f32_add(acc, nx_f32_mul(nx_le_read_u32(ha, hb + i * 4), nx_le_read_u32(qkvw, wb + i * 4))); i = i + 1 } 112 qkv[t * QKV + o] = acc 113 if (o - (o / K_MAGIC_2000) * K_MAGIC_2000) == 0 { zstage(t * K_MAGIC_100000 + o) } 114 o = o + 1 115 } 116 zstage(300 + t) 117 t = t + 1 118 } 119 120 zstage(2) 121 // qk-norm + rope on Q (offset 0) and K (offset 3840), per head 122 t = 0 123 while t < NT { 124 var hd: i64 = 0 125 while hd < H { 126 zav_qknorm_rope(qkv, t * QKV + hd * HD, qn, cosr, sinr, t, eps) 127 zav_qknorm_rope(qkv, t * QKV + D + hd * HD, kn, cosr, sinr, t, eps) 128 hd = hd + 1 129 } 130 t = t + 1 131 } 132 133 zstage(3) 134 // SDPA per head -> ao [NT, D] 135 let ao: *i64 = sys_mmap(NT * D * 8) as *i64 136 let sc: *i64 = sys_mmap(NT * 8) as *i64 137 var hd2: i64 = 0 138 while hd2 < H { 139 var t1: i64 = 0 140 while t1 < NT { 141 var mx: i64 = 0 142 var t2: i64 = 0 143 while t2 < NT { 144 var s: i64 = 0 145 var d: i64 = 0 146 while d < HD { s = nx_f32_add(s, nx_f32_mul(qkv[t1 * QKV + hd2 * HD + d], qkv[t2 * QKV + D + hd2 * HD + d])); d = d + 1 } 147 s = nx_f32_mul(s, scale) 148 sc[t2] = s 149 if t2 == 0 { mx = s } else { if nx_f32_lt(mx, s) == 1 { mx = s } } 150 t2 = t2 + 1 151 } 152 var sum: i64 = 0 153 t2 = 0 154 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 } 155 var d2: i64 = 0 156 while d2 < HD { 157 var a: i64 = 0 158 t2 = 0 159 while t2 < NT { a = nx_f32_add(a, nx_f32_mul(nx_f32_div(sc[t2], sum), qkv[t2 * QKV + D + D + hd2 * HD + d2])); t2 = t2 + 1 } 160 ao[t1 * D + hd2 * HD + d2] = a 161 d2 = d2 + 1 162 } 163 t1 = t1 + 1 164 } 165 hd2 = hd2 + 1 166 } 167 168 zstage(4) 169 // out = ao @ out_w^T ; verify vs ao_out golden 170 let tolc: i64 = nx_f32_div(nx_i32_to_f32(5), nx_i32_to_f32(100)) // 0.05 171 var fails: i64 = 0 172 var first_bad: i64 = 0 - 1 173 t = 0 174 while t < NT { 175 var o: i64 = 0 176 while o < D { 177 var acc: i64 = 0 178 var i: i64 = 0 179 let wb: i64 = o * D * 4 180 while i < D { acc = nx_f32_add(acc, nx_f32_mul(ao[t * D + i], nx_le_read_u32(outw, wb + i * 4))); i = i + 1 } 181 let g: i64 = nx_le_read_u32(aog, (t * D + o) * 4) 182 var thr: i64 = tolc 183 let ag: i64 = g & 0x7FFFFFFF 184 if nx_f32_lt(thr, nx_f32_mul(tolc, ag)) == 1 { thr = nx_f32_mul(tolc, ag) } 185 if (nx_f32_sub(acc, g) & 0x7FFFFFFF) >= thr { fails = fails + 1; if first_bad < 0 { first_bad = t * D + o } } 186 o = o + 1 187 } 188 t = t + 1 189 } 190 191 let ofd: i64 = sys_openat_wr("/tmp/zattn.txt" as *u8, 0x1a4) 192 if ofd >= 0 { 193 let dec: *u8 = sys_mmap(32) 194 sys_write(ofd, "fails=" as *u8, 6) 195 let n1: i64 = nx_strconv_format_i64(fails, dec) 196 sys_write(ofd, dec, n1) 197 sys_write(ofd, " first_bad=" as *u8, 11) 198 let n2: i64 = nx_strconv_format_i64(first_bad, dec) 199 sys_write(ofd, dec, n2) 200 sys_write(ofd, "\n" as *u8, 1) 201 sys_close(ofd) 202 } 203 if fails > 0 { return 20 } 204 return 0 205}