code wiki / (root) / nx_nofloat_blockfloat_gemm_gate.nx

nx_nofloat_blockfloat_gemm_gate.nx source

↩ module page · 141 lines · 7191 B

1// nx_nofloat_blockfloat_gemm_gate.nx -- BLOCK-FLOAT weights through the no-float GEMM: the format earns its place 2// in the actual compute path (the rung after nx_nofloat_blockfloat_gate proved the format round-trips). The GEMM 3// computes DIRECTLY on the integer mantissas + per-block power-of-2 shifts (no dequant materialisation) -- 4// C[m] = sum_over_K-blocks b ( (sum_{k in b} A[m][k]*q_W[k]) << e_b ). Proves: 5// 1 the direct block-float GEMM == dequant-then-matmul, BYTE-EXACT (the on-mantissa compute is correct) 6// 2 EXCEED (MEASURED): block-float matmul error vs full-precision < per-tensor INT8 matmul error 7// 3 DETERMINISTIC: the block-float GEMM run twice is bit-identical (integer sums + shifts, no float) 8// 4 per-tensor scaling ZEROED the small weight block (lost its whole matmul contribution); block-float kept it 9// expect_exit: 0 license_tier: ORIGINAL 10import "nx_syscalls.nx" 11import "nx_gate_verdict.nx" 12 13func bf_puts(s: *u8) -> i64 { var n: i64 = 0; while s[n] != (0 as u8) { n = n + 1 } sys_write(1, s, n); return 0 } 14func bf_putn(v: i64) -> i64 { 15 let b: *u8 = sys_mmap(28); var m: i64 = v 16 if m < 0 { m = 0 - m; sys_write(1, "-" as *u8, 1) } 17 let t: *u8 = sys_mmap(28); var k: i64 = 0 18 if m == 0 { t[0] = 48 as u8; k = 1 } 19 while m > 0 { t[k] = (48 + (m % 10)) as u8; m = m / 10; k = k + 1 } 20 var i: i64 = 0 21 while i < k { b[i] = t[k - 1 - i]; i = i + 1 } 22 sys_write(1, b, k); return 0 23} 24func bf_chk(name: *u8, ok: i64) -> i64 { 25 if ok == 1 { bf_puts(" PASS " as *u8); bf_puts(name); bf_puts("\n" as *u8); return 1 } 26 bf_puts(" FAIL " as *u8); bf_puts(name); bf_puts("\n" as *u8); return 0 27} 28func bf_bitlen(x: i64) -> i64 { var b: i64 = 0; var m: i64 = x; while m > 0 { m = m >> 1; b = b + 1 } return b } 29func bf_absdiff(a: i64, b: i64) -> i64 { if a >= b { return a - b } return b - a } 30func bf_scale(v: *i64, off: i64, B: i64, M: i64) -> i64 { 31 var amax: i64 = 0; var i: i64 = 0 32 while i < B { if v[off + i] > amax { amax = v[off + i] } i = i + 1 } 33 var e: i64 = bf_bitlen(amax) - M 34 if e < 0 { e = 0 } 35 return e 36} 37 38// plain integer matrix(M x K) . vector(K) -> out(M). Full-precision reference path. 39func mm(A: *i64, W: *i64, M: i64, K: i64, out: *i64) -> i64 { 40 var m: i64 = 0 41 while m < M { 42 var acc: i64 = 0; var k: i64 = 0 43 while k < K { acc = acc + A[m * K + k] * W[k]; k = k + 1 } 44 out[m] = acc; m = m + 1 45 } 46 return 0 47} 48 49// block-float GEMM: compute on the integer mantissas q_W + apply each K-block's power-of-2 scale e_bf[b]. 50func bf_gemm(A: *i64, q_W: *i64, e_bf: *i64, M: i64, K: i64, B: i64, out: *i64) -> i64 { 51 var m: i64 = 0 52 while m < M { 53 var acc: i64 = 0 54 var b: i64 = 0 55 while b < K / B { 56 var bsum: i64 = 0; var k: i64 = 0 57 while k < B { let idx: i64 = b * B + k; bsum = bsum + A[m * K + idx] * q_W[idx]; k = k + 1 } 58 acc = acc + (bsum << e_bf[b]) 59 b = b + 1 60 } 61 out[m] = acc; m = m + 1 62 } 63 return 0 64} 65 66func main() -> i64 { 67 let K: i64 = 8 68 let M: i64 = 2 69 let B: i64 = 4 70 let MB: i64 = 4 71 bf_puts("=== BLOCK-FLOAT weights through the no-float GEMM -- compute on mantissas+scales, vs per-tensor INT8 ===\n" as *u8) 72 73 let W: *i64 = sys_mmap(8 * K) as *i64 74 W[0] = 3; W[1] = 5; W[2] = 7; W[3] = 9 75 W[4] = 100; W[5] = 200; W[6] = 150; W[7] = 120 76 let A: *i64 = sys_mmap(8 * M * K) as *i64 77 A[0] = 2; A[1] = 1; A[2] = 3; A[3] = 1; A[4] = 1; A[5] = 1; A[6] = 1; A[7] = 1 78 A[8] = 1; A[9] = 2; A[10] = 1; A[11] = 2; A[12] = 2; A[13] = 1; A[14] = 2; A[15] = 1 79 80 // ---- quantize W: block-float (per K-block scale) + per-tensor (one scale) ---- 81 let q_W: *i64 = sys_mmap(8 * K) as *i64 82 let W_bf: *i64 = sys_mmap(8 * K) as *i64 83 let e_bf: *i64 = sys_mmap(8 * (K / B)) as *i64 84 var blk: i64 = 0 85 while blk < K / B { 86 let off: i64 = blk * B 87 let e: i64 = bf_scale(W, off, B, MB) 88 e_bf[blk] = e 89 var i: i64 = 0 90 while i < B { q_W[off + i] = W[off + i] >> e; W_bf[off + i] = q_W[off + i] << e; i = i + 1 } 91 blk = blk + 1 92 } 93 let W_pt: *i64 = sys_mmap(8 * K) as *i64 94 let eg: i64 = bf_scale(W, 0, K, MB) 95 var i2: i64 = 0 96 while i2 < K { W_pt[i2] = (W[i2] >> eg) << eg; i2 = i2 + 1 } 97 98 // ---- the four matmuls ---- 99 let C_ref: *i64 = sys_mmap(8 * M) as *i64; mm(A, W, M, K, C_ref) // full precision 100 let C_bf: *i64 = sys_mmap(8 * M) as *i64; bf_gemm(A, q_W, e_bf, M, K, B, C_bf) // block-float (on mantissas) 101 let C_deq: *i64 = sys_mmap(8 * M) as *i64; mm(A, W_bf, M, K, C_deq) // dequant-then-matmul 102 let C_pt: *i64 = sys_mmap(8 * M) as *i64; mm(A, W_pt, M, K, C_pt) // per-tensor INT8 103 104 var err_bf: i64 = 0; var err_pt: i64 = 0; var t1: i64 = 1; var m: i64 = 0 105 while m < M { 106 err_bf = err_bf + bf_absdiff(C_bf[m], C_ref[m]) 107 err_pt = err_pt + bf_absdiff(C_pt[m], C_ref[m]) 108 if C_bf[m] != C_deq[m] { t1 = 0 } 109 m = m + 1 110 } 111 bf_puts(" C_ref=[" as *u8); bf_putn(C_ref[0]); bf_puts(" " as *u8); bf_putn(C_ref[1]); bf_puts("] C_bf=[" as *u8); bf_putn(C_bf[0]); bf_puts(" " as *u8); bf_putn(C_bf[1]); bf_puts("] C_pt=[" as *u8); bf_putn(C_pt[0]); bf_puts(" " as *u8); bf_putn(C_pt[1]); bf_puts("]\n" as *u8) 112 bf_puts(" matmul error: BLOCK-FLOAT=" as *u8); bf_putn(err_bf); bf_puts(" PER-TENSOR=" as *u8); bf_putn(err_pt); bf_puts("\n" as *u8) 113 114 // determinism: recompute the block-float GEMM, compare bit-for-bit 115 let C_bf2: *i64 = sys_mmap(8 * M) as *i64; bf_gemm(A, q_W, e_bf, M, K, B, C_bf2) 116 var det_ok: i64 = 1; var d: i64 = 0 117 while d < M { if C_bf2[d] != C_bf[d] { det_ok = 0 } d = d + 1 } 118 119 // per-tensor zeroed the small block; block-float kept it 120 var t4: i64 = 1; var s: i64 = 0 121 while s < B { if W_bf[s] != W[s] { t4 = 0 } if W_pt[s] != 0 { t4 = 0 } s = s + 1 } 122 123 var pass: i64 = 0 124 var total: i64 = 0 125 total = total + 1; pass = pass + bf_chk("T1 direct block-float GEMM == dequant-then-matmul (byte-exact)" as *u8, t1) 126 var t2: i64 = 0; if err_bf < err_pt { t2 = 1 } 127 total = total + 1; pass = pass + bf_chk("T2 EXCEED: block-float matmul error < per-tensor (dynamic-range win)" as *u8, t2) 128 total = total + 1; pass = pass + bf_chk("T3 DETERMINISTIC: block-float GEMM twice == bit-identical" as *u8, det_ok) 129 total = total + 1; pass = pass + bf_chk("T4 per-tensor ZEROED small block; block-float kept it" as *u8, t4) 130 131 bf_puts("NX-NOFLOAT-BLOCKFLOAT-GEMM-GATE " as *u8); bf_putn(pass); bf_puts(" / " as *u8); bf_putn(total) 132 // MIGRATED onto nx_gate_verdict by nx_gate_dry_apply (D001, minimal form): every check 133 // row above is untouched, so the PASS/FAIL vector cannot change; only the hand-rolled 134 // verdict emission is replaced by the ONE shared base class. Proven by nx_gate_migrate verify. 135 let ctr__dry: *i64 = gv_ctr() 136 ctr__dry[0] = pass 137 ctr__dry[1] = total 138 let rc__dry: i64 = gv_verdict("NOFLOAT-BLOCKFLOAT-GEMM-GATE" as *u8, ctr__dry, "block-float weights in the GEMM: range of FP8 + bit-exact determinism, computed on mantissas)" as *u8) 139 sys_exit(rc__dry) 140 return rc__dry 141}