code wiki / _hdl_build / nx_fnet_mlm_gate.nx

nx_fnet_mlm_gate.nx source

↩ module page · 278 lines · 12959 B

1// nx_fnet_mlm_gate.nx -- GATE for MODEL-002: a tiny MASKED LANGUAGE MODEL on the FNet mixer. FNet's Fourier 2// mix is BIDIRECTIONAL (every position sees every other), so its native LM objective is masked-LM (BERT-style), 3// NOT causal next-token (which would let the model see the answer). Here: one position of a 4-token palindrome 4// [a,b,b,a] is replaced by a MASK token; the model fills it in -- which requires MIXING the mirror position into 5// the masked one (a pure bag-of-words model cannot). No attention anywhere. 6// tokens(+MASK) -> EMBED(trained) -> FNET mix -> relu FFN -> softmax-CE over the vocab at the masked slot 7// 8// G_train masked-token accuracy >= 7/8 AND final loss < first loss. 9// G_repro bit-exact: train twice -> identical accuracy + final-loss bits. 10// 11// Evidence -> knowledge/status/fnet_mlm.log (FNETMLMGATE authored=organ ... verdict=GREEN). license_tier: ORIGINAL 12import "nx_autograd_tensor.nx" 13import "nx_syscalls.nx" 14import "nx_gate_verdict.nx" 15 16const MV: i64 = 5 // input vocab: tokens 0..3 + MASK=4 17const MN: i64 = 4 // sequence length 18const MD: i64 = 4 // d_model 19const MH: i64 = 8 // FFN hidden 20const MC: i64 = 4 // output vocab (predict token 0..3) 21const MND: i64 = 16 // MN*MD 22 23const ML_LOG: *u8 = "knowledge/status/fnet_mlm.log" 24 25func ml_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 } 26func ml_wn(fd: i64, v: i64) -> i64 { 27 let bb: *u8 = sys_mmap(28); var m: i64 = v 28 if m < 0 { m = 0 - m; sys_write(fd, "-" as *u8, 1) } 29 let t: *u8 = sys_mmap(28); var k: i64 = 0 30 if m == 0 { t[0] = 48; k = 1 } 31 while m > 0 { t[k] = (48 + (m % 10)) as u8; m = m / 10; k = k + 1 } 32 var i: i64 = 0 33 while i < k { bb[i] = t[k - 1 - i]; i = i + 1 } 34 sys_write(fd, bb, k); return 0 35} 36 37func ml_embed(E: *i64, seq: *i64, soff: i64, xout: *i64) -> i64 { 38 var i: i64 = 0 39 while i < MN { 40 let tok: i64 = seq[soff + i] 41 var j: i64 = 0 42 while j < MD { xout[i * MD + j] = E[tok * MD + j]; j = j + 1 } 43 i = i + 1 44 } 45 return 0 46} 47 48func ml_fwd(tape: *i64, vals: *i64, st: *i64, x: *i64, nW1: i64, nb1: i64, nW2: i64, nb2: i64) -> i64 { 49 let xl: i64 = ta_leaf(tape, vals, st, MN, MD, x, 0) 50 let m: i64 = ta_fnet(tape, vals, st, xl) 51 let h: i64 = ta_relu(tape, vals, st, ta_vadd(tape, vals, st, ta_matvec(tape, vals, st, nW1, m), nb1)) 52 return ta_vadd(tape, vals, st, ta_matvec(tape, vals, st, nW2, h), nb2) 53} 54 55func ml_build(tape: *i64, vals: *i64, st: *i64, E: *i64, W1: *i64, b1: *i64, W2: *i64, b2: *i64, 56 seqs: *i64, labels: *i64, c8: *i64, wb: *i64, xl8: *i64) -> i64 { 57 st[0] = 0; st[1] = 0 58 let nW1: i64 = ta_leaf(tape, vals, st, MH, MND, W1, 0) 59 let nb1: i64 = ta_leaf(tape, vals, st, MH, 1, b1, 0) 60 let nW2: i64 = ta_leaf(tape, vals, st, MC, MH, W2, 0) 61 let nb2: i64 = ta_leaf(tape, vals, st, MC, 1, b2, 0) 62 wb[0] = nW1; wb[1] = nb1; wb[2] = nW2; wb[3] = nb2 63 let x: *i64 = (sys_mmap(MND * 8)) as *i64 64 let th: *i64 = (sys_mmap(MC * 8)) as *i64 65 var sumn: i64 = 0 - 1 66 var s: i64 = 0 67 while s < 8 { 68 ml_embed(E, seqs, s * MN, x) 69 let xl: i64 = ta_leaf(tape, vals, st, MN, MD, x, 0) 70 xl8[s] = xl 71 let m: i64 = ta_fnet(tape, vals, st, xl) 72 let h: i64 = ta_relu(tape, vals, st, ta_vadd(tape, vals, st, ta_matvec(tape, vals, st, nW1, m), nb1)) 73 let lo: i64 = ta_vadd(tape, vals, st, ta_matvec(tape, vals, st, nW2, h), nb2) 74 var j: i64 = 0 75 while j < MC { th[j] = TA_F32_ZERO; j = j + 1 } 76 th[labels[s]] = TA_F32_ONE 77 let tgt: i64 = ta_leaf(tape, vals, st, MC, 1, th, 0) 78 let ls: i64 = ta_softce(tape, vals, st, lo, tgt) 79 if s == 0 { sumn = ls } else { sumn = ta_vadd(tape, vals, st, sumn, ls) } 80 s = s + 1 81 } 82 let inv8: i64 = ta_leaf(tape, vals, st, 1, 1, c8, 0) 83 return ta_matvec(tape, vals, st, inv8, sumn) 84} 85 86func ml_predict(tape: *i64, vals: *i64, st: *i64, E: *i64, W1: *i64, b1: *i64, W2: *i64, b2: *i64, seqs: *i64, s: i64) -> i64 { 87 st[0] = 0; st[1] = 0 88 let nW1: i64 = ta_leaf(tape, vals, st, MH, MND, W1, 0) 89 let nb1: i64 = ta_leaf(tape, vals, st, MH, 1, b1, 0) 90 let nW2: i64 = ta_leaf(tape, vals, st, MC, MH, W2, 0) 91 let nb2: i64 = ta_leaf(tape, vals, st, MC, 1, b2, 0) 92 let x: *i64 = (sys_mmap(MND * 8)) as *i64 93 ml_embed(E, seqs, s * MN, x) 94 let lo: i64 = ml_fwd(tape, vals, st, x, nW1, nb1, nW2, nb2) 95 var best: i64 = 0 96 var bestv: i64 = ta_val(tape, vals, lo, 0) 97 var c: i64 = 1 98 while c < MC { 99 let v: i64 = ta_val(tape, vals, lo, c) 100 if nx_f32_gt(v, bestv) == 1 { bestv = v; best = c } 101 c = c + 1 102 } 103 return best 104} 105 106func ml_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 { 107 var i: i64 = 0 108 while i < n { 109 let gi: i64 = g[i] 110 m[i] = nx_f32_add(nx_f32_mul(beta1, m[i]), nx_f32_mul(om1, gi)) 111 v[i] = nx_f32_add(nx_f32_mul(beta2, v[i]), nx_f32_mul(om2, nx_f32_mul(gi, gi))) 112 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)))) 113 i = i + 1 114 } 115 return 0 116} 117 118func ml_train(tape: *i64, vals: *i64, grads: *i64, st: *i64, E: *i64, W1: *i64, b1: *i64, W2: *i64, b2: *i64, 119 seqs: *i64, labels: *i64, epochs: i64, lf: *i64, ll: *i64) -> i64 { 120 ta_det_init(E, MV * MD, 5) 121 ta_det_init(W1, MH * MND, 3) 122 ta_det_init(W2, MC * MH, 7) 123 var z: i64 = 0 124 while z < MH { b1[z] = TA_F32_ZERO; z = z + 1 } 125 z = 0 126 while z < MC { b2[z] = TA_F32_ZERO; z = z + 1 } 127 let mE: *i64 = (sys_mmap(MV * MD * 8)) as *i64; let vE: *i64 = (sys_mmap(MV * MD * 8)) as *i64 128 let mW1: *i64 = (sys_mmap(MH * MND * 8)) as *i64; let vW1: *i64 = (sys_mmap(MH * MND * 8)) as *i64 129 let mb1: *i64 = (sys_mmap(MH * 8)) as *i64; let vb1: *i64 = (sys_mmap(MH * 8)) as *i64 130 let mW2: *i64 = (sys_mmap(MC * MH * 8)) as *i64; let vW2: *i64 = (sys_mmap(MC * MH * 8)) as *i64 131 let mb2: *i64 = (sys_mmap(MC * 8)) as *i64; let vb2: *i64 = (sys_mmap(MC * 8)) as *i64 132 z = 0 133 while z < MV * MD { mE[z] = TA_F32_ZERO; vE[z] = TA_F32_ZERO; z = z + 1 } 134 z = 0 135 while z < MH * MND { mW1[z] = TA_F32_ZERO; vW1[z] = TA_F32_ZERO; z = z + 1 } 136 z = 0 137 while z < MH { mb1[z] = TA_F32_ZERO; vb1[z] = TA_F32_ZERO; z = z + 1 } 138 z = 0 139 while z < MC * MH { mW2[z] = TA_F32_ZERO; vW2[z] = TA_F32_ZERO; z = z + 1 } 140 z = 0 141 while z < MC { mb2[z] = TA_F32_ZERO; vb2[z] = TA_F32_ZERO; z = z + 1 } 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 c8: *i64 = (sys_mmap(8)) as *i64; c8[0] = ta_constf(1, 8) 147 let wb: *i64 = (sys_mmap(4 * 8)) as *i64 148 let xl8: *i64 = (sys_mmap(8 * 8)) as *i64 149 let gW1: *i64 = (sys_mmap(MH * MND * 8)) as *i64 150 let gb1: *i64 = (sys_mmap(MH * 8)) as *i64 151 let gW2: *i64 = (sys_mmap(MC * MH * 8)) as *i64 152 let gb2: *i64 = (sys_mmap(MC * 8)) as *i64 153 let dE: *i64 = (sys_mmap(MV * MD * 8)) as *i64 154 var ep: i64 = 0 155 while ep < epochs { 156 let loss: i64 = ml_build(tape, vals, st, E, W1, b1, W2, b2, seqs, labels, c8, wb, xl8) 157 ta_backward(tape, vals, grads, st[0], loss) 158 if ep == 0 { *lf = ta_val(tape, vals, loss, 0) } 159 *ll = ta_val(tape, vals, loss, 0) 160 var i: i64 = 0 161 while i < MH * MND { gW1[i] = ta_grad(tape, grads, wb[0], i); i = i + 1 } 162 i = 0 163 while i < MH { gb1[i] = ta_grad(tape, grads, wb[1], i); i = i + 1 } 164 i = 0 165 while i < MC * MH { gW2[i] = ta_grad(tape, grads, wb[2], i); i = i + 1 } 166 i = 0 167 while i < MC { gb2[i] = ta_grad(tape, grads, wb[3], i); i = i + 1 } 168 i = 0 169 while i < MV * MD { dE[i] = TA_F32_ZERO; i = i + 1 } 170 var s: i64 = 0 171 while s < 8 { 172 var pos: i64 = 0 173 while pos < MN { 174 let tok: i64 = seqs[s * MN + pos] 175 var j: i64 = 0 176 while j < MD { 177 dE[tok * MD + j] = nx_f32_add(dE[tok * MD + j], ta_grad(tape, grads, xl8[s], pos * MD + j)) 178 j = j + 1 179 } 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 ml_adamw(E, mE, vE, dE, MV * MD, lr, beta1, beta2, om1, om2, eps, c1, c2) 187 ml_adamw(W1, mW1, vW1, gW1, MH * MND, lr, beta1, beta2, om1, om2, eps, c1, c2) 188 ml_adamw(b1, mb1, vb1, gb1, MH, lr, beta1, beta2, om1, om2, eps, c1, c2) 189 ml_adamw(W2, mW2, vW2, gW2, MC * MH, lr, beta1, beta2, om1, om2, eps, c1, c2) 190 ml_adamw(b2, mb2, vb2, gb2, MC, lr, beta1, beta2, om1, om2, eps, c1, c2) 191 ep = ep + 1 192 } 193 return 0 194} 195 196func main() -> i64 { 197 var ok: i64 = 1 198 let tape: *i64 = (sys_mmap(2048 * 7 * 8)) as *i64 199 let vals: *i64 = (sys_mmap(8192 * 8)) as *i64 200 let grads: *i64 = (sys_mmap(8192 * 8)) as *i64 201 let st: *i64 = (sys_mmap(2 * 8)) as *i64 202 203 // 8 masked palindromes [a,b,b,a] (MASK=4 at one position); label = the true masked token (its mirror). 204 let seqs: *i64 = (sys_mmap(32 * 8)) as *i64 205 seqs[0]=0; seqs[1]=1; seqs[2]=1; seqs[3]=4 206 seqs[4]=4; seqs[5]=2; seqs[6]=2; seqs[7]=1 207 seqs[8]=2; seqs[9]=4; seqs[10]=3; seqs[11]=2 208 seqs[12]=3; seqs[13]=0; seqs[14]=4; seqs[15]=3 209 seqs[16]=0; seqs[17]=2; seqs[18]=2; seqs[19]=4 210 seqs[20]=4; seqs[21]=3; seqs[22]=3; seqs[23]=1 211 seqs[24]=2; seqs[25]=1; seqs[26]=4; seqs[27]=2 212 seqs[28]=3; seqs[29]=4; seqs[30]=0; seqs[31]=3 213 let labels: *i64 = (sys_mmap(8 * 8)) as *i64 214 labels[0]=0; labels[1]=1; labels[2]=3; labels[3]=0; labels[4]=0; labels[5]=1; labels[6]=1; labels[7]=0 215 216 let E: *i64 = (sys_mmap(MV * MD * 8)) as *i64 217 let W1: *i64 = (sys_mmap(MH * MND * 8)) as *i64 218 let b1: *i64 = (sys_mmap(MH * 8)) as *i64 219 let W2: *i64 = (sys_mmap(MC * MH * 8)) as *i64 220 let b2: *i64 = (sys_mmap(MC * 8)) as *i64 221 let lf: *i64 = (sys_mmap(8)) as *i64 222 let ll: *i64 = (sys_mmap(8)) as *i64 223 ml_train(tape, vals, grads, st, E, W1, b1, W2, b2, seqs, labels, 2000, lf, ll) 224 var acc: i64 = 0 225 var s: i64 = 0 226 while s < 8 { 227 if ml_predict(tape, vals, st, E, W1, b1, W2, b2, seqs, s) == labels[s] { acc = acc + 1 } 228 s = s + 1 229 } 230 var trainPass: i64 = 1 231 if acc < 7 { trainPass = 0 } 232 if nx_f32_lt(*ll, *lf) != 1 { trainPass = 0 } 233 if trainPass != 1 { ok = 0 } 234 235 let E2: *i64 = (sys_mmap(MV * MD * 8)) as *i64 236 let W1b: *i64 = (sys_mmap(MH * MND * 8)) as *i64 237 let b1b: *i64 = (sys_mmap(MH * 8)) as *i64 238 let W2b: *i64 = (sys_mmap(MC * MH * 8)) as *i64 239 let b2b: *i64 = (sys_mmap(MC * 8)) as *i64 240 let lf2: *i64 = (sys_mmap(8)) as *i64 241 let ll2: *i64 = (sys_mmap(8)) as *i64 242 ml_train(tape, vals, grads, st, E2, W1b, b1b, W2b, b2b, seqs, labels, 2000, lf2, ll2) 243 var acc2: i64 = 0 244 s = 0 245 while s < 8 { 246 if ml_predict(tape, vals, st, E2, W1b, b1b, W2b, b2b, seqs, s) == labels[s] { acc2 = acc2 + 1 } 247 s = s + 1 248 } 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(ML_LOG, 420) } 258 if out >= 0 { 259 ml_w(out, "FNETMLMGATE authored=organ model=masked-LM embed+fnet-mix+relu-ffn+softmaxCE task=palindrome-fill no-attention" as *u8) 260 ml_w(out, " | masked_token_accuracy=" as *u8); ml_wn(out, acc); ml_w(out, "/8" as *u8) 261 ml_w(out, " loss_first_milli=" as *u8); ml_wn(out, ta_f32_to_milli(*lf)) 262 ml_w(out, " loss_last_milli=" as *u8); ml_wn(out, ta_f32_to_milli(*ll)) 263 ml_w(out, " | bitexact_repro=" as *u8); ml_wn(out, reproPass) 264 if ok == 1 { ml_w(out, " verdict=GREEN\n" as *u8) } else { ml_w(out, " verdict=RED\n" as *u8) } 265 if fdi == 0 { sys_close(out) } 266 } 267 fdi = fdi - 1 268 } 269 // MIGRATED onto nx_gate_verdict by nx_gate_dry_apply (D001, minimal form): every check 270 // row above is untouched, so the PASS/FAIL vector cannot change; only the hand-rolled 271 // verdict emission is replaced by the ONE shared base class. Proven by nx_gate_migrate verify. 272 let ctr__dry: *i64 = gv_ctr() 273 ctr__dry[0] = ok 274 ctr__dry[1] = 1 275 let rc__dry: i64 = gv_verdict("FNET-MLM-GATE" as *u8, ctr__dry, "teeth unchanged; verdict emission migrated onto the shared base class" as *u8) 276 sys_exit(rc__dry) 277 return rc__dry 278}