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}