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}