code wiki / (root) / nx_zimage_ffn_down_verify.nx

nx_zimage_ffn_down_verify.nx source

↩ module page · 110 lines · 4455 B

1// nx_zimage_ffn_down_verify.nx -- SOVEREIGN Z-Image FFN sub-piece c (final): down-proj + norm + gate + residual. 2// ff = ffh @ w2.T ; ff = rmsnorm(ff, fn2) * tanh(gl) ; out = ff + x_after_attn. Verifies vs oracle `out`. 3// tanh built from exp (sign-stable). ff[] is a per-token local scratch. 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_f32_exp.nx" 10import "nx_strconv.nx" 11const K_MAGIC_3840: i64 = 3840 12const K_MAGIC_10240: i64 = 10240 13const K_MAGIC_100000: i64 = 100000 14 15func zfd_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 39// tanh(x) via sign-stable exp of a non-positive argument 40func zfd_tanh(x: i64) -> i64 { 41 let one: i64 = nx_i32_to_f32(1) 42 let two_x: i64 = nx_f32_add(x, x) 43 if (x & 0x80000000) == 0 { 44 let z: i64 = nx_f32_exp(nx_f32_neg(two_x)) 45 return nx_f32_div(nx_f32_sub(one, z), nx_f32_add(one, z)) 46 } 47 let z2: i64 = nx_f32_exp(two_x) 48 return nx_f32_div(nx_f32_sub(z2, one), nx_f32_add(z2, one)) 49} 50 51func main() -> i64 { 52 let D: i64 = K_MAGIC_3840 53 let FFN: i64 = K_MAGIC_10240 54 let NT: i64 = 4 55 let ffh: *u8 = zfd_load("ffh" as *u8, 3, NT * FFN) 56 let w2: *u8 = zfd_load("w2" as *u8, 2, D * FFN) 57 let fn2: *u8 = zfd_load("fn2" as *u8, 3, D) 58 let gl: *u8 = zfd_load("gl" as *u8, 2, D) 59 let xaa: *u8 = zfd_load("x_after_attn" as *u8, 12, NT * D) 60 let outg: *u8 = zfd_load("out" as *u8, 3, NT * D) 61 if (ffh as i64) == 0 { return 30 } 62 if (w2 as i64) == 0 { return 32 } 63 if (outg as i64) == 0 { return 31 } 64 65 let eps: i64 = nx_f32_div(nx_i32_to_f32(1), nx_i32_to_f32(K_MAGIC_100000)) 66 let tolc: i64 = nx_f32_div(nx_i32_to_f32(3), nx_i32_to_f32(100)) 67 let ff: *i64 = sys_mmap(D * 8) as *i64 68 var fails: i64 = 0 69 var first_bad: i64 = 0 - 1 70 71 var t: i64 = 0 72 while t < NT { 73 let tb: i64 = t * FFN * 4 74 var o: i64 = 0 75 while o < D { 76 let rowb: i64 = o * FFN * 4 77 var sum: i64 = 0 78 var j: i64 = 0 79 while j < FFN { sum = nx_f32_add(sum, nx_f32_mul(nx_le_read_u32(ffh, tb + j * 4), nx_le_read_u32(w2, rowb + j * 4))); j = j + 1 } 80 ff[o] = sum 81 o = o + 1 82 } 83 var ss: i64 = 0 84 o = 0 85 while o < D { ss = nx_f32_add(ss, nx_f32_mul(ff[o], ff[o])); o = o + 1 } 86 let rms: i64 = nx_f32_sqrt(nx_f32_add(nx_f32_div(ss, nx_i32_to_f32(D)), eps)) 87 o = 0 88 while o < D { 89 let ffn: i64 = nx_f32_mul(nx_f32_mul(nx_f32_div(ff[o], rms), nx_le_read_u32(fn2, o * 4)), zfd_tanh(nx_le_read_u32(gl, o * 4))) 90 let outv: i64 = nx_f32_add(ffn, nx_le_read_u32(xaa, (t * D + o) * 4)) 91 let g: i64 = nx_le_read_u32(outg, (t * D + o) * 4) 92 var thr: i64 = tolc 93 let ag: i64 = g & 0x7FFFFFFF 94 if nx_f32_lt(thr, nx_f32_mul(tolc, ag)) == 1 { thr = nx_f32_mul(tolc, ag) } 95 if (nx_f32_sub(outv, g) & 0x7FFFFFFF) >= thr { fails = fails + 1; if first_bad < 0 { first_bad = t * D + o } } 96 o = o + 1 97 } 98 t = t + 1 99 } 100 101 let ofd: i64 = sys_openat_wr("/tmp/zfd.txt" as *u8, 0x1a4) 102 if ofd >= 0 { 103 let dec: *u8 = sys_mmap(32) 104 sys_write(ofd, "fails=" as *u8, 6); let n1: i64 = nx_strconv_format_i64(fails, dec); sys_write(ofd, dec, n1) 105 sys_write(ofd, " first_bad=" as *u8, 11); let n2: i64 = nx_strconv_format_i64(first_bad, dec); sys_write(ofd, dec, n2) 106 sys_write(ofd, "\n" as *u8, 1); sys_close(ofd) 107 } 108 if fails > 0 { return 20 } 109 return 0 110}