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}