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}