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}