nx_q8_matmul_micro.nx source
↩ module page · 146 lines · 6130 B
1// nx_q8_matmul_micro.nx -- DECISIVE micro-probe of the REAL production Q8_0
2// matmul path (nx_f32_lazy_weight_new_q8_0 + nx_f32_lazy_matmul ->
3// _lw_q8_0_matmul). The forward's MATMUL bucket is 72% of a decode token
4// (462ms), yet the ISOLATED nx_q8_0_simd_gate clocks this same kernel at
5// ~20 GFLOP/s (=> lm_head should be ~14ms, but the forward pays ~143ms):
6// a 10x in-situ gap. This probe localizes it:
7//
8// WARM same weight, m=1, matmul in a tight REP loop (weights become
9// cache-hot after pass 1) -> per-call us + GFLOP/s + bytes/s.
10// COLD N DISTINCT weight buffers (aggregate >> LLC), one matmul each,
11// NEVER reused -> the real-forward access pattern (every weight
12// read exactly once from DRAM).
13//
14// If WARM ~ 20 GFLOP/s and COLD << that -> the forward is COLD-WEIGHT-
15// STREAMING bound (fix = fewer bytes or batch tokens, NOT a wider kernel).
16// If WARM is ALSO slow -> per-call overhead (mmap/pack/dispatch) or the
17// kernel itself. Shapes: W_gate (k=896,n=4864) + lm_head band (k=896,
18// n=32768). license_tier: ORIGINAL expect_exit: 0
19import "nx_syscalls.nx"
20import "nx_tier.nx"
21import "nx_le.nx"
22import "nx_f32.nx"
23import "nx_f32_cvt.nx"
24import "nx_thread_pool.nx"
25import "nx_f32_lazy_weight.nx"
26import "nx_fmt.nx"
27
28const MK: i64 = 896 // hidden
29const Q8B: i64 = 34
30const Q8V: i64 = 32
31
32func mm_nl() -> i64 { fmt_puts("\n" as *u8); return 0 }
33func mm_lcg(s: i64) -> i64 { var v: i64 = s * 1103515245 + 12345; v = v & 2147483647; return v }
34
35// build a Q8_0 weight buffer of n rows x (k/32) blocks; returns ptr.
36func mm_weight(k: i64, n: i64, seed: i64) -> *u8 {
37 let bpr: i64 = (k / Q8V) * Q8B
38 let w: *u8 = sys_mmap(n * bpr)
39 var s: i64 = seed
40 var r: i64 = 0
41 while r < n {
42 var b: i64 = 0
43 while b < k / Q8V {
44 let off: i64 = r * bpr + b * Q8B
45 w[off + 0] = 0x00 as u8
46 w[off + 1] = 0x2C as u8 // f16 d ~ 0.0625
47 var q: i64 = 0
48 while q < 32 { s = mm_lcg(s); w[off + 2 + q] = ((s % 17) - 8) as u8; q = q + 1 }
49 b = b + 1
50 }
51 r = r + 1
52 }
53 return w
54}
55
56func mm_report(tag: *u8, k: i64, n: i64, us: i64, calls: i64) -> i64 {
57 var u: i64 = us
58 if u < 1 { u = 1 }
59 let per: i64 = us / calls
60 let flop: i64 = 2 * k * n * calls // MACs*2
61 let mflops: i64 = flop / u // us -> MFLOP/s
62 let bytes: i64 = (n * (k / Q8V) * Q8B) * calls
63 let mbps: i64 = bytes / u // us -> MB/s
64 fmt_puts(tag)
65 fmt_puts(" per_call_us="); fmt_putn(per)
66 fmt_puts(" MFLOPs="); fmt_putn(mflops)
67 fmt_puts(" MBps="); fmt_putn(mbps)
68 mm_nl()
69 return 0
70}
71
72func main() -> i64 {
73 let A: *i64 = sys_mmap(MK * 8) as *i64
74 var s: i64 = 4242
75 var i: i64 = 0
76 while i < MK { s = mm_lcg(s); A[i] = nx_i32_to_f32((s % 9) - 4); i = i + 1 }
77
78 // warm the shared pool once (excluded from timings).
79 let pool: *NxThreadPool = nx_lw_shared_pool()
80
81 // ---- W_gate shape (k=896, n=4864): POOL vs SINGLE-THREAD ----
82 let NG: i64 = 4864
83 let Wg: *u8 = mm_weight(MK, NG, 111)
84 let lwg: *NxF32LazyWeight = nx_f32_lazy_weight_new_q8_0(Wg, 0, NG, MK)
85 let Cg: *i64 = sys_mmap(NG * 8) as *i64
86 _lw_q8_0_matmul_pool_force(lwg, A, Cg, 1, MK, NG) // prime
87 let REPS: i64 = 40
88 let t0: i64 = sys_now_us()
89 var r0: i64 = 0
90 while r0 < REPS { _lw_q8_0_matmul_pool_force(lwg, A, Cg, 1, MK, NG); r0 = r0 + 1 }
91 let usw: i64 = sys_now_us() - t0
92 mm_report("POOL Wgate(896x4864)" as *u8, MK, NG, usw, REPS)
93 _lw_q8_0_matmul_st(lwg, A, Cg, 1, MK, NG) // prime
94 let t0s: i64 = sys_now_us()
95 var r0s: i64 = 0
96 while r0s < REPS { _lw_q8_0_matmul_st(lwg, A, Cg, 1, MK, NG); r0s = r0s + 1 }
97 let usws: i64 = sys_now_us() - t0s
98 mm_report("ST Wgate(896x4864)" as *u8, MK, NG, usws, REPS)
99
100 // CORRECTNESS: ST result must be BIT-IDENTICAL to pool (same per-column
101 // dequant-dot order; no reduction reordering). Guards the forward.
102 let Cp: *i64 = sys_mmap(NG * 8) as *i64
103 let Cs: *i64 = sys_mmap(NG * 8) as *i64
104 _lw_q8_0_matmul_pool_force(lwg, A, Cp, 1, MK, NG)
105 _lw_q8_0_matmul_st(lwg, A, Cs, 1, MK, NG)
106 var mism: i64 = 0
107 var ci: i64 = 0
108 while ci < NG { if Cp[ci] != Cs[ci] { mism = mism + 1 } ci = ci + 1 }
109 fmt_puts("BITEXACT pool==ST mismatches="); fmt_putn(mism); mm_nl()
110 if mism != 0 { fmt_puts("FATAL ST diverges from pool"); mm_nl(); return 77 }
111
112 // ---- lm_head band (k=896, n=32768): POOL vs ST ----
113 let NL: i64 = 32768
114 let Wl: *u8 = mm_weight(MK, NL, 222)
115 let lwl: *NxF32LazyWeight = nx_f32_lazy_weight_new_q8_0(Wl, 0, NL, MK)
116 let Cl: *i64 = sys_mmap(NL * 8) as *i64
117 _lw_q8_0_matmul_pool_force(lwl, A, Cl, 1, MK, NL)
118 let t2: i64 = sys_now_us()
119 var r2: i64 = 0
120 while r2 < 10 { _lw_q8_0_matmul_pool_force(lwl, A, Cl, 1, MK, NL); r2 = r2 + 1 }
121 mm_report("POOL lmhead(896x32768)" as *u8, MK, NL, sys_now_us() - t2, 10)
122 _lw_q8_0_matmul_st(lwl, A, Cl, 1, MK, NL)
123 let t2s: i64 = sys_now_us()
124 var r2s: i64 = 0
125 while r2s < 10 { _lw_q8_0_matmul_st(lwl, A, Cl, 1, MK, NL); r2s = r2s + 1 }
126 mm_report("ST lmhead(896x32768)" as *u8, MK, NL, sys_now_us() - t2s, 10)
127
128 // ---- tiny K/V-proj (k=896, n=128): POOL vs ST -- the dispatch case ----
129 let NK: i64 = 128
130 let Wk: *u8 = mm_weight(MK, NK, 333)
131 let lwk: *NxF32LazyWeight = nx_f32_lazy_weight_new_q8_0(Wk, 0, NK, MK)
132 let Ck: *i64 = sys_mmap(NK * 8) as *i64
133 _lw_q8_0_matmul_pool_force(lwk, A, Ck, 1, MK, NK)
134 let t3: i64 = sys_now_us()
135 var r3: i64 = 0
136 while r3 < 200 { _lw_q8_0_matmul_pool_force(lwk, A, Ck, 1, MK, NK); r3 = r3 + 1 }
137 mm_report("POOL kvproj(896x128)" as *u8, MK, NK, sys_now_us() - t3, 200)
138 _lw_q8_0_matmul_st(lwk, A, Ck, 1, MK, NK)
139 let t3s: i64 = sys_now_us()
140 var r3s: i64 = 0
141 while r3s < 200 { _lw_q8_0_matmul_st(lwk, A, Ck, 1, MK, NK); r3s = r3s + 1 }
142 mm_report("ST kvproj(896x128)" as *u8, MK, NK, sys_now_us() - t3s, 200)
143
144 fmt_puts("Q8_MATMUL_MICRO DONE"); mm_nl()
145 return 0
146}