nx_zimage_qknorm_verify.nx source
↩ module page · 125 lines · 4838 B
1// nx_zimage_qknorm_verify.nx -- SOVEREIGN Z-Image attention sub-piece 2b: QK-RMSNorm + RoPE, verified.
2// Isolates norm+rope from SDPA: takes dumped qkv, norms+ropes Q,K, verifies vs oracle q_roped/k_roped.
3// license_tier: ORIGINAL
4import "nx_syscalls.nx"
5import "nx_le.nx"
6import "nx_f32.nx"
7import "nx_f32_div.nx"
8import "nx_f32_cvt.nx"
9import "nx_strconv.nx"
10const K_MAGIC_3840: i64 = 3840
11const K_MAGIC_11520: i64 = 11520
12const K_MAGIC_100000: i64 = 100000
13
14func zqk_load(name: *u8, nl: i64, n_floats: i64) -> *u8 {
15 let base: *u8 = "/mnt/c/Users/elder/AppData/Local/Temp/claude/C--Users-elder/7be78b15-304c-449e-afe8-4d5bd7ddaa9c/scratchpad/zblk/" as *u8
16 let path: *u8 = sys_mmap(256)
17 var p: i64 = 0
18 var i: i64 = 0
19 while base[i] != 0 { path[p] = base[i]; p = p + 1; i = i + 1 }
20 i = 0
21 while i < nl { path[p] = name[i]; p = p + 1; i = i + 1 }
22 path[p] = 0x2E; p = p + 1
23 path[p] = 0x66; p = p + 1
24 path[p] = 0x33; p = p + 1
25 path[p] = 0x32; p = p + 1
26 path[p] = 0
27 let fd: i64 = sys_openat_rd(path)
28 if fd < 0 { return 0 as *u8 }
29 let bytes: i64 = n_floats * 4
30 let buf: *u8 = sys_mmap(bytes + 64)
31 var tot: i64 = 0
32 var go: i64 = 1
33 while go == 1 { let r: i64 = sys_read(fd, ((buf as i64) + tot) as *u8, bytes - tot); if r <= 0 { go = 0 } else { tot = tot + r; if tot >= bytes { go = 0 } } }
34 sys_close(fd)
35 return buf
36}
37
38func zqk_norm_rope(qkv: *i64, off: i64, w: *u8, cosr: *u8, sinr: *u8, t: i64, eps: i64) -> i64 {
39 var ss: i64 = 0
40 var d: i64 = 0
41 while d < 128 { let v: i64 = qkv[off + d]; ss = nx_f32_add(ss, nx_f32_mul(v, v)); d = d + 1 }
42 let rms: i64 = nx_f32_sqrt(nx_f32_add(nx_f32_div(ss, nx_i32_to_f32(128)), eps))
43 d = 0
44 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 }
45 var j: i64 = 0
46 while j < 64 {
47 let c: i64 = nx_le_read_u32(cosr, (t * 64 + j) * 4)
48 let s: i64 = nx_le_read_u32(sinr, (t * 64 + j) * 4)
49 let x0: i64 = qkv[off + 2 * j]
50 let x1: i64 = qkv[off + 2 * j + 1]
51 qkv[off + 2 * j] = nx_f32_sub(nx_f32_mul(x0, c), nx_f32_mul(x1, s))
52 qkv[off + 2 * j + 1] = nx_f32_add(nx_f32_mul(x0, s), nx_f32_mul(x1, c))
53 j = j + 1
54 }
55 return 0
56}
57
58func main() -> i64 {
59 let D: i64 = K_MAGIC_3840
60 let H: i64 = 30
61 let HD: i64 = 128
62 let NT: i64 = 4
63 let QKV: i64 = K_MAGIC_11520
64 let qkvd: *u8 = zqk_load("qkv" as *u8, 3, NT * QKV)
65 let qn: *u8 = zqk_load("qn" as *u8, 2, HD)
66 let kn: *u8 = zqk_load("kn" as *u8, 2, HD)
67 let cosr: *u8 = zqk_load("cosr" as *u8, 4, NT * 64)
68 let sinr: *u8 = zqk_load("sinr" as *u8, 4, NT * 64)
69 let qg: *u8 = zqk_load("q_roped" as *u8, 7, NT * D)
70 let kg: *u8 = zqk_load("k_roped" as *u8, 7, NT * D)
71 if (qkvd as i64) == 0 { return 30 }
72 if (qg as i64) == 0 { return 31 }
73
74 let eps: i64 = nx_f32_div(nx_i32_to_f32(1), nx_i32_to_f32(K_MAGIC_100000))
75 let qkva: *i64 = sys_mmap(NT * QKV * 8) as *i64
76 var k0: i64 = 0
77 while k0 < NT * QKV { qkva[k0] = nx_le_read_u32(qkvd, k0 * 4); k0 = k0 + 1 }
78
79 var t: i64 = 0
80 while t < NT {
81 var hd: i64 = 0
82 while hd < H {
83 zqk_norm_rope(qkva, t * QKV + hd * HD, qn, cosr, sinr, t, eps)
84 zqk_norm_rope(qkva, t * QKV + D + hd * HD, kn, cosr, sinr, t, eps)
85 hd = hd + 1
86 }
87 t = t + 1
88 }
89
90 let tolc: i64 = nx_f32_div(nx_i32_to_f32(2), nx_i32_to_f32(100))
91 var fq: i64 = 0
92 var fk: i64 = 0
93 var fbq: i64 = 0 - 1
94 t = 0
95 while t < NT {
96 var hd: i64 = 0
97 while hd < H {
98 var d: i64 = 0
99 while d < HD {
100 let gi: i64 = (t * D + hd * HD + d) * 4
101 let qc: i64 = qkva[t * QKV + hd * HD + d]
102 let qgv: i64 = nx_le_read_u32(qg, gi)
103 if (nx_f32_sub(qc, qgv) & 0x7FFFFFFF) >= tolc { fq = fq + 1; if fbq < 0 { fbq = t * D + hd * HD + d } }
104 let kc: i64 = qkva[t * QKV + D + hd * HD + d]
105 let kgv: i64 = nx_le_read_u32(kg, gi)
106 if (nx_f32_sub(kc, kgv) & 0x7FFFFFFF) >= tolc { fk = fk + 1 }
107 d = d + 1
108 }
109 hd = hd + 1
110 }
111 t = t + 1
112 }
113
114 let ofd: i64 = sys_openat_wr("/tmp/zqk.txt" as *u8, 0x1a4)
115 if ofd >= 0 {
116 let dec: *u8 = sys_mmap(32)
117 sys_write(ofd, "fq=" as *u8, 3); let n1: i64 = nx_strconv_format_i64(fq, dec); sys_write(ofd, dec, n1)
118 sys_write(ofd, " fk=" as *u8, 4); let n2: i64 = nx_strconv_format_i64(fk, dec); sys_write(ofd, dec, n2)
119 sys_write(ofd, " fbq=" as *u8, 5); let n3: i64 = nx_strconv_format_i64(fbq, dec); sys_write(ofd, dec, n3)
120 sys_write(ofd, "\n" as *u8, 1); sys_close(ofd)
121 }
122 if fq > 0 { return 20 }
123 if fk > 0 { return 21 }
124 return 0
125}