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}