code wiki / _hdl_build / nx_train_r2_gate.nx
nx_train_r2_gate.nx source
↩ module page · 291 lines · 13816 B
1// nx_train_r2_gate.nx -- GATE for TRAIN-R2 (T6): tensor autograd. Proves, by RUNNING:
2// A MLP GRADCHECK: loss = MSE(W2*relu(W1*x+b1)+b2, t), W1 2x2/b1 2/W2 1x2/b2 1 = 9 params (values off the
3// relu kinks). Each param: analytic (reverse-mode) vs central finite diff (h=1/128), rel<1/32 floor 1/64.
4// Exercises every identity (matvec x2, vadd x2, relu, mse) through a real nonlinear composition.
5// B RECOVER AN AFFINE MAP: train W(2x2)+b(2) to recover y=A*x+c (A=[[3/2,-1/2],[1/4,1]], c=[-1/2,3/4]) from
6// 8 deterministic samples; full-batch GD lr=1/10, 400 epochs. Assert loss<1/1000 AND every W,b elt within
7// 1/16 of truth. (LINEAR model: zero-init is convex-safe; nonconvex needs an init strategy = next rung.)
8// C BIT-EXACT: run Gate-B training twice from scratch -> identical final bits for all 6 cells.
9// D AdamW: the same affine recovery via AdamW (m/v moments + bias correction + sqrt) also converges -- the
10// optimizer the FNet/transformer models will actually use.
11//
12// Evidence -> knowledge/status/train_r2.log (TRAINR2GATE authored=organ ... verdict=GREEN). license_tier: ORIGINAL
13import "nx_autograd_tensor.nx" // ta_* tensor tape + transitively nx_f32 / nx_f32_div / nx_f32_cvt / nx_syscalls
14import "nx_syscalls.nx"
15
16const T2_LOG: *u8 = "knowledge/status/train_r2.log"
17
18func t2_w(fd: i64, s: *u8) -> i64 { var n: i64 = 0; while s[n] != (0 as u8) { n = n + 1 } sys_write(fd, s, n); return 0 }
19func t2_wn(fd: i64, v: i64) -> i64 {
20 let bb: *u8 = sys_mmap(28); var m: i64 = v
21 if m < 0 { m = 0 - m; sys_write(fd, "-" as *u8, 1) }
22 let t: *u8 = sys_mmap(28); var k: i64 = 0
23 if m == 0 { t[0] = 48; k = 1 }
24 while m > 0 { t[k] = (48 + (m % 10)) as u8; m = m / 10; k = k + 1 }
25 var i: i64 = 0
26 while i < k { bb[i] = t[k - 1 - i]; i = i + 1 }
27 sys_write(fd, bb, k); return 0
28}
29
30// ===== Gate A helpers: the MLP loss graph =====
31func g_mlp_build(tape: *i64, vals: *i64, st: *i64, p: *i64, x: *i64, t: *i64, lv: *i64) -> i64 {
32 st[0] = 0; st[1] = 0
33 let W1: i64 = ta_leaf(tape, vals, st, 2, 2, p, 0)
34 let b1: i64 = ta_leaf(tape, vals, st, 2, 1, p, 4)
35 let W2: i64 = ta_leaf(tape, vals, st, 1, 2, p, 6)
36 let b2: i64 = ta_leaf(tape, vals, st, 1, 1, p, 8)
37 let xx: i64 = ta_leaf(tape, vals, st, 2, 1, x, 0)
38 let tt: i64 = ta_leaf(tape, vals, st, 1, 1, t, 0)
39 let h1: i64 = ta_vadd(tape, vals, st, ta_matvec(tape, vals, st, W1, xx), b1)
40 let rr: i64 = ta_relu(tape, vals, st, h1)
41 let h2: i64 = ta_vadd(tape, vals, st, ta_matvec(tape, vals, st, W2, rr), b2)
42 let loss: i64 = ta_mse(tape, vals, st, h2, tt)
43 lv[0] = W1; lv[1] = b1; lv[2] = W2; lv[3] = b2
44 return loss
45}
46func g_mlp_loss(tape: *i64, vals: *i64, st: *i64, p: *i64, x: *i64, t: *i64) -> i64 {
47 let lv: *i64 = (sys_mmap(4 * 8)) as *i64
48 let loss: i64 = g_mlp_build(tape, vals, st, p, x, t, lv)
49 return ta_val(tape, vals, loss, 0)
50}
51func g_param_grad(tape: *i64, grads: *i64, lv: *i64, pi: i64) -> i64 {
52 if pi < 4 { return ta_grad(tape, grads, lv[0], pi) }
53 if pi < 6 { return ta_grad(tape, grads, lv[1], pi - 4) }
54 if pi < 8 { return ta_grad(tape, grads, lv[2], pi - 6) }
55 return ta_grad(tape, grads, lv[3], 0)
56}
57
58// ===== shared affine forward (Gate B/C/D): loss = (1/8) sum_s MSE(W*x_s + b, y_s) =====
59func g_affine_build(tape: *i64, vals: *i64, st: *i64, p: *i64, xs: *i64, ys: *i64, c1: *i64, wb: *i64) -> i64 {
60 st[0] = 0; st[1] = 0
61 let W: i64 = ta_leaf(tape, vals, st, 2, 2, p, 0)
62 let b: i64 = ta_leaf(tape, vals, st, 2, 1, p, 4)
63 var sumn: i64 = 0 - 1
64 var s: i64 = 0
65 while s < 8 {
66 let xl: i64 = ta_leaf(tape, vals, st, 2, 1, xs, s * 2)
67 let yl: i64 = ta_leaf(tape, vals, st, 2, 1, ys, s * 2)
68 let pred: i64 = ta_vadd(tape, vals, st, ta_matvec(tape, vals, st, W, xl), b)
69 let msn: i64 = ta_mse(tape, vals, st, pred, yl)
70 if s == 0 { sumn = msn } else { sumn = ta_vadd(tape, vals, st, sumn, msn) }
71 s = s + 1
72 }
73 let inv8: i64 = ta_leaf(tape, vals, st, 1, 1, c1, 0)
74 let loss: i64 = ta_matvec(tape, vals, st, inv8, sumn) // (1/8)*sum
75 wb[0] = W; wb[1] = b
76 return loss
77}
78func g_read6(tape: *i64, grads: *i64, wb: *i64, g: *i64) -> i64 {
79 var i: i64 = 0
80 while i < 4 { g[i] = ta_grad(tape, grads, wb[0], i); i = i + 1 }
81 g[4] = ta_grad(tape, grads, wb[1], 0); g[5] = ta_grad(tape, grads, wb[1], 1)
82 return 0
83}
84
85// full-batch gradient descent. pout[0..5] = final params; *lossout = final loss.
86func g_train_gd(tape: *i64, vals: *i64, grads: *i64, st: *i64, xs: *i64, ys: *i64, epochs: i64, pout: *i64, lossout: *i64) -> i64 {
87 let p: *i64 = (sys_mmap(6 * 8)) as *i64
88 var i: i64 = 0
89 while i < 6 { p[i] = TA_F32_ZERO; i = i + 1 }
90 let lr: i64 = ta_constf(1, 10)
91 let c1: *i64 = (sys_mmap(8)) as *i64; c1[0] = ta_constf(1, 8)
92 let wb: *i64 = (sys_mmap(2 * 8)) as *i64
93 let g: *i64 = (sys_mmap(6 * 8)) as *i64
94 var ep: i64 = 0
95 while ep < epochs {
96 let loss: i64 = g_affine_build(tape, vals, st, p, xs, ys, c1, wb)
97 ta_backward(tape, vals, grads, st[0], loss)
98 *lossout = ta_val(tape, vals, loss, 0)
99 g_read6(tape, grads, wb, g)
100 i = 0
101 while i < 6 { p[i] = nx_f32_sub(p[i], nx_f32_mul(lr, g[i])); i = i + 1 }
102 ep = ep + 1
103 }
104 i = 0
105 while i < 6 { pout[i] = p[i]; i = i + 1 }
106 return 0
107}
108
109// AdamW. m,v moments + bias correction; weight decay 0 (= Adam) so it recovers the exact affine map.
110func g_train_adamw(tape: *i64, vals: *i64, grads: *i64, st: *i64, xs: *i64, ys: *i64, epochs: i64, pout: *i64, lossout: *i64) -> i64 {
111 let p: *i64 = (sys_mmap(6 * 8)) as *i64
112 let mm: *i64 = (sys_mmap(6 * 8)) as *i64
113 let vv: *i64 = (sys_mmap(6 * 8)) as *i64
114 var i: i64 = 0
115 while i < 6 { p[i] = TA_F32_ZERO; mm[i] = TA_F32_ZERO; vv[i] = TA_F32_ZERO; i = i + 1 }
116 let beta1: i64 = ta_constf(9, 10)
117 let beta2: i64 = ta_constf(999, 1000)
118 let om1: i64 = ta_constf(1, 10)
119 let om2: i64 = ta_constf(1, 1000)
120 let lr: i64 = ta_constf(1, 10)
121 let eps: i64 = ta_constf(1, 100000000)
122 var b1t: i64 = TA_F32_ONE
123 var b2t: i64 = TA_F32_ONE
124 let c1: *i64 = (sys_mmap(8)) as *i64; c1[0] = ta_constf(1, 8)
125 let wb: *i64 = (sys_mmap(2 * 8)) as *i64
126 let g: *i64 = (sys_mmap(6 * 8)) as *i64
127 var ep: i64 = 0
128 while ep < epochs {
129 let loss: i64 = g_affine_build(tape, vals, st, p, xs, ys, c1, wb)
130 ta_backward(tape, vals, grads, st[0], loss)
131 *lossout = ta_val(tape, vals, loss, 0)
132 g_read6(tape, grads, wb, g)
133 b1t = nx_f32_mul(b1t, beta1)
134 b2t = nx_f32_mul(b2t, beta2)
135 let c1corr: i64 = nx_f32_sub(TA_F32_ONE, b1t)
136 let c2corr: i64 = nx_f32_sub(TA_F32_ONE, b2t)
137 i = 0
138 while i < 6 {
139 let gi: i64 = g[i]
140 mm[i] = nx_f32_add(nx_f32_mul(beta1, mm[i]), nx_f32_mul(om1, gi))
141 vv[i] = nx_f32_add(nx_f32_mul(beta2, vv[i]), nx_f32_mul(om2, nx_f32_mul(gi, gi)))
142 let mhat: i64 = nx_f32_div(mm[i], c1corr)
143 let vhat: i64 = nx_f32_div(vv[i], c2corr)
144 let denom: i64 = nx_f32_add(nx_f32_sqrt(vhat), eps)
145 p[i] = nx_f32_sub(p[i], nx_f32_mul(lr, nx_f32_div(mhat, denom)))
146 i = i + 1
147 }
148 ep = ep + 1
149 }
150 i = 0
151 while i < 6 { pout[i] = p[i]; i = i + 1 }
152 return 0
153}
154
155func main() -> i64 {
156 var ok: i64 = 1
157 let tape: *i64 = (sys_mmap(1024 * 7 * 8)) as *i64
158 let vals: *i64 = (sys_mmap(4096 * 8)) as *i64
159 let grads: *i64 = (sys_mmap(4096 * 8)) as *i64
160 let st: *i64 = (sys_mmap(2 * 8)) as *i64
161
162 // ---------- Gate A: MLP gradcheck ----------
163 let p: *i64 = (sys_mmap(9 * 8)) as *i64
164 p[0] = ta_constf(2, 3); p[1] = ta_constf(0 - 1, 4); p[2] = ta_constf(1, 2); p[3] = ta_constf(1, 3)
165 p[4] = ta_constf(1, 4); p[5] = ta_constf(0 - 1, 8)
166 p[6] = ta_constf(3, 2); p[7] = ta_constf(0 - 1, 2)
167 p[8] = ta_constf(1, 4)
168 let xv: *i64 = (sys_mmap(2 * 8)) as *i64; xv[0] = ta_constf(3, 4); xv[1] = ta_constf(1, 2)
169 let tv: *i64 = (sys_mmap(1 * 8)) as *i64; tv[0] = TA_F32_ONE
170 let lv: *i64 = (sys_mmap(4 * 8)) as *i64
171 let lossA: i64 = g_mlp_build(tape, vals, st, p, xv, tv, lv)
172 ta_backward(tape, vals, grads, st[0], lossA)
173 let ana: *i64 = (sys_mmap(9 * 8)) as *i64
174 var pi: i64 = 0
175 while pi < 9 { ana[pi] = g_param_grad(tape, grads, lv, pi); pi = pi + 1 }
176 let h: i64 = ta_constf(1, 128)
177 let flo: i64 = ta_constf(1, 64)
178 let tol: i64 = ta_constf(1, 32)
179 var gradcheckA: i64 = 1
180 var worstA: i64 = 0
181 let pp: *i64 = (sys_mmap(9 * 8)) as *i64
182 let pm: *i64 = (sys_mmap(9 * 8)) as *i64
183 pi = 0
184 while pi < 9 {
185 var j: i64 = 0
186 while j < 9 { pp[j] = p[j]; pm[j] = p[j]; j = j + 1 }
187 pp[pi] = nx_f32_add(p[pi], h)
188 pm[pi] = nx_f32_sub(p[pi], h)
189 let lpv: i64 = g_mlp_loss(tape, vals, st, pp, xv, tv)
190 let lmv: i64 = g_mlp_loss(tape, vals, st, pm, xv, tv)
191 let fd: i64 = nx_f32_div(nx_f32_sub(lpv, lmv), nx_f32_add(h, h))
192 let num: i64 = nx_f32_abs(nx_f32_sub(fd, ana[pi]))
193 var den: i64 = nx_f32_abs(ana[pi])
194 if nx_f32_lt(den, flo) == 1 { den = flo }
195 if nx_f32_lt(num, nx_f32_mul(tol, den)) != 1 { gradcheckA = 0 }
196 let nm: i64 = ta_f32_to_milli(num)
197 if nm > worstA { worstA = nm }
198 pi = pi + 1
199 }
200 if gradcheckA != 1 { ok = 0 }
201
202 // ---------- samples for Gate B/C/D: y = A x + c ----------
203 let xi: *i64 = (sys_mmap(16 * 8)) as *i64 // 8 points * 2
204 xi[0] = nx_i32_to_f32(1); xi[1] = nx_i32_to_f32(0)
205 xi[2] = nx_i32_to_f32(0); xi[3] = nx_i32_to_f32(1)
206 xi[4] = nx_i32_to_f32(1); xi[5] = nx_i32_to_f32(1)
207 xi[6] = nx_i32_to_f32(0 - 1); xi[7] = nx_i32_to_f32(1)
208 xi[8] = nx_i32_to_f32(1); xi[9] = nx_i32_to_f32(0 - 1)
209 xi[10] = nx_i32_to_f32(0 - 1); xi[11] = nx_i32_to_f32(0)
210 xi[12] = nx_i32_to_f32(0); xi[13] = nx_i32_to_f32(0 - 1)
211 xi[14] = nx_i32_to_f32(2); xi[15] = nx_i32_to_f32(1)
212 let yi: *i64 = (sys_mmap(16 * 8)) as *i64
213 let A00: i64 = ta_constf(3, 2); let A01: i64 = ta_constf(1, 2); let A10: i64 = ta_constf(1, 4)
214 let c0: i64 = ta_constf(1, 2); let c1c: i64 = ta_constf(3, 4)
215 var s: i64 = 0
216 while s < 8 {
217 let x0: i64 = xi[s * 2]; let x1: i64 = xi[s * 2 + 1]
218 // y0 = 1.5 x0 - 0.5 x1 - 0.5 ; y1 = 0.25 x0 + x1 + 0.75
219 yi[s * 2] = nx_f32_sub(nx_f32_sub(nx_f32_mul(A00, x0), nx_f32_mul(A01, x1)), c0)
220 yi[s * 2 + 1] = nx_f32_add(nx_f32_add(nx_f32_mul(A10, x0), x1), c1c)
221 s = s + 1
222 }
223 // truth params for the within-1/16 check: W=[[1.5,-0.5],[0.25,1.0]], b=[-0.5,0.75]
224 let truth: *i64 = (sys_mmap(6 * 8)) as *i64
225 truth[0] = ta_constf(3, 2); truth[1] = ta_constf(0 - 1, 2); truth[2] = ta_constf(1, 4); truth[3] = TA_F32_ONE
226 truth[4] = ta_constf(0 - 1, 2); truth[5] = ta_constf(3, 4)
227 let tol16: i64 = ta_constf(1, 16)
228 let thou: i64 = ta_constf(1, 1000)
229
230 // ---------- Gate B: GD recovers the affine map ----------
231 let poutB: *i64 = (sys_mmap(6 * 8)) as *i64
232 let lossB: *i64 = (sys_mmap(8)) as *i64
233 g_train_gd(tape, vals, grads, st, xi, yi, 400, poutB, lossB)
234 var learnsB: i64 = 1
235 if nx_f32_lt(*lossB, thou) != 1 { learnsB = 0 }
236 var i: i64 = 0
237 while i < 6 {
238 if nx_f32_lt(nx_f32_abs(nx_f32_sub(poutB[i], truth[i])), tol16) != 1 { learnsB = 0 }
239 i = i + 1
240 }
241 if learnsB != 1 { ok = 0 }
242
243 // ---------- Gate C: bit-exact reproducible ----------
244 let poutC: *i64 = (sys_mmap(6 * 8)) as *i64
245 let lossC: *i64 = (sys_mmap(8)) as *i64
246 g_train_gd(tape, vals, grads, st, xi, yi, 400, poutC, lossC)
247 var reproC: i64 = 1
248 i = 0
249 while i < 6 { if poutC[i] != poutB[i] { reproC = 0 } i = i + 1 }
250 if reproC != 1 { ok = 0 }
251
252 // ---------- Gate D: AdamW also recovers ----------
253 let poutD: *i64 = (sys_mmap(6 * 8)) as *i64
254 let lossD: *i64 = (sys_mmap(8)) as *i64
255 g_train_adamw(tape, vals, grads, st, xi, yi, 400, poutD, lossD)
256 var learnsD: i64 = 1
257 if nx_f32_lt(*lossD, thou) != 1 { learnsD = 0 }
258 i = 0
259 while i < 6 {
260 if nx_f32_lt(nx_f32_abs(nx_f32_sub(poutD[i], truth[i])), tol16) != 1 { learnsD = 0 }
261 i = i + 1
262 }
263 if learnsD != 1 { ok = 0 }
264
265 // ---------- emit ----------
266 var fdi: i64 = 1
267 while fdi >= 0 {
268 var out: i64 = 1
269 if fdi == 0 { out = sys_openat_append(T2_LOG, 420) }
270 if out >= 0 {
271 t2_w(out, "TRAINR2GATE authored=organ engine=tensor-tape-autograd-f32" as *u8)
272 t2_w(out, " | A_mlp_gradcheck_pass=" as *u8); t2_wn(out, gradcheckA)
273 t2_w(out, " worst_|fd-analytic|_milli=" as *u8); t2_wn(out, worstA)
274 t2_w(out, " | B_affine_GD_pass=" as *u8); t2_wn(out, learnsB)
275 t2_w(out, " loss_milli=" as *u8); t2_wn(out, ta_f32_to_milli(*lossB))
276 t2_w(out, " W=[" as *u8); t2_wn(out, ta_f32_to_milli(poutB[0])); t2_w(out, "," as *u8); t2_wn(out, ta_f32_to_milli(poutB[1]))
277 t2_w(out, "," as *u8); t2_wn(out, ta_f32_to_milli(poutB[2])); t2_w(out, "," as *u8); t2_wn(out, ta_f32_to_milli(poutB[3]))
278 t2_w(out, "] b=[" as *u8); t2_wn(out, ta_f32_to_milli(poutB[4])); t2_w(out, "," as *u8); t2_wn(out, ta_f32_to_milli(poutB[5])); t2_w(out, "]" as *u8)
279 t2_w(out, " (truth W=[1500,-500,250,1000] b=[-500,750])" as *u8)
280 t2_w(out, " | C_bitexact_pass=" as *u8); t2_wn(out, reproC)
281 t2_w(out, " | D_AdamW_pass=" as *u8); t2_wn(out, learnsD)
282 t2_w(out, " loss_milli=" as *u8); t2_wn(out, ta_f32_to_milli(*lossD))
283 if ok == 1 { t2_w(out, " verdict=GREEN\n" as *u8) } else { t2_w(out, " verdict=RED\n" as *u8) }
284 if fdi == 0 { sys_close(out) }
285 }
286 fdi = fdi - 1
287 }
288
289 if ok == 1 { return 0 }
290 return 1
291}