code wiki / _hdl_build / nx_train_r3_gate.nx

nx_train_r3_gate.nx source

↩ module page · 284 lines · 13331 B

1// nx_train_r3_gate.nx -- GATE for TRAIN-R3 (T7): softmax + cross-entropy + nonlinear training. Proves, by RUNNING: 2// A GRADCHECK the new ops: softmax-cross-entropy (dlogits = softmax - onehot) AND standalone softmax 3// (Jacobian-vector product), analytic vs central finite difference (h=1/128, rel<1/32 floor 1/64). 4// B TRAIN A NONLINEAR CLASSIFIER -- XOR: a 2->8(relu)->2 MLP with a softmax-CE head and a DETERMINISTIC 5// symmetry-breaking init, trained by AdamW. XOR is the canonical proof that a LINEAR model CANNOT solve 6// it -- reaching 4/4 correct demonstrates real nonlinear learning. Assert accuracy==4 AND loss decreased. 7// C BIT-EXACT: train twice from the same init -> identical final bits for all 42 parameters. 8// 9// Evidence -> knowledge/status/train_r3.log (TRAINR3GATE authored=organ ... verdict=GREEN). license_tier: ORIGINAL 10import "nx_autograd_tensor.nx" // ta_* (now incl softmax/softce/det_init) + transitively the f32 tower / syscalls 11import "nx_syscalls.nx" 12 13const T3_LOG: *u8 = "knowledge/status/train_r3.log" 14 15func t3_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 } 16func t3_wn(fd: i64, v: i64) -> i64 { 17 let bb: *u8 = sys_mmap(28); var m: i64 = v 18 if m < 0 { m = 0 - m; sys_write(fd, "-" as *u8, 1) } 19 let t: *u8 = sys_mmap(28); var k: i64 = 0 20 if m == 0 { t[0] = 48; k = 1 } 21 while m > 0 { t[k] = (48 + (m % 10)) as u8; m = m / 10; k = k + 1 } 22 var i: i64 = 0 23 while i < k { bb[i] = t[k - 1 - i]; i = i + 1 } 24 sys_write(fd, bb, k); return 0 25} 26 27// ===== Gate A: softmax-CE gradcheck ===== 28func g_softce_loss(tape: *i64, vals: *i64, st: *i64, lg: *i64, tg: *i64) -> i64 { 29 st[0] = 0; st[1] = 0 30 let ln: i64 = ta_leaf(tape, vals, st, 3, 1, lg, 0) 31 let tn: i64 = ta_leaf(tape, vals, st, 3, 1, tg, 0) 32 let loss: i64 = ta_softce(tape, vals, st, ln, tn) 33 return ta_val(tape, vals, loss, 0) 34} 35func g_softce_grads(tape: *i64, vals: *i64, grads: *i64, st: *i64, lg: *i64, tg: *i64, gout: *i64) -> i64 { 36 st[0] = 0; st[1] = 0 37 let ln: i64 = ta_leaf(tape, vals, st, 3, 1, lg, 0) 38 let tn: i64 = ta_leaf(tape, vals, st, 3, 1, tg, 0) 39 let loss: i64 = ta_softce(tape, vals, st, ln, tn) 40 ta_backward(tape, vals, grads, st[0], loss) 41 gout[0] = ta_grad(tape, grads, ln, 0); gout[1] = ta_grad(tape, grads, ln, 1); gout[2] = ta_grad(tape, grads, ln, 2) 42 return 0 43} 44// ===== Gate A: standalone softmax gradcheck (loss = mse(softmax(x), tg)) ===== 45func g_smax_loss(tape: *i64, vals: *i64, st: *i64, x: *i64, tg: *i64) -> i64 { 46 st[0] = 0; st[1] = 0 47 let xn: i64 = ta_leaf(tape, vals, st, 3, 1, x, 0) 48 let sm: i64 = ta_softmax(tape, vals, st, xn) 49 let tn: i64 = ta_leaf(tape, vals, st, 3, 1, tg, 0) 50 let loss: i64 = ta_mse(tape, vals, st, sm, tn) 51 return ta_val(tape, vals, loss, 0) 52} 53func g_smax_grads(tape: *i64, vals: *i64, grads: *i64, st: *i64, x: *i64, tg: *i64, gout: *i64) -> i64 { 54 st[0] = 0; st[1] = 0 55 let xn: i64 = ta_leaf(tape, vals, st, 3, 1, x, 0) 56 let sm: i64 = ta_softmax(tape, vals, st, xn) 57 let tn: i64 = ta_leaf(tape, vals, st, 3, 1, tg, 0) 58 let loss: i64 = ta_mse(tape, vals, st, sm, tn) 59 ta_backward(tape, vals, grads, st[0], loss) 60 gout[0] = ta_grad(tape, grads, xn, 0); gout[1] = ta_grad(tape, grads, xn, 1); gout[2] = ta_grad(tape, grads, xn, 2) 61 return 0 62} 63 64// generic 3-vector gradcheck: returns 1 if all pass, writes worst |fd-analytic| milli into *worst. kind 0=softce,1=softmax. 65func g_gc3(tape: *i64, vals: *i64, grads: *i64, st: *i64, base: *i64, tg: *i64, kind: i64, worst: *i64) -> i64 { 66 let ana: *i64 = (sys_mmap(3 * 8)) as *i64 67 if kind == 0 { g_softce_grads(tape, vals, grads, st, base, tg, ana) } else { g_smax_grads(tape, vals, grads, st, base, tg, ana) } 68 let h: i64 = ta_constf(1, 128) 69 let flo: i64 = ta_constf(1, 64) 70 let tol: i64 = ta_constf(1, 32) 71 let pp: *i64 = (sys_mmap(3 * 8)) as *i64 72 let pm: *i64 = (sys_mmap(3 * 8)) as *i64 73 var pass: i64 = 1 74 var wm: i64 = 0 75 var pi: i64 = 0 76 while pi < 3 { 77 var j: i64 = 0 78 while j < 3 { pp[j] = base[j]; pm[j] = base[j]; j = j + 1 } 79 pp[pi] = nx_f32_add(base[pi], h) 80 pm[pi] = nx_f32_sub(base[pi], h) 81 var lp: i64 = 0 82 var lm: i64 = 0 83 if kind == 0 { lp = g_softce_loss(tape, vals, st, pp, tg); lm = g_softce_loss(tape, vals, st, pm, tg) } 84 else { lp = g_smax_loss(tape, vals, st, pp, tg); lm = g_smax_loss(tape, vals, st, pm, tg) } 85 let fd: i64 = nx_f32_div(nx_f32_sub(lp, lm), nx_f32_add(h, h)) 86 let num: i64 = nx_f32_abs(nx_f32_sub(fd, ana[pi])) 87 var den: i64 = nx_f32_abs(ana[pi]) 88 if nx_f32_lt(den, flo) == 1 { den = flo } 89 if nx_f32_lt(num, nx_f32_mul(tol, den)) != 1 { pass = 0 } 90 let nm: i64 = ta_f32_to_milli(num) 91 if nm > wm { wm = nm } 92 pi = pi + 1 93 } 94 *worst = wm 95 return pass 96} 97 98// ===== Gate B: XOR MLP (2 -> 8 relu -> 2, softmax-CE). params p[42]: W1[0..15] b1[16..23] W2[24..39] b2[40..41] ===== 99func g_xor_build(tape: *i64, vals: *i64, st: *i64, p: *i64, xs: *i64, ts: *i64, c4: *i64, wb: *i64) -> i64 { 100 st[0] = 0; st[1] = 0 101 let W1: i64 = ta_leaf(tape, vals, st, 8, 2, p, 0) 102 let b1: i64 = ta_leaf(tape, vals, st, 8, 1, p, 16) 103 let W2: i64 = ta_leaf(tape, vals, st, 2, 8, p, 24) 104 let b2: i64 = ta_leaf(tape, vals, st, 2, 1, p, 40) 105 var sumn: i64 = 0 - 1 106 var s: i64 = 0 107 while s < 4 { 108 let x: i64 = ta_leaf(tape, vals, st, 2, 1, xs, s * 2) 109 let h: i64 = ta_relu(tape, vals, st, ta_vadd(tape, vals, st, ta_matvec(tape, vals, st, W1, x), b1)) 110 let lo: i64 = ta_vadd(tape, vals, st, ta_matvec(tape, vals, st, W2, h), b2) 111 let tt: i64 = ta_leaf(tape, vals, st, 2, 1, ts, s * 2) 112 let ls: i64 = ta_softce(tape, vals, st, lo, tt) 113 if s == 0 { sumn = ls } else { sumn = ta_vadd(tape, vals, st, sumn, ls) } 114 s = s + 1 115 } 116 let inv4: i64 = ta_leaf(tape, vals, st, 1, 1, c4, 0) 117 let loss: i64 = ta_matvec(tape, vals, st, inv4, sumn) 118 wb[0] = W1; wb[1] = b1; wb[2] = W2; wb[3] = b2 119 return loss 120} 121func g_read42(tape: *i64, grads: *i64, wb: *i64, g: *i64) -> i64 { 122 var i: i64 = 0 123 while i < 16 { g[i] = ta_grad(tape, grads, wb[0], i); i = i + 1 } 124 i = 0 125 while i < 8 { g[16 + i] = ta_grad(tape, grads, wb[1], i); i = i + 1 } 126 i = 0 127 while i < 16 { g[24 + i] = ta_grad(tape, grads, wb[2], i); i = i + 1 } 128 g[40] = ta_grad(tape, grads, wb[3], 0); g[41] = ta_grad(tape, grads, wb[3], 1) 129 return 0 130} 131// argmax of the 2 output logits for sample s 132func g_xor_predict(tape: *i64, vals: *i64, st: *i64, p: *i64, xs: *i64, s: i64) -> i64 { 133 st[0] = 0; st[1] = 0 134 let W1: i64 = ta_leaf(tape, vals, st, 8, 2, p, 0) 135 let b1: i64 = ta_leaf(tape, vals, st, 8, 1, p, 16) 136 let W2: i64 = ta_leaf(tape, vals, st, 2, 8, p, 24) 137 let b2: i64 = ta_leaf(tape, vals, st, 2, 1, p, 40) 138 let x: i64 = ta_leaf(tape, vals, st, 2, 1, xs, s * 2) 139 let h: i64 = ta_relu(tape, vals, st, ta_vadd(tape, vals, st, ta_matvec(tape, vals, st, W1, x), b1)) 140 let lo: i64 = ta_vadd(tape, vals, st, ta_matvec(tape, vals, st, W2, h), b2) 141 if nx_f32_gt(ta_val(tape, vals, lo, 1), ta_val(tape, vals, lo, 0)) == 1 { return 1 } 142 return 0 143} 144// init the 42 params: W1 (seed 3) + W2 (seed 7) spread, biases zero. 145func g_xor_init(p: *i64) -> i64 { 146 var i: i64 = 0 147 while i < 42 { p[i] = TA_F32_ZERO; i = i + 1 } 148 i = 0 149 while i < 16 { p[i] = ta_constf(((i * 3 + 1) % 11) - 5, 8); i = i + 1 } 150 i = 0 151 while i < 16 { p[24 + i] = ta_constf(((i * 7 + 1) % 11) - 5, 8); i = i + 1 } 152 return 0 153} 154// AdamW training of the XOR MLP; writes final params + first/last loss. 155func g_xor_train(tape: *i64, vals: *i64, grads: *i64, st: *i64, xs: *i64, ts: *i64, epochs: i64, pout: *i64, lf: *i64, ll: *i64) -> i64 { 156 let p: *i64 = (sys_mmap(42 * 8)) as *i64 157 let mm: *i64 = (sys_mmap(42 * 8)) as *i64 158 let vv: *i64 = (sys_mmap(42 * 8)) as *i64 159 g_xor_init(p) 160 var i: i64 = 0 161 while i < 42 { mm[i] = TA_F32_ZERO; vv[i] = TA_F32_ZERO; i = i + 1 } 162 let beta1: i64 = ta_constf(9, 10) 163 let beta2: i64 = ta_constf(999, 1000) 164 let om1: i64 = ta_constf(1, 10) 165 let om2: i64 = ta_constf(1, 1000) 166 let lr: i64 = ta_constf(1, 20) 167 let eps: i64 = ta_constf(1, 100000000) 168 var b1t: i64 = TA_F32_ONE 169 var b2t: i64 = TA_F32_ONE 170 let c4: *i64 = (sys_mmap(8)) as *i64; c4[0] = ta_constf(1, 4) 171 let wb: *i64 = (sys_mmap(4 * 8)) as *i64 172 let g: *i64 = (sys_mmap(42 * 8)) as *i64 173 var ep: i64 = 0 174 while ep < epochs { 175 let loss: i64 = g_xor_build(tape, vals, st, p, xs, ts, c4, wb) 176 ta_backward(tape, vals, grads, st[0], loss) 177 if ep == 0 { *lf = ta_val(tape, vals, loss, 0) } 178 *ll = ta_val(tape, vals, loss, 0) 179 g_read42(tape, grads, wb, g) 180 b1t = nx_f32_mul(b1t, beta1) 181 b2t = nx_f32_mul(b2t, beta2) 182 let c1c: i64 = nx_f32_sub(TA_F32_ONE, b1t) 183 let c2c: i64 = nx_f32_sub(TA_F32_ONE, b2t) 184 i = 0 185 while i < 42 { 186 let gi: i64 = g[i] 187 mm[i] = nx_f32_add(nx_f32_mul(beta1, mm[i]), nx_f32_mul(om1, gi)) 188 vv[i] = nx_f32_add(nx_f32_mul(beta2, vv[i]), nx_f32_mul(om2, nx_f32_mul(gi, gi))) 189 let mhat: i64 = nx_f32_div(mm[i], c1c) 190 let vhat: i64 = nx_f32_div(vv[i], c2c) 191 p[i] = nx_f32_sub(p[i], nx_f32_mul(lr, nx_f32_div(mhat, nx_f32_add(nx_f32_sqrt(vhat), eps)))) 192 i = i + 1 193 } 194 ep = ep + 1 195 } 196 i = 0 197 while i < 42 { pout[i] = p[i]; i = i + 1 } 198 return 0 199} 200 201func main() -> i64 { 202 var ok: i64 = 1 203 let tape: *i64 = (sys_mmap(1024 * 7 * 8)) as *i64 204 let vals: *i64 = (sys_mmap(4096 * 8)) as *i64 205 let grads: *i64 = (sys_mmap(4096 * 8)) as *i64 206 let st: *i64 = (sys_mmap(2 * 8)) as *i64 207 208 // ----- Gate A: gradcheck softmax-CE + standalone softmax ----- 209 let lg: *i64 = (sys_mmap(3 * 8)) as *i64 210 lg[0] = ta_constf(1, 2); lg[1] = ta_constf(0 - 3, 10); lg[2] = ta_constf(4, 5) 211 let oneh: *i64 = (sys_mmap(3 * 8)) as *i64 212 oneh[0] = TA_F32_ZERO; oneh[1] = TA_F32_ONE; oneh[2] = TA_F32_ZERO 213 let wce: *i64 = (sys_mmap(8)) as *i64 214 let ceGc: i64 = g_gc3(tape, vals, grads, st, lg, oneh, 0, wce) 215 if ceGc != 1 { ok = 0 } 216 217 let xg: *i64 = (sys_mmap(3 * 8)) as *i64 218 xg[0] = ta_constf(2, 5); xg[1] = ta_constf(0 - 1, 5); xg[2] = ta_constf(7, 10) 219 let dist: *i64 = (sys_mmap(3 * 8)) as *i64 220 dist[0] = ta_constf(3, 10); dist[1] = ta_constf(2, 10); dist[2] = ta_constf(5, 10) 221 let wsm: *i64 = (sys_mmap(8)) as *i64 222 let smGc: i64 = g_gc3(tape, vals, grads, st, xg, dist, 1, wsm) 223 if smGc != 1 { ok = 0 } 224 225 // ----- Gate B: train XOR ----- 226 let xs: *i64 = (sys_mmap(8 * 8)) as *i64 227 xs[0] = TA_F32_ZERO; xs[1] = TA_F32_ZERO 228 xs[2] = TA_F32_ZERO; xs[3] = TA_F32_ONE 229 xs[4] = TA_F32_ONE; xs[5] = TA_F32_ZERO 230 xs[6] = TA_F32_ONE; xs[7] = TA_F32_ONE 231 let ts: *i64 = (sys_mmap(8 * 8)) as *i64 // onehot targets per XOR label 0,1,1,0 232 ts[0] = TA_F32_ONE; ts[1] = TA_F32_ZERO 233 ts[2] = TA_F32_ZERO; ts[3] = TA_F32_ONE 234 ts[4] = TA_F32_ZERO; ts[5] = TA_F32_ONE 235 ts[6] = TA_F32_ONE; ts[7] = TA_F32_ZERO 236 let labels: *i64 = (sys_mmap(4 * 8)) as *i64 237 labels[0] = 0; labels[1] = 1; labels[2] = 1; labels[3] = 0 238 let pB: *i64 = (sys_mmap(42 * 8)) as *i64 239 let lfB: *i64 = (sys_mmap(8)) as *i64 240 let llB: *i64 = (sys_mmap(8)) as *i64 241 g_xor_train(tape, vals, grads, st, xs, ts, 2500, pB, lfB, llB) 242 var acc: i64 = 0 243 var s: i64 = 0 244 while s < 4 { 245 if g_xor_predict(tape, vals, st, pB, xs, s) == labels[s] { acc = acc + 1 } 246 s = s + 1 247 } 248 var learnsB: i64 = 1 249 if acc != 4 { learnsB = 0 } 250 if nx_f32_lt(*llB, *lfB) != 1 { learnsB = 0 } 251 if learnsB != 1 { ok = 0 } 252 253 // ----- Gate C: bit-exact ----- 254 let pC: *i64 = (sys_mmap(42 * 8)) as *i64 255 let lfC: *i64 = (sys_mmap(8)) as *i64 256 let llC: *i64 = (sys_mmap(8)) as *i64 257 g_xor_train(tape, vals, grads, st, xs, ts, 2500, pC, lfC, llC) 258 var reproC: i64 = 1 259 var i: i64 = 0 260 while i < 42 { if pC[i] != pB[i] { reproC = 0 } i = i + 1 } 261 if reproC != 1 { ok = 0 } 262 263 // ----- emit ----- 264 var fdi: i64 = 1 265 while fdi >= 0 { 266 var out: i64 = 1 267 if fdi == 0 { out = sys_openat_append(T3_LOG, 420) } 268 if out >= 0 { 269 t3_w(out, "TRAINR3GATE authored=organ engine=tensor-autograd softmax+crossentropy+nonlinear" as *u8) 270 t3_w(out, " | A_softce_gradcheck=" as *u8); t3_wn(out, ceGc); t3_w(out, " worst_milli=" as *u8); t3_wn(out, wce[0]) 271 t3_w(out, " A_softmax_gradcheck=" as *u8); t3_wn(out, smGc); t3_w(out, " worst_milli=" as *u8); t3_wn(out, wsm[0]) 272 t3_w(out, " | B_XOR_accuracy=" as *u8); t3_wn(out, acc); t3_w(out, "/4" as *u8) 273 t3_w(out, " loss_first_milli=" as *u8); t3_wn(out, ta_f32_to_milli(*lfB)) 274 t3_w(out, " loss_last_milli=" as *u8); t3_wn(out, ta_f32_to_milli(*llB)) 275 t3_w(out, " | C_bitexact=" as *u8); t3_wn(out, reproC) 276 if ok == 1 { t3_w(out, " verdict=GREEN\n" as *u8) } else { t3_w(out, " verdict=RED\n" as *u8) } 277 if fdi == 0 { sys_close(out) } 278 } 279 fdi = fdi - 1 280 } 281 282 if ok == 1 { return 0 } 283 return 1 284}