nx_zimage_ffn_swiglu_verify.nx source
↩ module page · 93 lines · 3627 B
1// nx_zimage_ffn_swiglu_verify.nx -- SOVEREIGN Z-Image FFN sub-piece b: SwiGLU gate, verified vs oracle.
2// ffh = silu(h_ff @ w1.T) * (h_ff @ w3.T) ([4,3840]@[10240,3840]->[4,10240], gated). silu(x)=x/(1+e^-x).
3// Local accumulators (no scratch). Verifies vs dumped ffh.
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_10240: i64 = 10240
14
15func zsw_load(name: *u8, nl: i64, n_floats: i64) -> *u8 {
16 let base: *u8 = "/mnt/c/Users/elder/AppData/Local/Temp/claude/C--Users-elder/7be78b15-304c-449e-afe8-4d5bd7ddaa9c/scratchpad/zblk/" as *u8
17 let path: *u8 = sys_mmap(256)
18 var p: i64 = 0
19 var i: i64 = 0
20 while base[i] != 0 { path[p] = base[i]; p = p + 1; i = i + 1 }
21 i = 0
22 while i < nl { path[p] = name[i]; p = p + 1; i = i + 1 }
23 path[p] = 0x2E; p = p + 1
24 path[p] = 0x66; p = p + 1
25 path[p] = 0x33; p = p + 1
26 path[p] = 0x32; p = p + 1
27 path[p] = 0
28 let fd: i64 = sys_openat_rd(path)
29 if fd < 0 { return 0 as *u8 }
30 let bytes: i64 = n_floats * 4
31 let buf: *u8 = sys_mmap(bytes + 64)
32 var tot: i64 = 0
33 var go: i64 = 1
34 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 } } }
35 sys_close(fd)
36 return buf
37}
38
39func main() -> i64 {
40 let D: i64 = K_MAGIC_3840
41 let FFN: i64 = K_MAGIC_10240
42 let NT: i64 = 4
43 let h: *u8 = zsw_load("h_ff" as *u8, 4, NT * D)
44 let w1: *u8 = zsw_load("w1" as *u8, 2, FFN * D)
45 let w3: *u8 = zsw_load("w3" as *u8, 2, FFN * D)
46 let fg: *u8 = zsw_load("ffh" as *u8, 3, NT * FFN)
47 if (h as i64) == 0 { return 30 }
48 if (w1 as i64) == 0 { return 32 }
49 if (w3 as i64) == 0 { return 33 }
50 if (fg as i64) == 0 { return 31 }
51
52 let one: i64 = nx_i32_to_f32(1)
53 let tolc: i64 = nx_f32_div(nx_i32_to_f32(3), nx_i32_to_f32(100))
54 var fails: i64 = 0
55 var first_bad: i64 = 0 - 1
56
57 var t: i64 = 0
58 while t < NT {
59 let tb: i64 = t * D * 4
60 var j: i64 = 0
61 while j < FFN {
62 let rowb: i64 = j * D * 4
63 var gate: i64 = 0
64 var up: i64 = 0
65 var i: i64 = 0
66 while i < D {
67 let hv: i64 = nx_le_read_u32(h, tb + i * 4)
68 gate = nx_f32_add(gate, nx_f32_mul(hv, nx_le_read_u32(w1, rowb + i * 4)))
69 up = nx_f32_add(up, nx_f32_mul(hv, nx_le_read_u32(w3, rowb + i * 4)))
70 i = i + 1
71 }
72 let s: i64 = nx_f32_div(gate, nx_f32_add(one, nx_f32_exp(nx_f32_neg(gate))))
73 let ffhv: i64 = nx_f32_mul(s, up)
74 let g: i64 = nx_le_read_u32(fg, (t * FFN + j) * 4)
75 var thr: i64 = tolc
76 let ag: i64 = g & 0x7FFFFFFF
77 if nx_f32_lt(thr, nx_f32_mul(tolc, ag)) == 1 { thr = nx_f32_mul(tolc, ag) }
78 if (nx_f32_sub(ffhv, g) & 0x7FFFFFFF) >= thr { fails = fails + 1; if first_bad < 0 { first_bad = t * FFN + j } }
79 j = j + 1
80 }
81 t = t + 1
82 }
83
84 let ofd: i64 = sys_openat_wr("/tmp/zsw.txt" as *u8, 0x1a4)
85 if ofd >= 0 {
86 let dec: *u8 = sys_mmap(32)
87 sys_write(ofd, "fails=" as *u8, 6); let n1: i64 = nx_strconv_format_i64(fails, dec); sys_write(ofd, dec, n1)
88 sys_write(ofd, " first_bad=" as *u8, 11); let n2: i64 = nx_strconv_format_i64(first_bad, dec); sys_write(ofd, dec, n2)
89 sys_write(ofd, "\n" as *u8, 1); sys_close(ofd)
90 }
91 if fails > 0 { return 20 }
92 return 0
93}