code wiki / (root) / nx_blockfloat_block_gate.nx

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}