code wiki / (root) / nx_zimage_ffn_swiglu_verify.nx

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}