code wiki / _hdl_build / nx_lowrank_train.nx
nx_lowrank_train.nx source
↩ module page · 273 lines · 13813 B
1// nx_lowrank_train.nx -- M2: does a LOW-RANK bottleneck IMPROVE quality (not just preserve it)? The MLA/
2// inductive-bias thesis (operator: "quality UP, VRAM down -- not just quant shrinking quality"). Trains
3// models on a NOISY rank-2 target, scored on a CLEAN held-out test set, swept over many SEEDS (robustness)
4// AND over bottleneck RANK (the U-curve -- guards the over-claim "smaller is always better"):
5// FULL: y = W x (W is D x D = 64 params) -- can overfit the noise
6// LOWRANK: y = Wu (Wd x) (Wd r x D, Wu D x r = 2*r*D params, the rank-r bottleneck)
7// CLAIM (precise, not over-claimed): the OPTIMAL bottleneck rank MATCHES the data's true rank R0=2 -- too
8// small (r=1) UNDERFITS, too large (r=4) OVERFITS like full, and at r=R0 it beats full at HALF the VRAM.
9// Sovereign autograd (nx_tgrad_core tg_* + AdamW ad_step; the proven _t4 recipe). nx_cc UNTOUCHED (operator:
10// it's the parallel effort). license_tier: ORIGINAL
11import "nx_syscalls.nx"
12import "_hdl_build/nx_tgrad_core.nx"
13const M2_MAGIC_1103515245: i64 = 1103515245
14const M2_MAGIC_12345: i64 = 12345
15const M2_MAGIC_2048: i64 = 2048
16const M2_MAGIC_4096: i64 = 4096
17const M2_MAGIC_20480: i64 = 20480
18const M2_MAGIC_16384: i64 = 16384
19const M2_MAGIC_100000: i64 = 100000
20const M2_MAGIC_1000003: i64 = 1000003
21const M2_MAGIC_7919: i64 = 7919
22
23const M2_D: i64 = 8
24const M2_R0: i64 = 2 // TRUE target rank
25const M2_R: i64 = 2 // low-rank model bottleneck (matched to R0)
26const M2_RMAX: i64 = 4 // max bottleneck rank in the sweep (buffer sizing)
27const M2_NTR: i64 = 12 // train samples
28const M2_NTE: i64 = 8 // test samples
29const M2_EPOCHS: i64 = 500
30const M2_NSEED: i64 = 9 // seed sweep -- robustness, not one lucky draw
31
32func m2_puts(s: *u8) -> i64 { var n: i64 = 0; while s[n] != (0 as u8) { n = n + 1 } sys_write(1, s, n); return 0 }
33func m2_putn(v: i64) -> i64 {
34 if v == 0 { sys_write(1, "0" as *u8, 1); return 0 }
35 var m: i64 = v
36 if m < 0 { sys_write(1, "-" as *u8, 1); m = 0 - m }
37 let d: *u8 = sys_mmap(24); var k: i64 = 0
38 while m > 0 { d[k] = (48 + (m % 10)) as u8; m = m / 10; k = k + 1 }
39 var j: i64 = k - 1
40 while j >= 0 { sys_write(1, ((d as i64)+j) as *u8, 1); j = j - 1 }
41 return 0
42}
43func m2_lcg(st: *i64) -> i64 { st[0] = st[0] * M2_MAGIC_1103515245 + M2_MAGIC_12345; return (st[0] >> 20) & 0xfff }
44func m2_rf(st: *i64) -> i64 { return tg_q(m2_lcg(st) - M2_MAGIC_2048, M2_MAGIC_4096) } // ~[-0.5, 0.5]
45func m2_rsmall(st: *i64) -> i64 { return tg_q(m2_lcg(st) - M2_MAGIC_2048, M2_MAGIC_20480) } // ~[-0.1, 0.1] (weight init)
46
47// y = U @ (V @ x) -- the true rank-R0 map (U is D x R0, V is R0 x D)
48func m2_apply(x: *i64, U: *i64, V: *i64, y: *i64) -> i64 {
49 let vx: *i64 = sys_mmap(M2_R0 * 8) as *i64
50 var r: i64 = 0
51 while r < M2_R0 {
52 var acc: i64 = 0
53 var j: i64 = 0
54 while j < M2_D { acc = nx_f32_add(acc, nx_f32_mul(V[r * M2_D + j], x[j])); j = j + 1 }
55 vx[r] = acc; r = r + 1
56 }
57 var i: i64 = 0
58 while i < M2_D {
59 var acc2: i64 = 0
60 var r2: i64 = 0
61 while r2 < M2_R0 { acc2 = nx_f32_add(acc2, nx_f32_mul(U[i * M2_R0 + r2], vx[r2])); r2 = r2 + 1 }
62 y[i] = acc2; i = i + 1
63 }
64 return 0
65}
66func m2_gen(X: *i64, Y: *i64, n: i64, U: *i64, V: *i64, st: *i64, noisy: i64) -> i64 {
67 var s: i64 = 0
68 while s < n {
69 var j: i64 = 0
70 while j < M2_D { X[s * M2_D + j] = m2_rf(st); j = j + 1 }
71 m2_apply((X as i64 + s * M2_D * 8) as *i64, U, V, (Y as i64 + s * M2_D * 8) as *i64)
72 if noisy == 1 {
73 j = 0
74 while j < M2_D { Y[s * M2_D + j] = nx_f32_add(Y[s * M2_D + j], tg_q(m2_lcg(st) - M2_MAGIC_2048, M2_MAGIC_16384)); j = j + 1 }
75 }
76 s = s + 1
77 }
78 return 0
79}
80
81func m2_train_full(tape: *i64, nb: *i64, arena: *i64, ab: *i64, X: *i64, Y: *i64, W: *i64, mW: *i64, vW: *i64) -> i64 {
82 let lr: i64 = tg_q(1, 20); let b1: i64 = tg_q(9, 10); let b2: i64 = tg_q(999, 1000); let eps: i64 = tg_q(1, M2_MAGIC_100000)
83 let invn: *i64 = sys_mmap(8) as *i64; invn[0] = tg_q(1, M2_NTR)
84 var ep: i64 = 0
85 while ep < M2_EPOCHS {
86 nb[0] = 0; ab[0] = 0
87 let lW: i64 = tg_leaf(tape, nb, arena, ab, W, M2_D, M2_D)
88 var accn: i64 = 0 - 1
89 var s: i64 = 0
90 while s < M2_NTR {
91 let lx: i64 = tg_leaf(tape, nb, arena, ab, (X as i64 + s * M2_D * 8) as *i64, M2_D, 1)
92 let lt: i64 = tg_leaf(tape, nb, arena, ab, (Y as i64 + s * M2_D * 8) as *i64, M2_D, 1)
93 let yy: i64 = tg_matvec(tape, nb, arena, ab, lW, lx)
94 let mm: i64 = tg_mse(tape, nb, arena, ab, yy, lt)
95 if accn < 0 { accn = mm } else { accn = tg_addvec(tape, nb, arena, ab, accn, mm) }
96 s = s + 1
97 }
98 let li: i64 = tg_leaf(tape, nb, arena, ab, invn, 1, 1)
99 let loss: i64 = tg_smul(tape, nb, arena, ab, accn, li)
100 tg_backward(tape, nb[0], loss)
101 ad_step(W, tg_gradp(tape, lW), mW, vW, M2_D * M2_D, lr, b1, b2, eps, 0, ep + 1)
102 ep = ep + 1
103 }
104 return 0
105}
106// bottleneck rank r is a parameter -> the same code trains r=1 (underfit), r=2 (matched), r=4 (over-capacity)
107func m2_train_low(tape: *i64, nb: *i64, arena: *i64, ab: *i64, X: *i64, Y: *i64, Wd: *i64, Wu: *i64, mWd: *i64, vWd: *i64, mWu: *i64, vWu: *i64, r: i64) -> i64 {
108 let lr: i64 = tg_q(1, 20); let b1: i64 = tg_q(9, 10); let b2: i64 = tg_q(999, 1000); let eps: i64 = tg_q(1, M2_MAGIC_100000)
109 let invn: *i64 = sys_mmap(8) as *i64; invn[0] = tg_q(1, M2_NTR)
110 var ep: i64 = 0
111 while ep < M2_EPOCHS {
112 nb[0] = 0; ab[0] = 0
113 let lWd: i64 = tg_leaf(tape, nb, arena, ab, Wd, r, M2_D)
114 let lWu: i64 = tg_leaf(tape, nb, arena, ab, Wu, M2_D, r)
115 var accn: i64 = 0 - 1
116 var s: i64 = 0
117 while s < M2_NTR {
118 let lx: i64 = tg_leaf(tape, nb, arena, ab, (X as i64 + s * M2_D * 8) as *i64, M2_D, 1)
119 let lt: i64 = tg_leaf(tape, nb, arena, ab, (Y as i64 + s * M2_D * 8) as *i64, M2_D, 1)
120 let hh: i64 = tg_matvec(tape, nb, arena, ab, lWd, lx)
121 let yy: i64 = tg_matvec(tape, nb, arena, ab, lWu, hh)
122 let mm: i64 = tg_mse(tape, nb, arena, ab, yy, lt)
123 if accn < 0 { accn = mm } else { accn = tg_addvec(tape, nb, arena, ab, accn, mm) }
124 s = s + 1
125 }
126 let li: i64 = tg_leaf(tape, nb, arena, ab, invn, 1, 1)
127 let loss: i64 = tg_smul(tape, nb, arena, ab, accn, li)
128 tg_backward(tape, nb[0], loss)
129 ad_step(Wd, tg_gradp(tape, lWd), mWd, vWd, r * M2_D, lr, b1, b2, eps, 0, ep + 1)
130 ad_step(Wu, tg_gradp(tape, lWu), mWu, vWu, M2_D * r, lr, b1, b2, eps, 0, ep + 1)
131 ep = ep + 1
132 }
133 return 0
134}
135func m2_eval_full(tape: *i64, nb: *i64, arena: *i64, ab: *i64, X: *i64, Y: *i64, n: i64, W: *i64) -> i64 {
136 nb[0] = 0; ab[0] = 0
137 let lW: i64 = tg_leaf(tape, nb, arena, ab, W, M2_D, M2_D)
138 var accn: i64 = 0 - 1
139 var s: i64 = 0
140 while s < n {
141 let lx: i64 = tg_leaf(tape, nb, arena, ab, (X as i64 + s * M2_D * 8) as *i64, M2_D, 1)
142 let lt: i64 = tg_leaf(tape, nb, arena, ab, (Y as i64 + s * M2_D * 8) as *i64, M2_D, 1)
143 let yy: i64 = tg_matvec(tape, nb, arena, ab, lW, lx)
144 let mm: i64 = tg_mse(tape, nb, arena, ab, yy, lt)
145 if accn < 0 { accn = mm } else { accn = tg_addvec(tape, nb, arena, ab, accn, mm) }
146 s = s + 1
147 }
148 let invn: *i64 = sys_mmap(8) as *i64; invn[0] = tg_q(1, n)
149 let li: i64 = tg_leaf(tape, nb, arena, ab, invn, 1, 1)
150 let loss: i64 = tg_smul(tape, nb, arena, ab, accn, li)
151 let lv: *i64 = tg_valp(tape, loss)
152 return lv[0]
153}
154func m2_eval_low(tape: *i64, nb: *i64, arena: *i64, ab: *i64, X: *i64, Y: *i64, n: i64, Wd: *i64, Wu: *i64, r: i64) -> i64 {
155 nb[0] = 0; ab[0] = 0
156 let lWd: i64 = tg_leaf(tape, nb, arena, ab, Wd, r, M2_D)
157 let lWu: i64 = tg_leaf(tape, nb, arena, ab, Wu, M2_D, r)
158 var accn: i64 = 0 - 1
159 var s: i64 = 0
160 while s < n {
161 let lx: i64 = tg_leaf(tape, nb, arena, ab, (X as i64 + s * M2_D * 8) as *i64, M2_D, 1)
162 let lt: i64 = tg_leaf(tape, nb, arena, ab, (Y as i64 + s * M2_D * 8) as *i64, M2_D, 1)
163 let hh: i64 = tg_matvec(tape, nb, arena, ab, lWd, lx)
164 let yy: i64 = tg_matvec(tape, nb, arena, ab, lWu, hh)
165 let mm: i64 = tg_mse(tape, nb, arena, ab, yy, lt)
166 if accn < 0 { accn = mm } else { accn = tg_addvec(tape, nb, arena, ab, accn, mm) }
167 s = s + 1
168 }
169 let invn: *i64 = sys_mmap(8) as *i64; invn[0] = tg_q(1, n)
170 let li: i64 = tg_leaf(tape, nb, arena, ab, invn, 1, 1)
171 let loss: i64 = tg_smul(tape, nb, arena, ab, accn, li)
172 let lv: *i64 = tg_valp(tape, loss)
173 return lv[0]
174}
175
176// run the 9-seed sweep at bottleneck rank r; out[0]=avg low test_mse(milli), out[1]=wins vs full, out[2]=avg full test_mse(milli)
177func m2_sweep_rank(tape: *i64, nb: *i64, arena: *i64, ab: *i64, st: *i64, r: i64, out: *i64) -> i64 {
178 let U: *i64 = sys_mmap(M2_D * M2_R0 * 8) as *i64
179 let V: *i64 = sys_mmap(M2_R0 * M2_D * 8) as *i64
180 let Xtr: *i64 = sys_mmap(M2_NTR * M2_D * 8) as *i64
181 let Ytr: *i64 = sys_mmap(M2_NTR * M2_D * 8) as *i64
182 let Xte: *i64 = sys_mmap(M2_NTE * M2_D * 8) as *i64
183 let Yte: *i64 = sys_mmap(M2_NTE * M2_D * 8) as *i64
184 let W: *i64 = sys_mmap(M2_D * M2_D * 8) as *i64
185 let mW: *i64 = sys_mmap(M2_D * M2_D * 8) as *i64
186 let vW: *i64 = sys_mmap(M2_D * M2_D * 8) as *i64
187 let Wd: *i64 = sys_mmap(r * M2_D * 8) as *i64
188 let mWd: *i64 = sys_mmap(r * M2_D * 8) as *i64
189 let vWd: *i64 = sys_mmap(r * M2_D * 8) as *i64
190 let Wu: *i64 = sys_mmap(M2_D * r * 8) as *i64
191 let mWu: *i64 = sys_mmap(M2_D * r * 8) as *i64
192 let vWu: *i64 = sys_mmap(M2_D * r * 8) as *i64
193 var wins: i64 = 0
194 var sum_low: i64 = 0
195 var sum_full: i64 = 0
196 var sd: i64 = 0
197 while sd < M2_NSEED {
198 st[0] = M2_MAGIC_1000003 + sd * M2_MAGIC_7919
199 var i: i64 = 0
200 while i < M2_D * M2_R0 { U[i] = m2_rf(st); i = i + 1 }
201 i = 0
202 while i < M2_R0 * M2_D { V[i] = m2_rf(st); i = i + 1 }
203 m2_gen(Xtr, Ytr, M2_NTR, U, V, st, 1) // train: noisy
204 m2_gen(Xte, Yte, M2_NTE, U, V, st, 0) // test: clean
205 i = 0
206 while i < M2_D * M2_D { W[i] = m2_rsmall(st); mW[i] = 0; vW[i] = 0; i = i + 1 }
207 i = 0
208 while i < r * M2_D { Wd[i] = m2_rsmall(st); mWd[i] = 0; vWd[i] = 0; i = i + 1 }
209 i = 0
210 while i < M2_D * r { Wu[i] = m2_rsmall(st); mWu[i] = 0; vWu[i] = 0; i = i + 1 }
211
212 m2_train_full(tape, nb, arena, ab, Xtr, Ytr, W, mW, vW)
213 m2_train_low(tape, nb, arena, ab, Xtr, Ytr, Wd, Wu, mWd, vWd, mWu, vWu, r)
214
215 let full_te: i64 = m2_eval_full(tape, nb, arena, ab, Xte, Yte, M2_NTE, W)
216 let low_te: i64 = m2_eval_low(tape, nb, arena, ab, Xte, Yte, M2_NTE, Wd, Wu, r)
217 sum_low = sum_low + tg_milli(low_te)
218 sum_full = sum_full + tg_milli(full_te)
219 if nx_f32_lt(low_te, full_te) == 1 { wins = wins + 1 }
220 sd = sd + 1
221 }
222 out[0] = sum_low / M2_NSEED
223 out[1] = wins
224 out[2] = sum_full / M2_NSEED
225 return 0
226}
227
228func main() -> i64 {
229 let tape: *i64 = sys_mmap(256 * 8 * 8) as *i64
230 let arena: *i64 = sys_mmap(M2_MAGIC_16384 * 8) as *i64
231 let nb: *i64 = sys_mmap(8) as *i64
232 let ab: *i64 = sys_mmap(8) as *i64
233 let st: *i64 = sys_mmap(8) as *i64
234
235 m2_puts("=== M2: low-rank bottleneck quality -- noisy rank-2 target, CLEAN test, "); m2_putn(M2_NSEED)
236 m2_puts(" seeds x rank sweep ===\n")
237 m2_puts(" true rank R0=2; FULL=64 params; LOWRANK at rank r = 2*r*D params\n")
238
239 let out: *i64 = sys_mmap(3 * 8) as *i64
240 let rks: *i64 = sys_mmap(3 * 8) as *i64
241 rks[0] = 1; rks[1] = 2; rks[2] = 4
242 let te1: *i64 = sys_mmap(8) as *i64 // avg test_mse per rank
243 let te2: *i64 = sys_mmap(8) as *i64
244 let te4: *i64 = sys_mmap(8) as *i64
245 var full_ref: i64 = 0
246 var ri: i64 = 0
247 while ri < 3 {
248 let r: i64 = rks[ri]
249 m2_sweep_rank(tape, nb, arena, ab, st, r, out)
250 if ri == 0 { te1[0] = out[0] }
251 if ri == 1 { te2[0] = out[0]; full_ref = out[2] }
252 if ri == 2 { te4[0] = out[0] }
253 m2_puts(" rank r="); m2_putn(r); m2_puts(" ("); m2_putn(2 * r * M2_D); m2_puts(" params): avg test_mse=")
254 m2_putn(out[0]); m2_puts("milli beats-full "); m2_putn(out[1]); m2_puts("/"); m2_putn(M2_NSEED)
255 if r < M2_R0 { m2_puts(" (UNDER-rank: too small to fit rank-2)\n") }
256 if r == M2_R0 { m2_puts(" (MATCHED to true rank)\n") }
257 if r > M2_R0 { m2_puts(" (OVER-rank: capacity to overfit noise)\n") }
258 ri = ri + 1
259 }
260 m2_puts(" FULL (64 params): avg test_mse="); m2_putn(full_ref); m2_puts("milli (reference)\n")
261 m2_puts(" U-CURVE test_mse(milli): r1="); m2_putn(te1[0]); m2_puts(" -> r2="); m2_putn(te2[0])
262 m2_puts(" -> r4="); m2_putn(te4[0]); m2_puts("\n")
263
264 var pass: i64 = 0
265 var fail: i64 = 0
266 if te2[0] < te1[0] { m2_puts(" T1 matched r=2 < under-rank r=1 (need enough rank): PASS\n"); pass = pass + 1 } else { m2_puts(" T1 under-rank: FAIL\n"); fail = fail + 1 }
267 if te2[0] < te4[0] { m2_puts(" T2 matched r=2 < over-rank r=4 (too much overfits): PASS\n"); pass = pass + 1 } else { m2_puts(" T2 over-rank: FAIL\n"); fail = fail + 1 }
268 if te2[0] < full_ref { m2_puts(" T3 matched r=2 < FULL (beats full at half VRAM): PASS\n"); pass = pass + 1 } else { m2_puts(" T3 vs full: FAIL\n"); fail = fail + 1 }
269
270 m2_puts("\n PASS="); m2_putn(pass); m2_puts("/3 ")
271 if fail == 0 { m2_puts("VERDICT=GREEN (U-curve: optimal bottleneck rank MATCHES the data's true rank, and there beats full at half the VRAM -- quality UP, VRAM down, MEASURED + not over-claimed)\n"); sys_exit(0); return 0 }
272 m2_puts("VERDICT=RED\n"); sys_exit(1); return 1
273}