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}