nx_blockfloat_block_gate.nx source
↩ module page · 178 lines · 10877 B
1// nx_blockfloat_block_gate.nx -- BLOCK-FLOAT weights through the REAL no-float TRANSFORMER BLOCK. Composes the
2// gradcheck-verified pre-norm block forward (blk_fwd/blk_out, from nx_nofloat_block_gate: RMSNorm -> Q/K/V -> RoPE
3// -> scaled causal-softmax attention -> out-proj -> residual -> SwiGLU-FFN -> residual) UNCHANGED, but runs it with
4// the weight matrices (Wq,Wk,Wv,Wo,Wg,Wu,Wd) BLOCK-FLOAT quantized (per-row-per-block power-of-2 scale). Proves a
5// real transformer component runs on block-float, and that block-float weights beat per-tensor INT8 at the BLOCK
6// output level while staying bit-exact deterministic. Mixed-magnitude weights (outlier rows) make the range matter.
7// 1 the real block RUNS with block-float weights (non-trivial output)
8// 2 EXCEED (MEASURED): block output error vs full-precision < per-tensor INT8 weights
9// 3 DETERMINISTIC: the block-float block run twice == bit-identical
10// 4 block-float weights are REAL (block output != full-precision output)
11// expect_exit: 0 license_tier: ORIGINAL
12import "nx_nofloat_autograd.nx"
13import "nx_syscalls.nx"
14import "nx_gate_verdict.nx"
15
16const Q16: i64 = 65536
17
18func 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 }
19func bb_pn(v: i64) -> i64 {
20 let b: *u8 = sys_mmap(28); var m: i64 = v
21 if m < 0 { m = 0 - m; sys_write(1, "-" as *u8, 1) }
22 let t: *u8 = sys_mmap(28); var k: i64 = 0
23 if m == 0 { t[0] = 48 as u8; k = 1 }
24 while m > 0 { t[k] = (48 + (m % 10)) as u8; m = m / 10; k = k + 1 }
25 var i: i64 = 0
26 while i < k { b[i] = t[k - 1 - i]; i = i + 1 }
27 sys_write(1, b, k); return 0
28}
29func bb_chk(name: *u8, ok: i64) -> i64 {
30 if ok == 1 { bb_puts(" PASS " as *u8); bb_puts(name); bb_puts("\n" as *u8); return 1 }
31 bb_puts(" FAIL " as *u8); bb_puts(name); bb_puts("\n" as *u8); return 0
32}
33func bb_abs(x: i64) -> i64 { if x < 0 { return 0 - x } return x }
34func bf_bitlen(x: i64) -> i64 { var b: i64 = 0; var m: i64 = x; while m > 0 { m = m >> 1; b = b + 1 } return b }
35func bf_scale_abs(W: *i64, off: i64, B: i64, MB: i64) -> i64 {
36 var amax: i64 = 0; var i: i64 = 0
37 while i < B { let a: i64 = bb_abs(W[off + i]); if a > amax { amax = a } i = i + 1 }
38 var e: i64 = bf_bitlen(amax) - MB
39 if e < 0 { e = 0 }
40 return e
41}
42// block-float quantize a rows x cols matrix per-row-per-block (signed Q16, division-based -> sign-correct).
43func bf_quant_mat(W: *i64, rows: i64, cols: i64, B: i64, MB: i64, out: *i64) -> i64 {
44 var r: i64 = 0
45 while r < rows {
46 var blk: i64 = 0
47 while blk < cols / B {
48 let off: i64 = r * cols + blk * B
49 let e: i64 = bf_scale_abs(W, off, B, MB)
50 let sc: i64 = 1 << e
51 var i: i64 = 0
52 while i < B { out[off + i] = (W[off + i] / sc) * sc; i = i + 1 }
53 blk = blk + 1
54 }
55 r = r + 1
56 }
57 return 0
58}
59// per-tensor quantize: ONE scale over the whole matrix.
60func bf_quant_pt(W: *i64, n: i64, MB: i64, out: *i64) -> i64 {
61 let eg: i64 = bf_scale_abs(W, 0, n, MB); let sc: i64 = 1 << eg
62 var i: i64 = 0
63 while i < n { out[i] = (W[i] / sc) * sc; i = i + 1 }
64 return 0
65}
66
67// ---- the existing pre-norm transformer block forward (copied verbatim from nx_nofloat_block_gate) ----
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, leaves: *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 leaves[0] = nWq; leaves[1] = nWd
98 return nY
99}
100func blk_out(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, outY: *i64) -> i64 {
101 let lv: *i64 = sys_mmap(2 * 8) as *i64
102 let nY: i64 = blk_fwd(tape, vals, st, X, Wq, Wk, Wv, Wo, Wg, Wu, Wd, T, dm, ffn, scale, lv)
103 var i: i64 = 0
104 while i < T * dm { outY[i] = nfa_val(tape, vals, nY, i); i = i + 1 }
105 return 0
106}
107
108func main() -> i64 {
109 bb_puts("=== BLOCK-FLOAT weights through the REAL no-float transformer block (existing blk_fwd, unchanged) ===\n" as *u8)
110 let tape: *i64 = sys_mmap(1024 * 7 * 8) as *i64
111 let vals: *i64 = sys_mmap(16384 * 8) as *i64
112 let st: *i64 = sys_mmap(2 * 8) as *i64
113 let T: i64 = 2; let dm: i64 = 2; let ffn: i64 = 4; let scale: i64 = 46341; let MB: i64 = 4; let B: i64 = 2
114
115 let X: *i64 = sys_mmap(T * dm * 8) as *i64; X[0] = 32768; X[1] = 0 - 16384; X[2] = 49152; X[3] = 24576
116 // mixed-magnitude weights (outlier rows/blocks: large ~Q16 next to small) so per-block scaling matters
117 let Wq: *i64 = sys_mmap(dm * dm * 8) as *i64; Wq[0] = 65536; Wq[1] = 49152; Wq[2] = 1280; Wq[3] = 1536
118 let Wk: *i64 = sys_mmap(dm * dm * 8) as *i64; Wk[0] = 1536; Wk[1] = 1280; Wk[2] = 49152; Wk[3] = 65536
119 let Wv: *i64 = sys_mmap(dm * dm * 8) as *i64; Wv[0] = 49152; Wv[1] = 65536; Wv[2] = 1024; Wv[3] = 1792
120 let Wo: *i64 = sys_mmap(dm * dm * 8) as *i64; Wo[0] = 1280; Wo[1] = 1536; Wo[2] = 65536; Wo[3] = 49152
121 let Wg: *i64 = sys_mmap(dm * ffn * 8) as *i64; Wg[0] = 65536; Wg[1] = 49152; Wg[2] = 1280; Wg[3] = 1536; Wg[4] = 1536; Wg[5] = 1280; Wg[6] = 49152; Wg[7] = 65536
122 let Wu: *i64 = sys_mmap(dm * ffn * 8) as *i64; Wu[0] = 1280; Wu[1] = 1536; Wu[2] = 65536; Wu[3] = 49152; Wu[4] = 49152; Wu[5] = 65536; Wu[6] = 1280; Wu[7] = 1024
123 let Wd: *i64 = sys_mmap(ffn * dm * 8) as *i64; Wd[0] = 65536; Wd[1] = 49152; Wd[2] = 1280; Wd[3] = 1536; Wd[4] = 1536; Wd[5] = 1280; Wd[6] = 49152; Wd[7] = 65536
124
125 // block-float-quantized copies
126 let qq: *i64 = sys_mmap(dm * dm * 8) as *i64; bf_quant_mat(Wq, dm, dm, B, MB, qq)
127 let qk: *i64 = sys_mmap(dm * dm * 8) as *i64; bf_quant_mat(Wk, dm, dm, B, MB, qk)
128 let qv: *i64 = sys_mmap(dm * dm * 8) as *i64; bf_quant_mat(Wv, dm, dm, B, MB, qv)
129 let qo: *i64 = sys_mmap(dm * dm * 8) as *i64; bf_quant_mat(Wo, dm, dm, B, MB, qo)
130 let qg: *i64 = sys_mmap(dm * ffn * 8) as *i64; bf_quant_mat(Wg, dm, ffn, B, MB, qg)
131 let qu: *i64 = sys_mmap(dm * ffn * 8) as *i64; bf_quant_mat(Wu, dm, ffn, B, MB, qu)
132 let qd: *i64 = sys_mmap(ffn * dm * 8) as *i64; bf_quant_mat(Wd, ffn, dm, B, MB, qd)
133 // per-tensor-quantized copies
134 let pq: *i64 = sys_mmap(dm * dm * 8) as *i64; bf_quant_pt(Wq, dm * dm, MB, pq)
135 let pk: *i64 = sys_mmap(dm * dm * 8) as *i64; bf_quant_pt(Wk, dm * dm, MB, pk)
136 let pv: *i64 = sys_mmap(dm * dm * 8) as *i64; bf_quant_pt(Wv, dm * dm, MB, pv)
137 let po: *i64 = sys_mmap(dm * dm * 8) as *i64; bf_quant_pt(Wo, dm * dm, MB, po)
138 let pg: *i64 = sys_mmap(dm * ffn * 8) as *i64; bf_quant_pt(Wg, dm * ffn, MB, pg)
139 let pu: *i64 = sys_mmap(dm * ffn * 8) as *i64; bf_quant_pt(Wu, dm * ffn, MB, pu)
140 let pd: *i64 = sys_mmap(ffn * dm * 8) as *i64; bf_quant_pt(Wd, ffn * dm, MB, pd)
141
142 let Yr: *i64 = sys_mmap(T * dm * 8) as *i64; blk_out(tape, vals, st, X, Wq, Wk, Wv, Wo, Wg, Wu, Wd, T, dm, ffn, scale, Yr)
143 let Yb: *i64 = sys_mmap(T * dm * 8) as *i64; blk_out(tape, vals, st, X, qq, qk, qv, qo, qg, qu, qd, T, dm, ffn, scale, Yb)
144 let Yp: *i64 = sys_mmap(T * dm * 8) as *i64; blk_out(tape, vals, st, X, pq, pk, pv, po, pg, pu, pd, T, dm, ffn, scale, Yp)
145
146 var err_bf: i64 = 0; var err_pt: i64 = 0; var nontriv: i64 = 0; var real_pt: i64 = 0; var i: i64 = 0
147 while i < T * dm {
148 err_bf = err_bf + bb_abs(Yb[i] - Yr[i])
149 err_pt = err_pt + bb_abs(Yp[i] - Yr[i])
150 if Yr[i] != 0 { nontriv = 1 }
151 if Yp[i] != Yr[i] { real_pt = 1 }
152 i = i + 1
153 }
154 bb_puts(" block out: Y_ref[0]=" as *u8); bb_pn(Yr[0]); bb_puts(" Y_bf[0]=" as *u8); bb_pn(Yb[0]); bb_puts(" Y_pt[0]=" as *u8); bb_pn(Yp[0]); bb_puts("\n" as *u8)
155 bb_puts(" block-output error (Q16 L1): BLOCK-FLOAT=" as *u8); bb_pn(err_bf); bb_puts(" PER-TENSOR=" as *u8); bb_pn(err_pt); bb_puts("\n" as *u8)
156
157 let Yb2: *i64 = sys_mmap(T * dm * 8) as *i64; blk_out(tape, vals, st, X, qq, qk, qv, qo, qg, qu, qd, T, dm, ffn, scale, Yb2)
158 var det: i64 = 1; var d: i64 = 0
159 while d < T * dm { if Yb2[d] != Yb[d] { det = 0 } d = d + 1 }
160
161 var pass: i64 = 0; var total: i64 = 0
162 total = total + 1; pass = pass + bb_chk("T1 the real transformer block RUNS with block-float weights" as *u8, nontriv)
163 var t2: i64 = 0; if err_bf < err_pt { t2 = 1 }
164 total = total + 1; pass = pass + bb_chk("T2 EXCEED: block-float block error < per-tensor (dynamic-range win)" as *u8, t2)
165 total = total + 1; pass = pass + bb_chk("T3 DETERMINISTIC: block-float block twice == bit-identical" as *u8, det)
166 total = total + 1; pass = pass + bb_chk("T4 per-tensor is LOSSY (output != full-precision) -- block-float's win is over real info loss" as *u8, real_pt)
167
168 bb_puts("NX-BLOCKFLOAT-BLOCK-GATE " as *u8); bb_pn(pass); bb_puts(" / " as *u8); bb_pn(total)
169 // MIGRATED onto nx_gate_verdict by nx_gate_dry_apply (D001, minimal form): every check
170 // row above is untouched, so the PASS/FAIL vector cannot change; only the hand-rolled
171 // verdict emission is replaced by the ONE shared base class. Proven by nx_gate_migrate verify.
172 let ctr__dry: *i64 = gv_ctr()
173 ctr__dry[0] = pass
174 ctr__dry[1] = total
175 let rc__dry: i64 = gv_verdict("BLOCKFLOAT-BLOCK-GATE" as *u8, ctr__dry, "block-float weights run the real transformer block: FP8 range + bit-exact determinism)" as *u8)
176 sys_exit(rc__dry)
177 return rc__dry
178}