code wiki / _hdl_build / nx_ssm_lm_gate.nx

nx_ssm_lm_gate.nx source

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