code wiki / (root) / nx_q8_matmul_micro.nx

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}