nx_q4k_matmul_bench_test.nx source
↩ module page · 188 lines · 6095 B
1// nx_q4k_matmul_bench_test.nx -- the fused-vs-materialized bench.
2
3import "nx_syscalls.nx"
4import "nx_tier.nx"
5import "nx_le.nx"
6import "nx_strconv.nx"
7import "nx_gguf.nx"
8import "nx_gguf_load.nx"
9import "nx_dequant_iter.nx"
10import "nx_q4k_matmul.nx"
11import "nx_q4k_matmul_bench.nx"
12
13func _emit(fd: i64, label: *u8, label_len: nx_int,
14 k_dim: i64, ns_per_run: i64, ns_per_element: i64) -> i64 {
15 let line: *u8 = sys_mmap(160)
16 var lo: i64 = 0
17 var i: nx_int = 0
18 while i < label_len { line[lo] = label[i]; lo = lo + 1; i = i + 1 }
19 line[lo] = 0x09; lo = lo + 1
20
21 let dec: *u8 = sys_mmap(32)
22 var n: i64 = nx_strconv_format_i64(k_dim, dec)
23 var j: i64 = 0
24 while j < n { line[lo] = dec[j]; lo = lo + 1; j = j + 1 }
25 line[lo] = 0x09; lo = lo + 1
26
27 n = nx_strconv_format_i64(ns_per_run, dec)
28 j = 0
29 while j < n { line[lo] = dec[j]; lo = lo + 1; j = j + 1 }
30 line[lo] = 0x09; lo = lo + 1
31
32 n = nx_strconv_format_i64(ns_per_element, dec)
33 j = 0
34 while j < n { line[lo] = dec[j]; lo = lo + 1; j = j + 1 }
35 line[lo] = 0x0A; lo = lo + 1
36
37 return sys_write(fd, line, lo)
38}
39
40// Build N pseudo-random Q4_K blocks.
41func _build_blocks(buf: *u8, n: i64) -> nx_int {
42 var b: i64 = 0
43 while b < n {
44 let base: i64 = b * NX_GL_Q4_K_BPB
45 let pat: i64 = b - (b / 3) * 3
46 var d_raw: i64 = 0x3C00
47 if pat == 0 { d_raw = 0x3800 }
48 if pat == 2 { d_raw = 0x4000 }
49 nx_le_write_u16(buf, base + 0, d_raw)
50 nx_le_write_u16(buf, base + 2, 0x0000)
51 var i: i64 = 0
52 while i < 12 {
53 buf[base + 4 + i] = (b * 7 + i * 3) & 0xFF
54 i = i + 1
55 }
56 var jj: i64 = 0
57 while jj < 128 {
58 buf[base + 16 + jj] = (b * 11 + jj * 5 + 1) & 0xFF
59 jj = jj + 1
60 }
61 b = b + 1
62 }
63 return 0
64}
65
66func main() -> i64 {
67 let n_blocks: i64 = NX_QMB_N_BLOCKS
68 let k_dim: i64 = n_blocks * 256
69
70 let buf: *u8 = sys_mmap(n_blocks * NX_GL_Q4_K_BPB + 64)
71 _build_blocks(buf, n_blocks)
72
73 let col: *i64 = sys_mmap(k_dim * 8) as *i64
74 var ci: nx_int = 0
75 while ci < k_dim {
76 col[ci] = (ci * 17 + 3) & 0x7FFF // bounded Q10 ints
77 ci = ci + 1
78 }
79
80 let scratch: *i64 = sys_mmap(k_dim * 8) as *i64
81 let it: *NxQ4KBlockIter = nx_q4k_iter_alloc()
82
83 // ----- Warmup -----
84 nx_gguf_dequant_q4_k(buf, 0, k_dim, scratch)
85 var warm_sum: i64 = 0
86 var wi: nx_int = 0
87 while wi < k_dim {
88 warm_sum = warm_sum + scratch[wi] * col[wi]
89 wi = wi + 1
90 }
91 let _warm2: i64 = nx_q4k_dot_row_col(buf, 0, n_blocks, col, it)
92
93 // ----- Sanity: results match -----
94 let truth_a: i64 = warm_sum
95 let truth_b: i64 = _warm2
96 if truth_a != truth_b { return 20 }
97
98 // ----- Path A: dequant then manual dot -----
99 let t_a0: *i64 = sys_mmap(16) as *i64
100 sys_clock_gettime_mono(t_a0)
101 let a0_s: i64 = t_a0[0]
102 let a0_n: i64 = t_a0[1]
103
104 var sink_a: i64 = 0
105 var r: i64 = 0
106 while r < NX_QMB_K_RUNS {
107 nx_gguf_dequant_q4_k(buf, 0, k_dim, scratch)
108 var dot_a: i64 = 0
109 var di: nx_int = 0
110 while di < k_dim {
111 dot_a = dot_a + scratch[di] * col[di]
112 di = di + 1
113 }
114 sink_a = sink_a + dot_a
115 r = r + 1
116 }
117 let t_a1: *i64 = sys_mmap(16) as *i64
118 sys_clock_gettime_mono(t_a1)
119 let ns_a_total: i64 = (t_a1[0] - a0_s) * 1000000000 + (t_a1[1] - a0_n)
120 let ns_a_per_run: i64 = ns_a_total / NX_QMB_K_RUNS
121 let ns_a_per_elt: i64 = ns_a_per_run / k_dim
122
123 // ----- Path B: fused dot -----
124 let t_b0: *i64 = sys_mmap(16) as *i64
125 sys_clock_gettime_mono(t_b0)
126 let b0_s: i64 = t_b0[0]
127 let b0_n: i64 = t_b0[1]
128
129 var sink_b: i64 = 0
130 var rb: i64 = 0
131 while rb < NX_QMB_K_RUNS {
132 sink_b = sink_b + nx_q4k_dot_row_col(buf, 0, n_blocks, col, it)
133 rb = rb + 1
134 }
135 let t_b1: *i64 = sys_mmap(16) as *i64
136 sys_clock_gettime_mono(t_b1)
137 let ns_b_total: i64 = (t_b1[0] - b0_s) * 1000000000 + (t_b1[1] - b0_n)
138 let ns_b_per_run: i64 = ns_b_total / NX_QMB_K_RUNS
139 let ns_b_per_elt: i64 = ns_b_per_run / k_dim
140
141 // ----- Sinks must equal (sanity) -----
142 if sink_a != sink_b { return 30 }
143
144 // ----- Write TSV -----
145 let path: *u8 = sys_mmap(64)
146 path[0]=0x2F; path[1]=0x74; path[2]=0x6D; path[3]=0x70; path[4]=0x2F
147 path[5]=0x6E; path[6]=0x78; path[7]=0x5F
148 path[8]=0x71; path[9]=0x34; path[10]=0x6B; path[11]=0x5F
149 path[12]=0x6D; path[13]=0x61; path[14]=0x74; path[15]=0x6D
150 path[16]=0x75; path[17]=0x6C; path[18]=0x5F
151 path[19]=0x62; path[20]=0x65; path[21]=0x6E; path[22]=0x63
152 path[23]=0x68; path[24]=0x2E; path[25]=0x74; path[26]=0x73
153 path[27]=0x76 // ".tsv"
154 path[28]=0
155 // Full: /tmp/nx_q4k_matmul_bench.tsv
156
157 let fd: i64 = sys_openat_wr(path, 0x1A4)
158 if fd < 0 { return 50 }
159
160 // Header.
161 let hdr: *u8 = sys_mmap(96)
162 hdr[0]=0x70; hdr[1]=0x61; hdr[2]=0x74; hdr[3]=0x68 // "path"
163 hdr[4]=0x09
164 hdr[5]=0x6B; hdr[6]=0x5F; hdr[7]=0x64; hdr[8]=0x69
165 hdr[9]=0x6D // "k_dim"
166 hdr[10]=0x09
167 hdr[11]=0x6E; hdr[12]=0x73; hdr[13]=0x5F; hdr[14]=0x72
168 hdr[15]=0x75; hdr[16]=0x6E // "ns_run"
169 hdr[17]=0x09
170 hdr[18]=0x6E; hdr[19]=0x73; hdr[20]=0x5F; hdr[21]=0x65
171 hdr[22]=0x6C; hdr[23]=0x74 // "ns_elt"
172 hdr[24]=0x0A
173 sys_write(fd, hdr, 25)
174
175 let lab_mat: *u8 = sys_mmap(16)
176 lab_mat[0]=0x6D; lab_mat[1]=0x61; lab_mat[2]=0x74
177 lab_mat[3]=0x65; lab_mat[4]=0x72; lab_mat[5]=0x69
178 lab_mat[6]=0x61; lab_mat[7]=0x6C // "material"
179 _emit(fd, lab_mat, 8, k_dim, ns_a_per_run, ns_a_per_elt)
180
181 let lab_fus: *u8 = sys_mmap(16)
182 lab_fus[0]=0x66; lab_fus[1]=0x75; lab_fus[2]=0x73
183 lab_fus[3]=0x65; lab_fus[4]=0x64 // "fused"
184 _emit(fd, lab_fus, 5, k_dim, ns_b_per_run, ns_b_per_elt)
185
186 sys_close(fd)
187 return 0
188}