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}