code wiki / (root) / nx_blockfloat_lm_gate.nx

nx_blockfloat_lm_gate.nx source

↩ module page · 203 lines · 11893 B

1// nx_blockfloat_lm_gate.nx -- FULL LM FORWARD with block-float weights: the first rung of the model-assembly arc. 2// token ids -> EMBED (lookup E) -> transformer block (the gated blk_fwd: RMSNorm/attn/SwiGLU-FFN/residuals) -> 3// final RMSNorm -> LM HEAD (matmul W_lm) -> logits -> argmax = next token. ALL weights (E, block, W_lm) block-float. 4// Proven on a tiny synthetic Qwen-shaped model so a REAL Qwen-0.6B's weights+config slot into the same forward. 5// 1 the LM forward RUNS: ids -> logits -> a valid next-token id in [0,vocab) 6// 2 EXCEED: block-float logits error vs full-precision < per-tensor INT8 7// 3 DETERMINISTIC: block-float logits twice == bit-identical 8// 4 block-float PRESERVES the prediction: its argmax next-token == full-precision's (per-tensor may flip it) 9// expect_exit: 0 license_tier: ORIGINAL 10import "nx_nofloat_autograd.nx" 11import "nx_syscalls.nx" 12import "nx_gate_verdict.nx" 13 14func bb_puts(s: *u8) -> i64 { var n: i64 = 0; while s[n] != (0 as u8) { n = n + 1 } sys_write(1, s, n); return 0 } 15func bb_pn(v: i64) -> i64 { 16 let b: *u8 = sys_mmap(28); var m: i64 = v 17 if m < 0 { m = 0 - m; sys_write(1, "-" as *u8, 1) } 18 let t: *u8 = sys_mmap(28); var k: i64 = 0 19 if m == 0 { t[0] = 48 as u8; k = 1 } 20 while m > 0 { t[k] = (48 + (m % 10)) as u8; m = m / 10; k = k + 1 } 21 var i: i64 = 0 22 while i < k { b[i] = t[k - 1 - i]; i = i + 1 } 23 sys_write(1, b, k); return 0 24} 25func bb_chk(name: *u8, ok: i64) -> i64 { 26 if ok == 1 { bb_puts(" PASS " as *u8); bb_puts(name); bb_puts("\n" as *u8); return 1 } 27 bb_puts(" FAIL " as *u8); bb_puts(name); bb_puts("\n" as *u8); return 0 28} 29func bb_abs(x: i64) -> i64 { if x < 0 { return 0 - x } return x } 30func bf_bitlen(x: i64) -> i64 { var b: i64 = 0; var m: i64 = x; while m > 0 { m = m >> 1; b = b + 1 } return b } 31func bf_scale_abs(W: *i64, off: i64, B: i64, MB: i64) -> i64 { 32 var amax: i64 = 0; var i: i64 = 0 33 while i < B { let a: i64 = bb_abs(W[off + i]); if a > amax { amax = a } i = i + 1 } 34 var e: i64 = bf_bitlen(amax) - MB; if e < 0 { e = 0 } return e 35} 36func bf_quant_mat(W: *i64, rows: i64, cols: i64, B: i64, MB: i64, out: *i64) -> i64 { 37 var r: i64 = 0 38 while r < rows { 39 var blk: i64 = 0 40 while blk < cols / B { 41 let off: i64 = r * cols + blk * B 42 let e: i64 = bf_scale_abs(W, off, B, MB); let sc: i64 = 1 << e 43 var i: i64 = 0 44 while i < B { out[off + i] = (W[off + i] / sc) * sc; i = i + 1 } 45 blk = blk + 1 46 } 47 r = r + 1 48 } 49 return 0 50} 51func bf_quant_pt(W: *i64, n: i64, MB: i64, out: *i64) -> i64 { 52 let eg: i64 = bf_scale_abs(W, 0, n, MB); let sc: i64 = 1 << eg 53 var i: i64 = 0 54 while i < n { out[i] = (W[i] / sc) * sc; i = i + 1 } 55 return 0 56} 57func fill_mixed(W: *i64, n: i64, B: i64, sgn: i64) -> i64 { 58 var i: i64 = 0 59 while i < n { 60 let blk: i64 = i / B; let pos: i64 = i % B 61 if (blk % 2) == 0 { W[i] = sgn * (65536 - pos * 16384) } else { W[i] = sgn * (1280 + pos * 256) } 62 i = i + 1 63 } 64 return 0 65} 66 67// the gated pre-norm block forward (copied from nx_nofloat_block_gate): attention + SwiGLU-FFN + residuals. 68func blk_fwd(tape: *i64, vals: *i64, st: *i64, X: *i64, Wq: *i64, Wk: *i64, Wv: *i64, Wo: *i64, Wg: *i64, Wu: *i64, Wd: *i64, T: i64, dm: i64, ffn: i64, scale: i64) -> i64 { 69 st[0] = 0; st[1] = 0 70 let nX: i64 = nfa_leaf(tape, vals, st, T, dm, X, 0) 71 let nWq: i64 = nfa_leaf(tape, vals, st, dm, dm, Wq, 0) 72 let nWk: i64 = nfa_leaf(tape, vals, st, dm, dm, Wk, 0) 73 let nWv: i64 = nfa_leaf(tape, vals, st, dm, dm, Wv, 0) 74 let nWo: i64 = nfa_leaf(tape, vals, st, dm, dm, Wo, 0) 75 let nWg: i64 = nfa_leaf(tape, vals, st, dm, ffn, Wg, 0) 76 let nWu: i64 = nfa_leaf(tape, vals, st, dm, ffn, Wu, 0) 77 let nWd: i64 = nfa_leaf(tape, vals, st, ffn, dm, Wd, 0) 78 let nXn: i64 = nfa_rmsnorm_rows(tape, vals, st, nX) 79 let nQ: i64 = nfa_matmul(tape, vals, st, nXn, nWq) 80 let nK: i64 = nfa_matmul(tape, vals, st, nXn, nWk) 81 let nV: i64 = nfa_matmul(tape, vals, st, nXn, nWv) 82 let nQr: i64 = nfa_rope(tape, vals, st, nQ) 83 let nKr: i64 = nfa_rope(tape, vals, st, nK) 84 let nS: i64 = nfa_matmul_nt(tape, vals, st, nQr, nKr) 85 let nSs: i64 = nfa_cmul(tape, vals, st, nS, scale) 86 let nA: i64 = nfa_softmax_rows(tape, vals, st, nSs, 1) 87 let nO: i64 = nfa_matmul(tape, vals, st, nA, nV) 88 let nOp: i64 = nfa_matmul(tape, vals, st, nO, nWo) 89 let nH: i64 = nfa_vadd(tape, vals, st, nX, nOp) 90 let nHn: i64 = nfa_rmsnorm_rows(tape, vals, st, nH) 91 let nG: i64 = nfa_matmul(tape, vals, st, nHn, nWg) 92 let nU: i64 = nfa_matmul(tape, vals, st, nHn, nWu) 93 let nSg: i64 = nfa_silu(tape, vals, st, nG) 94 let nHs: i64 = nfa_hadamard(tape, vals, st, nSg, nU) 95 let nD: i64 = nfa_matmul(tape, vals, st, nHs, nWd) 96 let nY: i64 = nfa_vadd(tape, vals, st, nH, nD) 97 return nY 98} 99 100// embed: X[t] = E[ids[t]] (row lookup). E is vocab x dm. 101func embed(E: *i64, ids: *i64, T: i64, dm: i64, X: *i64) -> i64 { 102 var t: i64 = 0 103 while t < T { var d: i64 = 0; while d < dm { X[t * dm + d] = E[ids[t] * dm + d]; d = d + 1 } t = t + 1 } 104 return 0 105} 106// full LM forward -> logits (T x vocab). Block + final RMSNorm + LM head (matmul with W_lm: dm x vocab). 107func lm_logits(tape: *i64, vals: *i64, st: *i64, X: *i64, Wq: *i64, Wk: *i64, Wv: *i64, Wo: *i64, Wg: *i64, Wu: *i64, Wd: *i64, Wlm: *i64, T: i64, dm: i64, ffn: i64, vocab: i64, scale: i64, outLog: *i64) -> i64 { 108 let nY: i64 = blk_fwd(tape, vals, st, X, Wq, Wk, Wv, Wo, Wg, Wu, Wd, T, dm, ffn, scale) 109 let nFn: i64 = nfa_rmsnorm_rows(tape, vals, st, nY) 110 let nWlm: i64 = nfa_leaf(tape, vals, st, dm, vocab, Wlm, 0) 111 let nLog: i64 = nfa_matmul(tape, vals, st, nFn, nWlm) 112 var i: i64 = 0 113 while i < T * vocab { outLog[i] = nfa_val(tape, vals, nLog, i); i = i + 1 } 114 return 0 115} 116func argmax(logits: *i64, off: i64, n: i64) -> i64 { 117 var bi: i64 = 0; var i: i64 = 1 118 while i < n { if logits[off + i] > logits[off + bi] { bi = i } i = i + 1 } 119 return bi 120} 121func l1err(a: *i64, b: *i64, n: i64) -> i64 { var e: i64 = 0; var i: i64 = 0; while i < n { e = e + bb_abs(a[i] - b[i]); i = i + 1 } return e } 122 123func main() -> i64 { 124 bb_puts("=== FULL LM FORWARD with block-float weights (token ids -> embed -> block -> norm -> LM head -> token) ===\n" as *u8) 125 let tape: *i64 = sys_mmap(2048 * 7 * 8) as *i64 126 let vals: *i64 = sys_mmap(32768 * 8) as *i64 127 let st: *i64 = sys_mmap(2 * 8) as *i64 128 let vocab: i64 = 8; let dm: i64 = 4; let ffn: i64 = 8; let T: i64 = 3; let scale: i64 = 46341; let MB: i64 = 4; let B: i64 = 2 129 130 let ids: *i64 = sys_mmap(T * 8) as *i64; ids[0] = 1; ids[1] = 3; ids[2] = 5 131 let E: *i64 = sys_mmap(vocab * dm * 8) as *i64; fill_mixed(E, vocab * dm, B, 1) 132 let Wq: *i64 = sys_mmap(dm * dm * 8) as *i64; fill_mixed(Wq, dm * dm, B, 1) 133 let Wk: *i64 = sys_mmap(dm * dm * 8) as *i64; fill_mixed(Wk, dm * dm, B, 0 - 1) 134 let Wv: *i64 = sys_mmap(dm * dm * 8) as *i64; fill_mixed(Wv, dm * dm, B, 1) 135 let Wo: *i64 = sys_mmap(dm * dm * 8) as *i64; fill_mixed(Wo, dm * dm, B, 0 - 1) 136 let Wg: *i64 = sys_mmap(dm * ffn * 8) as *i64; fill_mixed(Wg, dm * ffn, B, 1) 137 let Wu: *i64 = sys_mmap(dm * ffn * 8) as *i64; fill_mixed(Wu, dm * ffn, B, 0 - 1) 138 let Wd: *i64 = sys_mmap(ffn * dm * 8) as *i64; fill_mixed(Wd, ffn * dm, B, 1) 139 let Wlm: *i64 = sys_mmap(dm * vocab * 8) as *i64; fill_mixed(Wlm, dm * vocab, B, 1) 140 141 // block-float quant of every weight matrix 142 let qE: *i64 = sys_mmap(vocab*dm*8) as *i64; bf_quant_mat(E, vocab, dm, B, MB, qE) 143 let qq: *i64 = sys_mmap(dm*dm*8) as *i64; bf_quant_mat(Wq, dm, dm, B, MB, qq) 144 let qk: *i64 = sys_mmap(dm*dm*8) as *i64; bf_quant_mat(Wk, dm, dm, B, MB, qk) 145 let qv: *i64 = sys_mmap(dm*dm*8) as *i64; bf_quant_mat(Wv, dm, dm, B, MB, qv) 146 let qo: *i64 = sys_mmap(dm*dm*8) as *i64; bf_quant_mat(Wo, dm, dm, B, MB, qo) 147 let qg: *i64 = sys_mmap(dm*ffn*8) as *i64; bf_quant_mat(Wg, dm, ffn, B, MB, qg) 148 let qu: *i64 = sys_mmap(dm*ffn*8) as *i64; bf_quant_mat(Wu, dm, ffn, B, MB, qu) 149 let qd: *i64 = sys_mmap(ffn*dm*8) as *i64; bf_quant_mat(Wd, ffn, dm, B, MB, qd) 150 let qlm: *i64 = sys_mmap(dm*vocab*8) as *i64; bf_quant_mat(Wlm, dm, vocab, B, MB, qlm) 151 // per-tensor quant 152 let pE: *i64 = sys_mmap(vocab*dm*8) as *i64; bf_quant_pt(E, vocab*dm, MB, pE) 153 let pq: *i64 = sys_mmap(dm*dm*8) as *i64; bf_quant_pt(Wq, dm*dm, MB, pq) 154 let pk: *i64 = sys_mmap(dm*dm*8) as *i64; bf_quant_pt(Wk, dm*dm, MB, pk) 155 let pv: *i64 = sys_mmap(dm*dm*8) as *i64; bf_quant_pt(Wv, dm*dm, MB, pv) 156 let po: *i64 = sys_mmap(dm*dm*8) as *i64; bf_quant_pt(Wo, dm*dm, MB, po) 157 let pg: *i64 = sys_mmap(dm*ffn*8) as *i64; bf_quant_pt(Wg, dm*ffn, MB, pg) 158 let pu: *i64 = sys_mmap(dm*ffn*8) as *i64; bf_quant_pt(Wu, dm*ffn, MB, pu) 159 let pd: *i64 = sys_mmap(ffn*dm*8) as *i64; bf_quant_pt(Wd, ffn*dm, MB, pd) 160 let plm: *i64 = sys_mmap(dm*vocab*8) as *i64; bf_quant_pt(Wlm, dm*vocab, MB, plm) 161 162 let Xr: *i64 = sys_mmap(T*dm*8) as *i64; embed(E, ids, T, dm, Xr) 163 let Xb: *i64 = sys_mmap(T*dm*8) as *i64; embed(qE, ids, T, dm, Xb) 164 let Xp: *i64 = sys_mmap(T*dm*8) as *i64; embed(pE, ids, T, dm, Xp) 165 let Lr: *i64 = sys_mmap(T*vocab*8) as *i64; lm_logits(tape,vals,st,Xr,Wq,Wk,Wv,Wo,Wg,Wu,Wd,Wlm,T,dm,ffn,vocab,scale,Lr) 166 let Lb: *i64 = sys_mmap(T*vocab*8) as *i64; lm_logits(tape,vals,st,Xb,qq,qk,qv,qo,qg,qu,qd,qlm,T,dm,ffn,vocab,scale,Lb) 167 let Lp: *i64 = sys_mmap(T*vocab*8) as *i64; lm_logits(tape,vals,st,Xp,pq,pk,pv,po,pg,pu,pd,plm,T,dm,ffn,vocab,scale,Lp) 168 169 let tok_r: i64 = argmax(Lr, (T-1)*vocab, vocab) 170 let tok_b: i64 = argmax(Lb, (T-1)*vocab, vocab) 171 let tok_p: i64 = argmax(Lp, (T-1)*vocab, vocab) 172 let err_bf: i64 = l1err(Lb, Lr, T*vocab) 173 let err_pt: i64 = l1err(Lp, Lr, T*vocab) 174 bb_puts(" next-token id: full-precision="); bb_pn(tok_r); bb_puts(" block-float="); bb_pn(tok_b); bb_puts(" per-tensor="); bb_pn(tok_p); bb_puts("\n" as *u8) 175 bb_puts(" logits error (Q16 L1): BLOCK-FLOAT="); bb_pn(err_bf); bb_puts(" PER-TENSOR="); bb_pn(err_pt); bb_puts("\n" as *u8) 176 177 let Lb2: *i64 = sys_mmap(T*vocab*8) as *i64; lm_logits(tape,vals,st,Xb,qq,qk,qv,qo,qg,qu,qd,qlm,T,dm,ffn,vocab,scale,Lb2) 178 var det: i64 = 1; var d: i64 = 0 179 while d < T*vocab { if Lb2[d] != Lb[d] { det = 0 } d = d + 1 } 180 var t1: i64 = 0; if tok_r >= 0 { if tok_r < vocab { t1 = 1 } } 181 var nontriv: i64 = 0; var z: i64 = 0 182 while z < T*vocab { if Lr[z] != 0 { nontriv = 1 } z = z + 1 } 183 if nontriv == 0 { t1 = 0 } 184 185 var pass: i64 = 0; var total: i64 = 0 186 total = total + 1; pass = pass + bb_chk("T1 LM forward runs: ids -> logits -> valid next-token id" as *u8, t1) 187 var t2: i64 = 0; if err_bf < err_pt { t2 = 1 } 188 total = total + 1; pass = pass + bb_chk("T2 EXCEED: block-float logits error < per-tensor (model-level dynamic-range win)" as *u8, t2) 189 total = total + 1; pass = pass + bb_chk("T3 DETERMINISTIC: block-float logits twice == bit-identical" as *u8, det) 190 var t4: i64 = 0; if tok_b == tok_r { t4 = 1 } 191 total = total + 1; pass = pass + bb_chk("T4 block-float PRESERVES the prediction (its argmax == full-precision's)" as *u8, t4) 192 193 bb_puts("NX-BLOCKFLOAT-LM-GATE "); bb_pn(pass); bb_puts(" / "); bb_pn(total) 194 // MIGRATED onto nx_gate_verdict by nx_gate_dry_apply (D001, minimal form): every check 195 // row above is untouched, so the PASS/FAIL vector cannot change; only the hand-rolled 196 // verdict emission is replaced by the ONE shared base class. Proven by nx_gate_migrate verify. 197 let ctr__dry: *i64 = gv_ctr() 198 ctr__dry[0] = pass 199 ctr__dry[1] = total 200 let rc__dry: i64 = gv_verdict("BLOCKFLOAT-LM-GATE" as *u8, ctr__dry, "a full LM forward runs on block-float: token in -> token out, deterministic -- the model-assembly skeleton)" as *u8) 201 sys_exit(rc__dry) 202 return rc__dry 203}