code wiki / _hdl_build / nx_ssm_lm_gate.nx
nx_ssm_lm_gate.nx source
↩ module page · 271 lines · 12645 B
1// nx_ssm_lm_gate.nx -- GATE for MODEL-003: a CAUSAL AUTOREGRESSIVE next-token language model (GPT-shaped,
2// decoder) with NO attention. The causal SSM scan makes next-token prediction legal (each position's state
3// sees only the past):
4// tokens -> EMBED(trained) -> TA_SSM causal mix -> per-position SLICE -> shared relu-FFN head -> softmax-CE
5// predicting token t+1 from tokens 0..t.
6// AdamW trains the embedding + the SSM decay + the head jointly. Corpus = 4 cyclic-shift sequences over vocab 4
7// (a simple deterministic language, next = (cur+1) mod 4); 12 next-token predictions.
8//
9// G_train next-token accuracy >= 11/12 AND final loss < first loss.
10// G_repro bit-exact: train twice -> identical accuracy + final-loss bits.
11//
12// Evidence -> knowledge/status/ssm_lm.log (SSMLMGATE authored=organ ... verdict=GREEN). license_tier: ORIGINAL
13import "nx_autograd_tensor.nx"
14import "nx_syscalls.nx"
15
16const LV: i64 = 4 // vocab
17const LN: i64 = 4 // sequence length
18const LD: i64 = 4 // d_model
19const LH: i64 = 8 // FFN hidden
20const LP: i64 = 3 // predictions per sequence (positions 0..LN-2)
21const LM_LOG: *u8 = "knowledge/status/ssm_lm.log"
22
23func lm_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 }
24func lm_wn(fd: i64, v: i64) -> i64 {
25 let bb: *u8 = sys_mmap(28); var m: i64 = v
26 if m < 0 { m = 0 - m; sys_write(fd, "-" as *u8, 1) }
27 let t: *u8 = sys_mmap(28); var k: i64 = 0
28 if m == 0 { t[0] = 48; k = 1 }
29 while m > 0 { t[k] = (48 + (m % 10)) as u8; m = m / 10; k = k + 1 }
30 var i: i64 = 0
31 while i < k { bb[i] = t[k - 1 - i]; i = i + 1 }
32 sys_write(fd, bb, k); return 0
33}
34
35func lm_embed(E: *i64, seq: *i64, soff: i64, xout: *i64) -> i64 {
36 var i: i64 = 0
37 while i < LN {
38 let tok: i64 = seq[soff + i]
39 var j: i64 = 0
40 while j < LD { xout[i * LD + j] = E[tok * LD + j]; j = j + 1 }
41 i = i + 1
42 }
43 return 0
44}
45
46// head over one position's state vector [LD] -> logits [LV]
47func lm_head(tape: *i64, vals: *i64, st: *i64, row: i64, nW1: i64, nb1: i64, nW2: i64, nb2: i64) -> i64 {
48 let h: i64 = ta_relu(tape, vals, st, ta_vadd(tape, vals, st, ta_matvec(tape, vals, st, nW1, row), nb1))
49 return ta_vadd(tape, vals, st, ta_matvec(tape, vals, st, nW2, h), nb2)
50}
51
52func lm_build(tape: *i64, vals: *i64, st: *i64, E: *i64, ad: *i64, W1: *i64, b1: *i64, W2: *i64, b2: *i64,
53 seqs: *i64, c12: *i64, wb: *i64, xl4: *i64) -> i64 {
54 st[0] = 0; st[1] = 0
55 let na: i64 = ta_leaf(tape, vals, st, LD, 1, ad, 0)
56 let nW1: i64 = ta_leaf(tape, vals, st, LH, LD, W1, 0)
57 let nb1: i64 = ta_leaf(tape, vals, st, LH, 1, b1, 0)
58 let nW2: i64 = ta_leaf(tape, vals, st, LV, LH, W2, 0)
59 let nb2: i64 = ta_leaf(tape, vals, st, LV, 1, b2, 0)
60 wb[0] = na; wb[1] = nW1; wb[2] = nb1; wb[3] = nW2; wb[4] = nb2
61 let x: *i64 = (sys_mmap(LN * LD * 8)) as *i64
62 let th: *i64 = (sys_mmap(LV * 8)) as *i64
63 var sumn: i64 = 0 - 1
64 var s: i64 = 0
65 while s < 4 {
66 lm_embed(E, seqs, s * LN, x)
67 let xl: i64 = ta_leaf(tape, vals, st, LN, LD, x, 0)
68 xl4[s] = xl
69 let m: i64 = ta_ssm(tape, vals, st, na, xl)
70 var i: i64 = 0
71 while i < LP {
72 let row: i64 = ta_slice(tape, vals, st, m, i)
73 let lo: i64 = lm_head(tape, vals, st, row, nW1, nb1, nW2, nb2)
74 var jj: i64 = 0
75 while jj < LV { th[jj] = TA_F32_ZERO; jj = jj + 1 }
76 th[seqs[s * LN + i + 1]] = TA_F32_ONE
77 let tgt: i64 = ta_leaf(tape, vals, st, LV, 1, th, 0)
78 let ls: i64 = ta_softce(tape, vals, st, lo, tgt)
79 if sumn < 0 { sumn = ls } else { sumn = ta_vadd(tape, vals, st, sumn, ls) }
80 i = i + 1
81 }
82 s = s + 1
83 }
84 let inv: i64 = ta_leaf(tape, vals, st, 1, 1, c12, 0)
85 return ta_matvec(tape, vals, st, inv, sumn)
86}
87
88func lm_predict(tape: *i64, vals: *i64, st: *i64, E: *i64, ad: *i64, W1: *i64, b1: *i64, W2: *i64, b2: *i64, seqs: *i64, s: i64, pos: i64) -> i64 {
89 st[0] = 0; st[1] = 0
90 let na: i64 = ta_leaf(tape, vals, st, LD, 1, ad, 0)
91 let nW1: i64 = ta_leaf(tape, vals, st, LH, LD, W1, 0)
92 let nb1: i64 = ta_leaf(tape, vals, st, LH, 1, b1, 0)
93 let nW2: i64 = ta_leaf(tape, vals, st, LV, LH, W2, 0)
94 let nb2: i64 = ta_leaf(tape, vals, st, LV, 1, b2, 0)
95 let x: *i64 = (sys_mmap(LN * LD * 8)) as *i64
96 lm_embed(E, seqs, s * LN, x)
97 let xl: i64 = ta_leaf(tape, vals, st, LN, LD, x, 0)
98 let m: i64 = ta_ssm(tape, vals, st, na, xl)
99 let lo: i64 = lm_head(tape, vals, st, ta_slice(tape, vals, st, m, pos), nW1, nb1, nW2, nb2)
100 var best: i64 = 0
101 var bestv: i64 = ta_val(tape, vals, lo, 0)
102 var c: i64 = 1
103 while c < LV {
104 let v: i64 = ta_val(tape, vals, lo, c)
105 if nx_f32_gt(v, bestv) == 1 { bestv = v; best = c }
106 c = c + 1
107 }
108 return best
109}
110
111func lm_adamw(p: *i64, m: *i64, v: *i64, g: *i64, n: i64, lr: i64, beta1: i64, beta2: i64, om1: i64, om2: i64, eps: i64, c1: i64, c2: i64) -> i64 {
112 var i: i64 = 0
113 while i < n {
114 let gi: i64 = g[i]
115 m[i] = nx_f32_add(nx_f32_mul(beta1, m[i]), nx_f32_mul(om1, gi))
116 v[i] = nx_f32_add(nx_f32_mul(beta2, v[i]), nx_f32_mul(om2, nx_f32_mul(gi, gi)))
117 p[i] = nx_f32_sub(p[i], nx_f32_mul(lr, nx_f32_div(nx_f32_div(m[i], c1), nx_f32_add(nx_f32_sqrt(nx_f32_div(v[i], c2)), eps))))
118 i = i + 1
119 }
120 return 0
121}
122
123func lm_zero(a: *i64, n: i64) -> i64 { var i: i64 = 0; while i < n { a[i] = TA_F32_ZERO; i = i + 1 } return 0 }
124
125func lm_train(tape: *i64, vals: *i64, grads: *i64, st: *i64, E: *i64, ad: *i64, W1: *i64, b1: *i64, W2: *i64, b2: *i64,
126 seqs: *i64, epochs: i64, lf: *i64, ll: *i64) -> i64 {
127 ta_det_init(E, LV * LD, 5)
128 ta_det_init(W1, LH * LD, 3)
129 ta_det_init(W2, LV * LH, 7)
130 lm_zero(b1, LH); lm_zero(b2, LV)
131 var i: i64 = 0
132 while i < LD { ad[i] = ta_constf(1, 2); i = i + 1 } // decay init 0.5 (stable, in (0,1))
133 let mE: *i64 = (sys_mmap(LV * LD * 8)) as *i64; let vE: *i64 = (sys_mmap(LV * LD * 8)) as *i64
134 let mA: *i64 = (sys_mmap(LD * 8)) as *i64; let vA: *i64 = (sys_mmap(LD * 8)) as *i64
135 let mW1: *i64 = (sys_mmap(LH * LD * 8)) as *i64; let vW1: *i64 = (sys_mmap(LH * LD * 8)) as *i64
136 let mb1: *i64 = (sys_mmap(LH * 8)) as *i64; let vb1: *i64 = (sys_mmap(LH * 8)) as *i64
137 let mW2: *i64 = (sys_mmap(LV * LH * 8)) as *i64; let vW2: *i64 = (sys_mmap(LV * LH * 8)) as *i64
138 let mb2: *i64 = (sys_mmap(LV * 8)) as *i64; let vb2: *i64 = (sys_mmap(LV * 8)) as *i64
139 lm_zero(mE, LV*LD); lm_zero(vE, LV*LD); lm_zero(mA, LD); lm_zero(vA, LD)
140 lm_zero(mW1, LH*LD); lm_zero(vW1, LH*LD); lm_zero(mb1, LH); lm_zero(vb1, LH)
141 lm_zero(mW2, LV*LH); lm_zero(vW2, LV*LH); lm_zero(mb2, LV); lm_zero(vb2, LV)
142 let beta1: i64 = ta_constf(9, 10); let beta2: i64 = ta_constf(999, 1000)
143 let om1: i64 = ta_constf(1, 10); let om2: i64 = ta_constf(1, 1000)
144 let lr: i64 = ta_constf(1, 50); let eps: i64 = ta_constf(1, 100000000)
145 var b1t: i64 = TA_F32_ONE; var b2t: i64 = TA_F32_ONE
146 let c12: *i64 = (sys_mmap(8)) as *i64; c12[0] = ta_constf(1, 12)
147 let wb: *i64 = (sys_mmap(5 * 8)) as *i64
148 let xl4: *i64 = (sys_mmap(4 * 8)) as *i64
149 let gA: *i64 = (sys_mmap(LD * 8)) as *i64
150 let gW1: *i64 = (sys_mmap(LH * LD * 8)) as *i64
151 let gb1: *i64 = (sys_mmap(LH * 8)) as *i64
152 let gW2: *i64 = (sys_mmap(LV * LH * 8)) as *i64
153 let gb2: *i64 = (sys_mmap(LV * 8)) as *i64
154 let dE: *i64 = (sys_mmap(LV * LD * 8)) as *i64
155 var ep: i64 = 0
156 while ep < epochs {
157 let loss: i64 = lm_build(tape, vals, st, E, ad, W1, b1, W2, b2, seqs, c12, wb, xl4)
158 ta_backward(tape, vals, grads, st[0], loss)
159 if ep == 0 { *lf = ta_val(tape, vals, loss, 0) }
160 *ll = ta_val(tape, vals, loss, 0)
161 i = 0
162 while i < LD { gA[i] = ta_grad(tape, grads, wb[0], i); i = i + 1 }
163 i = 0
164 while i < LH * LD { gW1[i] = ta_grad(tape, grads, wb[1], i); i = i + 1 }
165 i = 0
166 while i < LH { gb1[i] = ta_grad(tape, grads, wb[2], i); i = i + 1 }
167 i = 0
168 while i < LV * LH { gW2[i] = ta_grad(tape, grads, wb[3], i); i = i + 1 }
169 i = 0
170 while i < LV { gb2[i] = ta_grad(tape, grads, wb[4], i); i = i + 1 }
171 lm_zero(dE, LV * LD)
172 var s: i64 = 0
173 while s < 4 {
174 var pos: i64 = 0
175 while pos < LN {
176 let tok: i64 = seqs[s * LN + pos]
177 var j: i64 = 0
178 while j < LD { dE[tok * LD + j] = nx_f32_add(dE[tok * LD + j], ta_grad(tape, grads, xl4[s], pos * LD + j)); j = j + 1 }
179 pos = pos + 1
180 }
181 s = s + 1
182 }
183 b1t = nx_f32_mul(b1t, beta1); b2t = nx_f32_mul(b2t, beta2)
184 let c1: i64 = nx_f32_sub(TA_F32_ONE, b1t); let c2: i64 = nx_f32_sub(TA_F32_ONE, b2t)
185 lm_adamw(E, mE, vE, dE, LV * LD, lr, beta1, beta2, om1, om2, eps, c1, c2)
186 lm_adamw(ad, mA, vA, gA, LD, lr, beta1, beta2, om1, om2, eps, c1, c2)
187 lm_adamw(W1, mW1, vW1, gW1, LH * LD, lr, beta1, beta2, om1, om2, eps, c1, c2)
188 lm_adamw(b1, mb1, vb1, gb1, LH, lr, beta1, beta2, om1, om2, eps, c1, c2)
189 lm_adamw(W2, mW2, vW2, gW2, LV * LH, lr, beta1, beta2, om1, om2, eps, c1, c2)
190 lm_adamw(b2, mb2, vb2, gb2, LV, lr, beta1, beta2, om1, om2, eps, c1, c2)
191 ep = ep + 1
192 }
193 return 0
194}
195
196func lm_accuracy(tape: *i64, vals: *i64, st: *i64, E: *i64, ad: *i64, W1: *i64, b1: *i64, W2: *i64, b2: *i64, seqs: *i64) -> i64 {
197 var acc: i64 = 0
198 var s: i64 = 0
199 while s < 4 {
200 var pos: i64 = 0
201 while pos < LP {
202 if lm_predict(tape, vals, st, E, ad, W1, b1, W2, b2, seqs, s, pos) == seqs[s * LN + pos + 1] { acc = acc + 1 }
203 pos = pos + 1
204 }
205 s = s + 1
206 }
207 return acc
208}
209
210func main() -> i64 {
211 var ok: i64 = 1
212 let tape: *i64 = (sys_mmap(2048 * 7 * 8)) as *i64
213 let vals: *i64 = (sys_mmap(16384 * 8)) as *i64
214 let grads: *i64 = (sys_mmap(16384 * 8)) as *i64
215 let st: *i64 = (sys_mmap(2 * 8)) as *i64
216
217 // corpus: cyclic shifts (next = (cur+1) mod 4)
218 let seqs: *i64 = (sys_mmap(16 * 8)) as *i64
219 seqs[0]=0; seqs[1]=1; seqs[2]=2; seqs[3]=3
220 seqs[4]=1; seqs[5]=2; seqs[6]=3; seqs[7]=0
221 seqs[8]=2; seqs[9]=3; seqs[10]=0; seqs[11]=1
222 seqs[12]=3; seqs[13]=0; seqs[14]=1; seqs[15]=2
223
224 let E: *i64 = (sys_mmap(LV * LD * 8)) as *i64
225 let ad: *i64 = (sys_mmap(LD * 8)) as *i64
226 let W1: *i64 = (sys_mmap(LH * LD * 8)) as *i64
227 let b1: *i64 = (sys_mmap(LH * 8)) as *i64
228 let W2: *i64 = (sys_mmap(LV * LH * 8)) as *i64
229 let b2: *i64 = (sys_mmap(LV * 8)) as *i64
230 let lf: *i64 = (sys_mmap(8)) as *i64
231 let ll: *i64 = (sys_mmap(8)) as *i64
232 lm_train(tape, vals, grads, st, E, ad, W1, b1, W2, b2, seqs, 2000, lf, ll)
233 let acc: i64 = lm_accuracy(tape, vals, st, E, ad, W1, b1, W2, b2, seqs)
234 var trainPass: i64 = 1
235 if acc < 11 { trainPass = 0 }
236 if nx_f32_lt(*ll, *lf) != 1 { trainPass = 0 }
237 if trainPass != 1 { ok = 0 }
238
239 let E2: *i64 = (sys_mmap(LV * LD * 8)) as *i64
240 let ad2: *i64 = (sys_mmap(LD * 8)) as *i64
241 let W1b: *i64 = (sys_mmap(LH * LD * 8)) as *i64
242 let b1b: *i64 = (sys_mmap(LH * 8)) as *i64
243 let W2b: *i64 = (sys_mmap(LV * LH * 8)) as *i64
244 let b2b: *i64 = (sys_mmap(LV * 8)) as *i64
245 let lf2: *i64 = (sys_mmap(8)) as *i64
246 let ll2: *i64 = (sys_mmap(8)) as *i64
247 lm_train(tape, vals, grads, st, E2, ad2, W1b, b1b, W2b, b2b, seqs, 2000, lf2, ll2)
248 let acc2: i64 = lm_accuracy(tape, vals, st, E2, ad2, W1b, b1b, W2b, b2b, seqs)
249 var reproPass: i64 = 1
250 if acc2 != acc { reproPass = 0 }
251 if *ll2 != *ll { reproPass = 0 }
252 if reproPass != 1 { ok = 0 }
253
254 var fdi: i64 = 1
255 while fdi >= 0 {
256 var out: i64 = 1
257 if fdi == 0 { out = sys_openat_append(LM_LOG, 420) }
258 if out >= 0 {
259 lm_w(out, "SSMLMGATE authored=organ model=causal-autoregressive-LM embed+ssm+slice+ffn+softmaxCE no-attention" as *u8)
260 lm_w(out, " | next_token_accuracy=" as *u8); lm_wn(out, acc); lm_w(out, "/12" as *u8)
261 lm_w(out, " loss_first_milli=" as *u8); lm_wn(out, ta_f32_to_milli(*lf))
262 lm_w(out, " loss_last_milli=" as *u8); lm_wn(out, ta_f32_to_milli(*ll))
263 lm_w(out, " | bitexact_repro=" as *u8); lm_wn(out, reproPass)
264 if ok == 1 { lm_w(out, " verdict=GREEN\n" as *u8) } else { lm_w(out, " verdict=RED\n" as *u8) }
265 if fdi == 0 { sys_close(out) }
266 }
267 fdi = fdi - 1
268 }
269 if ok == 1 { return 0 }
270 return 1
271}