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