nx_zimage_sdpaonly_verify.nx source
↩ module page · 131 lines · 5644 B
1// nx_zimage_sdpaonly_verify.nx -- SOVEREIGN Z-Image attention sub-piece 2c: SDPA, verified vs oracle.
2// Rewritten with NT=4 UNROLLED locals (no sc[] scratch, no nested t2 loop) to dodge the codegen crash.
3// Reads oracle q_roped/k_roped + qkv(V) -> per-head softmax(q.k/sqrt128) v -> ao_pre.
4// license_tier: ORIGINAL
5import "nx_syscalls.nx"
6import "nx_le.nx"
7import "nx_f32.nx"
8import "nx_f32_div.nx"
9import "nx_f32_cvt.nx"
10import "nx_f32_exp.nx"
11import "nx_strconv.nx"
12const K_MAGIC_3840: i64 = 3840
13const K_MAGIC_11520: i64 = 11520
14const K_MAGIC_7680: i64 = 7680
15
16func zso_load(name: *u8, nl: i64, n_floats: i64) -> *u8 {
17 let base: *u8 = "/mnt/c/Users/elder/AppData/Local/Temp/claude/C--Users-elder/7be78b15-304c-449e-afe8-4d5bd7ddaa9c/scratchpad/zblk/" as *u8
18 let path: *u8 = sys_mmap(256)
19 var p: i64 = 0
20 var i: i64 = 0
21 while base[i] != 0 { path[p] = base[i]; p = p + 1; i = i + 1 }
22 i = 0
23 while i < nl { path[p] = name[i]; p = p + 1; i = i + 1 }
24 path[p] = 0x2E; p = p + 1
25 path[p] = 0x66; p = p + 1
26 path[p] = 0x33; p = p + 1
27 path[p] = 0x32; p = p + 1
28 path[p] = 0
29 let fd: i64 = sys_openat_rd(path)
30 if fd < 0 { return 0 as *u8 }
31 let bytes: i64 = n_floats * 4
32 let buf: *u8 = sys_mmap(bytes + 64)
33 var tot: i64 = 0
34 var go: i64 = 1
35 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 } } }
36 sys_close(fd)
37 return buf
38}
39
40func main() -> i64 {
41 let D: i64 = K_MAGIC_3840
42 let H: i64 = 30
43 let HD: i64 = 128
44 let NT: i64 = 4
45 let QKV: i64 = K_MAGIC_11520
46 let qr: *u8 = zso_load("q_roped" as *u8, 7, NT * D)
47 let kr: *u8 = zso_load("k_roped" as *u8, 7, NT * D)
48 let qkvd: *u8 = zso_load("qkv" as *u8, 3, NT * QKV)
49 let aog: *u8 = zso_load("ao_pre" as *u8, 6, NT * D)
50 if (qr as i64) == 0 { return 30 }
51 if (kr as i64) == 0 { return 32 }
52 if (qkvd as i64) == 0 { return 33 }
53 if (aog as i64) == 0 { return 31 }
54
55 let scale: i64 = nx_f32_div(nx_i32_to_f32(1), nx_f32_sqrt(nx_i32_to_f32(128)))
56 let tolc: i64 = nx_f32_div(nx_i32_to_f32(3), nx_i32_to_f32(100))
57 let vbase: i64 = D + D // V region offset = K_MAGIC_7680
58 var fails: i64 = 0
59 var first_bad: i64 = 0 - 1
60
61 var hd: i64 = 0
62 while hd < H {
63 var t1: i64 = 0
64 while t1 < NT {
65 let qb: i64 = (t1 * D + hd * HD) * 4
66 let k0b: i64 = (hd * HD) * 4
67 let k1b: i64 = (D + hd * HD) * 4
68 let k2b: i64 = (2 * D + hd * HD) * 4
69 let k3b: i64 = (3 * D + hd * HD) * 4
70 var s0: i64 = 0
71 var s1: i64 = 0
72 var s2: i64 = 0
73 var s3: i64 = 0
74 var d: i64 = 0
75 while d < HD {
76 let qd: i64 = nx_le_read_u32(qr, qb + d * 4)
77 s0 = nx_f32_add(s0, nx_f32_mul(qd, nx_le_read_u32(kr, k0b + d * 4)))
78 s1 = nx_f32_add(s1, nx_f32_mul(qd, nx_le_read_u32(kr, k1b + d * 4)))
79 s2 = nx_f32_add(s2, nx_f32_mul(qd, nx_le_read_u32(kr, k2b + d * 4)))
80 s3 = nx_f32_add(s3, nx_f32_mul(qd, nx_le_read_u32(kr, k3b + d * 4)))
81 d = d + 1
82 }
83 s0 = nx_f32_mul(s0, scale)
84 s1 = nx_f32_mul(s1, scale)
85 s2 = nx_f32_mul(s2, scale)
86 s3 = nx_f32_mul(s3, scale)
87 var mx: i64 = s0
88 if nx_f32_lt(mx, s1) == 1 { mx = s1 }
89 if nx_f32_lt(mx, s2) == 1 { mx = s2 }
90 if nx_f32_lt(mx, s3) == 1 { mx = s3 }
91 let e0: i64 = nx_f32_exp(nx_f32_sub(s0, mx))
92 let e1: i64 = nx_f32_exp(nx_f32_sub(s1, mx))
93 let e2: i64 = nx_f32_exp(nx_f32_sub(s2, mx))
94 let e3: i64 = nx_f32_exp(nx_f32_sub(s3, mx))
95 let sum: i64 = nx_f32_add(nx_f32_add(e0, e1), nx_f32_add(e2, e3))
96 let a0: i64 = nx_f32_div(e0, sum)
97 let a1: i64 = nx_f32_div(e1, sum)
98 let a2: i64 = nx_f32_div(e2, sum)
99 let a3: i64 = nx_f32_div(e3, sum)
100 let v0b: i64 = (vbase + hd * HD) * 4
101 let v1b: i64 = (QKV + vbase + hd * HD) * 4
102 let v2b: i64 = (2 * QKV + vbase + hd * HD) * 4
103 let v3b: i64 = (3 * QKV + vbase + hd * HD) * 4
104 var d2: i64 = 0
105 while d2 < HD {
106 var a: i64 = nx_f32_mul(a0, nx_le_read_u32(qkvd, v0b + d2 * 4))
107 a = nx_f32_add(a, nx_f32_mul(a1, nx_le_read_u32(qkvd, v1b + d2 * 4)))
108 a = nx_f32_add(a, nx_f32_mul(a2, nx_le_read_u32(qkvd, v2b + d2 * 4)))
109 a = nx_f32_add(a, nx_f32_mul(a3, nx_le_read_u32(qkvd, v3b + d2 * 4)))
110 let g: i64 = nx_le_read_u32(aog, (t1 * D + hd * HD + d2) * 4)
111 var thr: i64 = tolc
112 let ag: i64 = g & 0x7FFFFFFF
113 if nx_f32_lt(thr, nx_f32_mul(tolc, ag)) == 1 { thr = nx_f32_mul(tolc, ag) }
114 if (nx_f32_sub(a, g) & 0x7FFFFFFF) >= thr { fails = fails + 1; if first_bad < 0 { first_bad = t1 * D + hd * HD + d2 } }
115 d2 = d2 + 1
116 }
117 t1 = t1 + 1
118 }
119 hd = hd + 1
120 }
121
122 let ofd: i64 = sys_openat_wr("/tmp/zso.txt" as *u8, 0x1a4)
123 if ofd >= 0 {
124 let dec: *u8 = sys_mmap(32)
125 sys_write(ofd, "fails=" as *u8, 6); let n1: i64 = nx_strconv_format_i64(fails, dec); sys_write(ofd, dec, n1)
126 sys_write(ofd, " first_bad=" as *u8, 11); let n2: i64 = nx_strconv_format_i64(first_bad, dec); sys_write(ofd, dec, n2)
127 sys_write(ofd, "\n" as *u8, 1); sys_close(ofd)
128 }
129 if fails > 0 { return 20 }
130 return 0
131}