code wiki / (root) / nx_q4k_matmul_bench_test.nx

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}