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}