code wiki / (root) / nx_q4k_matmul_rate.nx

nx_q4k_matmul_rate.nx source

↩ module page · 74 lines · 3336 B

1// nx_q4k_matmul_rate.nx -- isolate the SSE Q4_K matmul rate (no 491MB model load) to pinpoint why the LM 2// forward is slow (408s/8tok). Times nx_f32_q4k_matmul on a real LM size (W_gate: m=8, k=896, n=4864) on 3// DUMMY Q4_K bytes (timing is layout-correct regardless of byte values). MFLOP/s tells us: ~30 = matmul is 4// OVERHEAD-bound (dequant + i64-boxing + indexing dominate; scalar SSE can't help) -> packed SIMD/codegen is 5// the lever; ~300 = matmul is fast and the forward's slowness is elsewhere (attention/softmax/rmsnorm). 6import "nx_syscalls.nx" 7import "nx_f32_q4k_matmul.nx" 8 9func bn_puts(s: *u8) -> i64 { var n: i64 = 0; while s[n] != (0 as u8) { n = n + 1 } sys_write(1, s, n); return 0 } 10func bn_putn(v: i64) -> i64 { 11 if v == 0 { sys_write(1, "0" as *u8, 1); return 0 } 12 var m: i64 = v 13 if m < 0 { sys_write(1, "-" as *u8, 1); m = 0 - m } 14 let d: *u8 = sys_mmap(24); var k: i64 = 0 15 while m > 0 { d[k] = (48 + (m % 10)) as u8; m = m / 10; k = k + 1 } 16 var j: i64 = k - 1 17 while j >= 0 { sys_write(1, ((d as i64)+j) as *u8, 1); j = j - 1 } 18 return 0 19} 20 21func main() -> i64 { 22 let m: i64 = 8 23 let k: i64 = 896 24 let n: i64 = 4864 // = 19 * 256 (Q4_K super-block aligned), a real Qwen W_gate dim 25 bn_puts("=== Q4_K matmul micro-bench (SSE) m=8 k=896 n=4864 (real W_gate size) ===\n") 26 27 let A: *i64 = sys_mmap(m * k * 8) as *i64 28 var a: i64 = 0 29 while a < m * k { A[a] = 0x3F800000; a = a + 1 } // 1.0f bits 30 let bpr: i64 = (n / 256) * 144 31 let total_b: i64 = k * bpr 32 let B: *u8 = sys_mmap(total_b) 33 var b: i64 = 0 34 while b < total_b { B[b] = (b & 0xff) as u8; b = b + 1 } 35 let C: *i64 = sys_mmap(m * n * 8) as *i64 36 37 var reps: i64 = 1 38 var dt: i64 = 0 39 var go: i64 = 1 40 while go == 1 { 41 let t0: i64 = sys_now_ms() 42 var r: i64 = 0 43 while r < reps { nx_f32_q4k_matmul(A, B, 0, C, m, k, n); r = r + 1 } 44 let t1: i64 = sys_now_ms() 45 dt = t1 - t0 46 if dt >= 250 { go = 0 } else { if reps >= 256 { go = 0 } else { reps = reps * 2 } } 47 } 48 if dt <= 0 { dt = 1 } 49 let macs: i64 = m * k * n 50 let mflops: i64 = (2 * macs * reps) / (dt * 1000) 51 bn_puts(" WARM (reuse 1 weight slab, cache-resident): "); bn_putn(mflops); bn_puts(" MFLOP/s ms/matmul="); bn_putn(dt / reps); bn_puts("\n") 52 53 // COLD: cycle through a pool >> L3 cache so each weight read MISSES cache, like the real forward's 491MB. 54 let NPOOL: i64 = 16 55 let pool: *u8 = sys_mmap(NPOOL * total_b) 56 var pb: i64 = 0 57 while pb < NPOOL * total_b { pool[pb] = (pb & 0xff) as u8; pb = pb + 1 } 58 let cold_reps: i64 = 24 59 let ct0: i64 = sys_now_ms() 60 var cr: i64 = 0 61 while cr < cold_reps { 62 nx_f32_q4k_matmul(A, pool, (cr % NPOOL) * total_b, C, m, k, n) 63 cr = cr + 1 64 } 65 let ct1: i64 = sys_now_ms() 66 var cdt: i64 = ct1 - ct0 67 if cdt <= 0 { cdt = 1 } 68 let cmflops: i64 = (2 * macs * cold_reps) / (cdt * 1000) 69 bn_puts(" COLD ("); bn_putn(NPOOL * total_b / 1048576); bn_puts("MB pool > L3, cache-cold): "); bn_putn(cmflops) 70 bn_puts(" MFLOP/s ms/matmul="); bn_putn(cdt / cold_reps); bn_puts("\n") 71 bn_puts(" => COLL<<WARM => forward is MEMORY-bound (locality lever); COLD~=WARM => codegen-bound (nx_cc lever). forward eff ~15.\n") 72 sys_exit(0) 73 return 0 74}