code wiki / (root) / nx_ta_transformer_gradcheck_gate.nx

nx_ta_transformer_gradcheck_gate.nx source

↩ module page · 207 lines · 10564 B

1// nx_ta_transformer_gradcheck_gate.nx -- LIAR-KILL for the 9 f32 transformer ops added to nx_autograd_tensor 2// (2026-07-09, for the R3e neural reader). Each backward identity was HAND-PORTED from its Q16 nfa_* twin; 3// a transposed index or flipped sign would train the reader to garbage and waste the slow f32 run. This gate 4// verifies every one by FINITE DIFFERENCE: build op -> scalar loss (mse-vs-0, or the CE itself for softce) -> 5// analytic grad (ta_backward) vs numeric grad ((L(+h)-L(-h))/2h). A wrong identity => order-of-magnitude/sign 6// mismatch. Requires each op to show a NON-TRIVIAL gradient (|g|>=20 milli) so it cannot pass at 0==0. 7// PASS <op> for each of: matmul, matmul_nt, cmul, softmax_rows(causal), rope_tab, hadamard, silu, 8// rmsnorm_rows, softce_rows. 9// expect_exit: 0 license_tier: ORIGINAL Sovereign: nx_autograd_tensor + nx_syscalls. 10import "nx_autograd_tensor.nx" 11import "nx_syscalls.nx" 12 13func gt_puts(s: *u8) -> i64 { var n: i64=0; while s[n]!=(0 as u8){n=n+1} sys_write(1,s,n); return 0 } 14func gt_pn(v: i64) -> i64 { let b: *u8=sys_mmap(28); var x: i64=v; if x<0{b[0]=45;sys_write(1,b,1);x=0-x} if x==0{b[0]=48;sys_write(1,b,1);return 0} var d: i64=0; var y: i64=x; while y>0{d=d+1;y=y/10} var i: i64=d-1; y=x; while i>=0{b[i]=(48+(y%10)) as u8;y=y/10;i=i-1} sys_write(1,b,d); return 0 } 15 16func gt_abs(x: i64) -> i64 { if x < 0 { return 0 - x } return x } 17// finite-diff compare in milli: pass if |ana-num|*100 <= |ana|*20 + 400 (20% rel + 0.4 abs, generous for f32 18// finite diff but catches any sign/index bug). returns pass(0/1); *nt set to 1 if this cell is non-trivial. 19func gt_cell(Lp: i64, Lm: i64, h2: i64, ana: i64, nt: *i64) -> i64 { 20 let num: i64 = nx_f32_div(nx_f32_sub(Lp, Lm), h2) 21 let am: i64 = ta_f32_to_milli(ana) 22 let nm: i64 = ta_f32_to_milli(num) 23 gt_puts(" ana="); gt_pn(am); gt_puts(" num="); gt_pn(nm); gt_puts("\n" as *u8) 24 let aa: i64 = gt_abs(am) 25 if aa >= 20 { nt[0] = 1 } 26 let diff: i64 = gt_abs(am - nm) 27 if diff*100 <= aa*20 + 400 { return 1 } 28 return 0 29} 30 31// ---- per-op tape builders: return the scalar-root node. leaf 0 = the checked input (grad cells at arena 0). ---- 32func bld_matmul(tape: *i64, vals: *i64, st: *i64, A: *i64, B: *i64, Z: *i64) -> i64 { 33 st[0]=0; st[1]=0 34 let nA: i64 = ta_leaf(tape,vals,st,2,3,A,0) 35 let nB: i64 = ta_leaf(tape,vals,st,3,2,B,0) 36 let nC: i64 = ta_matmul(tape,vals,st,nA,nB) 37 let nZ: i64 = ta_leaf(tape,vals,st,2,2,Z,0) 38 return ta_mse(tape,vals,st,nC,nZ) 39} 40func bld_matmul_nt(tape: *i64, vals: *i64, st: *i64, A: *i64, B: *i64, Z: *i64) -> i64 { 41 st[0]=0; st[1]=0 42 let nA: i64 = ta_leaf(tape,vals,st,2,3,A,0) 43 let nB: i64 = ta_leaf(tape,vals,st,2,3,B,0) // B[p=2, k=3] -> S[2,2] 44 let nS: i64 = ta_matmul_nt(tape,vals,st,nA,nB) 45 let nZ: i64 = ta_leaf(tape,vals,st,2,2,Z,0) 46 return ta_mse(tape,vals,st,nS,nZ) 47} 48func bld_cmul(tape: *i64, vals: *i64, st: *i64, A: *i64, Z: *i64, c: i64) -> i64 { 49 st[0]=0; st[1]=0 50 let nA: i64 = ta_leaf(tape,vals,st,2,3,A,0) 51 let nY: i64 = ta_cmul(tape,vals,st,nA,c) 52 let nZ: i64 = ta_leaf(tape,vals,st,2,3,Z,0) 53 return ta_mse(tape,vals,st,nY,nZ) 54} 55func bld_smrows(tape: *i64, vals: *i64, st: *i64, A: *i64, Z: *i64) -> i64 { 56 st[0]=0; st[1]=0 57 let nA: i64 = ta_leaf(tape,vals,st,3,3,A,0) 58 let nY: i64 = ta_softmax_rows(tape,vals,st,nA,1) // causal 59 let nZ: i64 = ta_leaf(tape,vals,st,3,3,Z,0) 60 return ta_mse(tape,vals,st,nY,nZ) 61} 62func bld_rope(tape: *i64, vals: *i64, st: *i64, A: *i64, Z: *i64, tab: *i64) -> i64 { 63 st[0]=0; st[1]=0 64 let nA: i64 = ta_leaf(tape,vals,st,2,4,A,0) // T=2, hd=4 65 let nY: i64 = ta_rope_tab(tape,vals,st,nA,tab) 66 let nZ: i64 = ta_leaf(tape,vals,st,2,4,Z,0) 67 return ta_mse(tape,vals,st,nY,nZ) 68} 69func bld_had(tape: *i64, vals: *i64, st: *i64, A: *i64, B: *i64, Z: *i64) -> i64 { 70 st[0]=0; st[1]=0 71 let nA: i64 = ta_leaf(tape,vals,st,2,3,A,0) 72 let nB: i64 = ta_leaf(tape,vals,st,2,3,B,0) 73 let nY: i64 = ta_hadamard(tape,vals,st,nA,nB) 74 let nZ: i64 = ta_leaf(tape,vals,st,2,3,Z,0) 75 return ta_mse(tape,vals,st,nY,nZ) 76} 77func bld_silu(tape: *i64, vals: *i64, st: *i64, A: *i64, Z: *i64) -> i64 { 78 st[0]=0; st[1]=0 79 let nA: i64 = ta_leaf(tape,vals,st,2,3,A,0) 80 let nY: i64 = ta_silu(tape,vals,st,nA) 81 let nZ: i64 = ta_leaf(tape,vals,st,2,3,Z,0) 82 return ta_mse(tape,vals,st,nY,nZ) 83} 84func bld_rms(tape: *i64, vals: *i64, st: *i64, A: *i64, Z: *i64) -> i64 { 85 st[0]=0; st[1]=0 86 let nA: i64 = ta_leaf(tape,vals,st,2,3,A,0) 87 let nY: i64 = ta_rmsnorm_rows(tape,vals,st,nA) 88 let nZ: i64 = ta_leaf(tape,vals,st,2,3,Z,0) 89 return ta_mse(tape,vals,st,nY,nZ) 90} 91func bld_softce(tape: *i64, vals: *i64, st: *i64, A: *i64, ids: *i64) -> i64 { 92 st[0]=0; st[1]=0 93 let nA: i64 = ta_leaf(tape,vals,st,2,3,A,0) // logits [T=2, V=3] 94 return ta_softce_rows(tape,vals,st,nA,ids) 95} 96 97func main() -> i64 { 98 gt_puts("nx_ta_transformer_gradcheck_gate (finite-diff gradcheck of the 9 f32 transformer ops behind the R3e reader)\n" as *u8) 99 let tape: *i64 = sys_mmap(512*7*8) as *i64 100 let vals: *i64 = sys_mmap(4096*8) as *i64 101 let grads: *i64 = sys_mmap(4096*8) as *i64 102 let st: *i64 = sys_mmap(2*8) as *i64 103 let A: *i64 = sys_mmap(64*8) as *i64 104 let B: *i64 = sys_mmap(64*8) as *i64 105 let Z: *i64 = sys_mmap(64*8) as *i64 106 let ids: *i64 = sys_mmap(8*8) as *i64 107 let tab: *i64 = sys_mmap((2 + 2*2*2)*8) as *i64 // T=2, np=2 108 var i: i64 = 0 109 // Z = a NON-ZERO mse target. mse-vs-ZERO is scale-invariant under rmsnorm/softmax (‖normalized‖²≈const -> 110 // gradient structurally ~0 -> a degenerate test), so use a real target: the gradient is then non-trivial 111 // and the analytic-vs-numeric match becomes a meaningful liar-kill for every op. 112 while i < 64 { A[i] = ta_constf(((i*7+1)%9)-4, 6); B[i] = ta_constf(((i*5+2)%9)-4, 6); Z[i] = ta_constf(((i*3+1)%7)-3, 5); i = i + 1 } 113 ids[0] = 1; ids[1] = 2 114 ta_rope_build_tab(tab, 2, 2) 115 let h: i64 = ta_constf(1, 256) 116 let h2: i64 = nx_f32_mul(nx_i32_to_f32(2), h) 117 let chalf: i64 = ta_constf(1, 2) // cmul test constant (0.5; not the tiny perturbation h) 118 var pass: i64 = 0 119 var total: i64 = 0 120 121 // ---- helper: run one op's gradcheck over `ncell` cells of leaf-0's source `src`. opid selects the builder. ---- 122 // (dispatch inline; no fn-ptr. cells checked = indices cs[0..ncell).) 123 let cs: *i64 = sys_mmap(8*8) as *i64 124 cs[0]=0; cs[1]=2; cs[2]=4 125 126 // ===== per-op driver, unrolled by opid (0..8) ===== 127 var opid: i64 = 0 128 while opid < 9 { 129 // pick name + src + ncell 130 var name: *u8 = "?" as *u8 131 var src: *i64 = A 132 var ncell: i64 = 3 133 if opid==0 { name = "matmul" as *u8 } 134 if opid==1 { name = "matmul_nt" as *u8 } 135 if opid==2 { name = "cmul" as *u8 } 136 if opid==3 { name = "softmax_rows" as *u8 } 137 if opid==4 { name = "rope_tab" as *u8 } 138 if opid==5 { name = "hadamard" as *u8 } 139 if opid==6 { name = "silu" as *u8 } 140 if opid==7 { name = "rmsnorm_rows" as *u8 } 141 if opid==8 { name = "softce_rows" as *u8 } 142 gt_puts(" ["); gt_puts(name); gt_puts("]\n" as *u8) 143 144 // analytic: build once, backward, read grads of leaf-0 (arena offset 0) 145 var root: i64 = 0 146 if opid==0 { root = bld_matmul(tape,vals,st,A,B,Z) } 147 if opid==1 { root = bld_matmul_nt(tape,vals,st,A,B,Z) } 148 if opid==2 { root = bld_cmul(tape,vals,st,A,Z,chalf) } 149 if opid==3 { root = bld_smrows(tape,vals,st,A,Z) } 150 if opid==4 { root = bld_rope(tape,vals,st,A,Z,tab) } 151 if opid==5 { root = bld_had(tape,vals,st,A,B,Z) } 152 if opid==6 { root = bld_silu(tape,vals,st,A,Z) } 153 if opid==7 { root = bld_rms(tape,vals,st,A,Z) } 154 if opid==8 { root = bld_softce(tape,vals,st,A,ids) } 155 ta_backward(tape, vals, grads, st[0], root) 156 157 var allok: i64 = 1 158 var nt: *i64 = sys_mmap(8) 159 nt[0] = 0 160 var ci: i64 = 0 161 while ci < ncell { 162 let c: i64 = cs[ci] 163 let ana: i64 = grads[0 + c] // leaf-0 grad cell c 164 let sv: i64 = src[c] 165 // L(+h) 166 src[c] = nx_f32_add(sv, h) 167 var rp: i64 = 0 168 if opid==0 { rp = bld_matmul(tape,vals,st,A,B,Z) } 169 if opid==1 { rp = bld_matmul_nt(tape,vals,st,A,B,Z) } 170 if opid==2 { rp = bld_cmul(tape,vals,st,A,Z,chalf) } 171 if opid==3 { rp = bld_smrows(tape,vals,st,A,Z) } 172 if opid==4 { rp = bld_rope(tape,vals,st,A,Z,tab) } 173 if opid==5 { rp = bld_had(tape,vals,st,A,B,Z) } 174 if opid==6 { rp = bld_silu(tape,vals,st,A,Z) } 175 if opid==7 { rp = bld_rms(tape,vals,st,A,Z) } 176 if opid==8 { rp = bld_softce(tape,vals,st,A,ids) } 177 let Lp: i64 = ta_val(tape, vals, rp, 0) 178 // L(-h) 179 src[c] = nx_f32_sub(sv, h) 180 var rm: i64 = 0 181 if opid==0 { rm = bld_matmul(tape,vals,st,A,B,Z) } 182 if opid==1 { rm = bld_matmul_nt(tape,vals,st,A,B,Z) } 183 if opid==2 { rm = bld_cmul(tape,vals,st,A,Z,chalf) } 184 if opid==3 { rm = bld_smrows(tape,vals,st,A,Z) } 185 if opid==4 { rm = bld_rope(tape,vals,st,A,Z,tab) } 186 if opid==5 { rm = bld_had(tape,vals,st,A,B,Z) } 187 if opid==6 { rm = bld_silu(tape,vals,st,A,Z) } 188 if opid==7 { rm = bld_rms(tape,vals,st,A,Z) } 189 if opid==8 { rm = bld_softce(tape,vals,st,A,ids) } 190 let Lm: i64 = ta_val(tape, vals, rm, 0) 191 src[c] = sv 192 let ok: i64 = gt_cell(Lp, Lm, h2, ana, nt) 193 if ok == 0 { allok = 0 } 194 ci = ci + 1 195 } 196 total = total + 1 197 var v: i64 = 0 198 if allok == 1 { if nt[0] == 1 { v = 1 } } 199 if v == 1 { pass = pass + 1; gt_puts(" PASS " as *u8); gt_puts(name); gt_puts("\n" as *u8) } else { gt_puts(" FAIL " as *u8); gt_puts(name); gt_puts("\n" as *u8) } 200 opid = opid + 1 201 } 202 203 gt_puts("---- nx_ta_transformer_gradcheck_gate: passed "); gt_pn(pass); gt_puts(" / "); gt_pn(total); gt_puts("\n" as *u8) 204 if pass == total { gt_puts("TA-TRANSFORMER GRADCHECK GREEN -- all 9 f32 transformer backward identities match finite-difference; the R3e f32 reader's gradients are CORRECT.\n" as *u8); return 0 } 205 gt_puts("RED -- a backward identity is WRONG (see the ana/num pair that diverged) -- fix before trusting the f32 reader run.\n" as *u8) 206 return 1 207}