code wiki / (root) / nx_q5k_dot_simd.nx

nx_q5k_dot_simd.nx source

↩ module page · 254 lines · 9977 B

1// nx_q5k_dot_simd.nx -- SIMD-accelerated Q5_K fused dot (vpmaddwd), bit-exact vs a scalar Q5_K dot. 2// 3// Extends the perf-exceed to the Q5_K layer matmuls (attn_v/output/ffn_down). Q5_K = Q4_K + a 5th bit from 4// qh: q5 = q4 + (qh_bit ? 16 : 0). Reformulation is the same: dot += d1*sq - m0*sc, sq = Σ q5*col via 5// vpmaddwd (q5 i16 0..31, col i16 Q10). Read-once + sc-precompute like nx_q4k_dot_simd. Includes an inline 6// scalar Q5_K dot as the bit-exact oracle. Verified on real blk.0.attn_v + measured. 7// license_tier: ORIGINAL 8import "nx_syscalls.nx" 9import "nx_tier.nx" 10import "nx_le.nx" 11import "nx_strconv.nx" 12import "nx_tensor.nx" 13import "nx_gguf.nx" 14import "nx_gguf_load.nx" 15import "nx_gguf_meta.nx" 16import "nx_placement.nx" 17import "nx_gguf_load_lazy.nx" 18import "nx_clock.nx" 19 20func q5s_pack4(a: i64, b: i64, c: i64, d: i64) -> i64 { 21 return (a & 0xFFFF) | ((b & 0xFFFF) << 16) | ((c & 0xFFFF) << 32) | ((d & 0xFFFF) << 48) 22} 23 24func q5s_hsum(acc: *i64) -> i64 { 25 var sum: i64 = 0 26 var i: i64 = 0 27 while i < 4 { 28 let v: i64 = acc[i] 29 var lo: i64 = v & 0xFFFFFFFF 30 if lo >= 0x80000000 { lo = lo - 0x100000000 } 31 var hi: i64 = (v >> 32) & 0xFFFFFFFF 32 if hi >= 0x80000000 { hi = hi - 0x100000000 } 33 sum = sum + lo + hi 34 i = i + 1 35 } 36 return sum 37} 38 39// unpack (sc,m) for sub-block j from the 12-byte packed scales at scales_off 40func q5s_scmin(buf: *u8, scales_off: i64, j: i64, out: *i64) -> i64 { 41 if j < 4 { 42 out[0] = nx_le_read_u8(buf, scales_off + j) & 0x3F 43 out[1] = nx_le_read_u8(buf, scales_off + j + 4) & 0x3F 44 } else { 45 let k: i64 = j - 4 46 out[0] = ((nx_le_read_u8(buf, scales_off + k) >> 6) << 4) | (nx_le_read_u8(buf, scales_off + 8 + k) & 0x0F) 47 out[1] = ((nx_le_read_u8(buf, scales_off + 4 + k) >> 6) << 4) | (nx_le_read_u8(buf, scales_off + 8 + k) >> 4) 48 } 49 return 0 50} 51 52// scalar Q5_K dot oracle: Σ (d*sc*q5 - dmin*m) * col, in Q34. 53func q5k_dot_scalar(buf: *u8, base_off: i64, n_blocks: i64, col_q10: *i64, sm: *i64) -> i64 { 54 var dot: i64 = 0 55 var blk: i64 = 0 56 while blk < n_blocks { 57 let sb: i64 = base_off + blk * 176 58 let d_q24: i64 = _gguf_f16_to_q24(nx_le_read_u16(buf, sb)) 59 let dmin_q24: i64 = _gguf_f16_to_q24(nx_le_read_u16(buf, sb + 2)) 60 let scales_off: i64 = sb + 4 61 let qh_off: i64 = sb + 16 62 let qs_off: i64 = sb + 48 63 let col_base: i64 = blk * 256 64 var g: i64 = 0 65 while g < 4 { 66 let is0: i64 = g + g 67 let is1: i64 = is0 + 1 68 q5s_scmin(buf, scales_off, is0, sm) 69 let d1: i64 = d_q24 * sm[0] 70 let m0v: i64 = dmin_q24 * sm[1] 71 q5s_scmin(buf, scales_off, is1, sm) 72 let d2: i64 = d_q24 * sm[0] 73 let m1v: i64 = dmin_q24 * sm[1] 74 let grp: i64 = qs_off + g * 32 75 let u1: i64 = 1 << (g + g) 76 let u2: i64 = 1 << (g + g + 1) 77 var l: i64 = 0 78 while l < 32 { 79 let byte_v: i64 = nx_le_read_u8(buf, grp + l) 80 let qh_l: i64 = nx_le_read_u8(buf, qh_off + l) 81 var q_lo: i64 = byte_v & 0x0F 82 var q_hi: i64 = byte_v >> 4 83 if (qh_l & u1) != 0 { q_lo = q_lo + 16 } 84 if (qh_l & u2) != 0 { q_hi = q_hi + 16 } 85 dot = dot + (d1 * q_lo - m0v) * col_q10[col_base + is0 * 32 + l] 86 dot = dot + (d2 * q_hi - m1v) * col_q10[col_base + is1 * 32 + l] 87 l = l + 1 88 } 89 g = g + 1 90 } 91 blk = blk + 1 92 } 93 return dot 94} 95 96// SIMD Q5_K dot: sq = Σ q5*col via vpmaddwd; sc from precompute. 97func q5k_dot_simd(buf: *u8, base_off: i64, n_blocks: i64, col_i16: *i64, 98 qpk: *i64, qhi: *i64, acc: *i64, sc_pre: *i64, sm: *i64) -> i64 { 99 var dot: i64 = 0 100 var blk: i64 = 0 101 while blk < n_blocks { 102 let sb: i64 = base_off + blk * 176 103 let d_q24: i64 = _gguf_f16_to_q24(nx_le_read_u16(buf, sb)) 104 let dmin_q24: i64 = _gguf_f16_to_q24(nx_le_read_u16(buf, sb + 2)) 105 let scales_off: i64 = sb + 4 106 let qh_off: i64 = sb + 16 107 let qs_off: i64 = sb + 48 108 let col_base: i64 = blk * 256 109 var g: i64 = 0 110 while g < 4 { 111 let is0: i64 = g + g 112 let is1: i64 = is0 + 1 113 q5s_scmin(buf, scales_off, is0, sm) 114 let d1: i64 = d_q24 * sm[0] 115 let m0v: i64 = dmin_q24 * sm[1] 116 q5s_scmin(buf, scales_off, is1, sm) 117 let d2: i64 = d_q24 * sm[0] 118 let m1v: i64 = dmin_q24 * sm[1] 119 let grp: i64 = qs_off + g * 32 120 let u1: i64 = 1 << (g + g) 121 let u2: i64 = 1 << (g + g + 1) 122 123 var jl: i64 = 0 124 while jl < 32 { 125 let byte_v: i64 = nx_le_read_u8(buf, grp + jl) 126 let qh_l: i64 = nx_le_read_u8(buf, qh_off + jl) 127 var q_lo: i64 = byte_v & 0x0F 128 var q_hi: i64 = byte_v >> 4 129 if (qh_l & u1) != 0 { q_lo = q_lo + 16 } 130 if (qh_l & u2) != 0 { q_hi = q_hi + 16 } 131 let w: i64 = jl / 4 132 let sh: i64 = (jl - w * 4) * 16 133 if sh == 0 { qpk[w] = q_lo; qhi[w] = q_hi } 134 else { qpk[w] = qpk[w] | (q_lo << sh); qhi[w] = qhi[w] | (q_hi << sh) } 135 jl = jl + 1 136 } 137 138 acc[0] = 0; acc[1] = 0; acc[2] = 0; acc[3] = 0 139 let cl0: i64 = (col_i16 as i64) + ((col_base + is0 * 32) / 4) * 8 140 __i16x16_madd(acc as *i64, qpk as *i64, cl0 as *i64) 141 __i16x16_madd(acc as *i64, ((qpk as i64) + 32) as *i64, (cl0 + 32) as *i64) 142 let sq_lo: i64 = q5s_hsum(acc) 143 acc[0] = 0; acc[1] = 0; acc[2] = 0; acc[3] = 0 144 let cl1: i64 = (col_i16 as i64) + ((col_base + is1 * 32) / 4) * 8 145 __i16x16_madd(acc as *i64, qhi as *i64, cl1 as *i64) 146 __i16x16_madd(acc as *i64, ((qhi as i64) + 32) as *i64, (cl1 + 32) as *i64) 147 let sq_hi: i64 = q5s_hsum(acc) 148 149 dot = dot + d1 * sq_lo - m0v * sc_pre[blk * 8 + is0] + d2 * sq_hi - m1v * sc_pre[blk * 8 + is1] 150 g = g + 1 151 } 152 blk = blk + 1 153 } 154 return dot 155} 156 157func q5s_emit(fd: i64, key: *u8, kl: i64, v: i64) -> i64 { 158 let line: *u8 = sys_mmap(64) 159 var lo: i64 = 0 160 var i: i64 = 0 161 while i < kl { line[lo] = key[i]; lo = lo + 1; i = i + 1 } 162 line[lo] = 0x3D; lo = lo + 1 163 let dec: *u8 = sys_mmap(32) 164 let nd: i64 = nx_strconv_format_i64(v, dec) 165 var k: i64 = 0 166 while k < nd { line[lo] = dec[k]; lo = lo + 1; k = k + 1 } 167 line[lo] = 0x0A; lo = lo + 1 168 return sys_write(fd, line, lo) 169} 170 171func main() -> i64 { 172 let path: *u8 = "/mnt/c/Users/elder/elder-ai-platform/models/unified/text_encoder/Huihui-Qwen3-4B-Instruct-2507-abliterated-Q4_K_M.gguf" as *u8 173 let fd: i64 = sys_openat_rd(path) 174 if fd < 0 { return 30 } 175 let CAP: i64 = 1153433600 176 let buf: *u8 = sys_mmap(CAP) 177 var total: i64 = 0 178 var go: i64 = 1 179 while go == 1 { 180 let r: i64 = sys_read(fd, ((buf as i64) + total) as *u8, CAP - total) 181 if r <= 0 { go = 0 } else { total = total + r; if total >= CAP { go = 0 } } 182 } 183 sys_close(fd) 184 if total < 100000000 { return 31 } 185 let hdr: *NxGgufHeader = sys_mmap(NX_GGUF_HDR_BYTES) as *NxGgufHeader 186 if nx_gguf_parse(buf, total, hdr) != NX_GGUF_OK { return 40 } 187 let vi: nx_int = nx_gguf_find_tensor(hdr, "blk.0.attn_v.weight" as *u8, 19) 188 if vi < 0 { return 60 } 189 let ti: *NxGgufTensorInfo = nx_gguf_tensor_at(hdr, vi) 190 if ti.ggml_type != 13 { return 61 } 191 let IN: i64 = ti.dim_0 192 let w_off: i64 = hdr.data_off + ti.offset 193 let n_blocks: i64 = IN / 256 194 if w_off + n_blocks * 176 > total { return 63 } 195 196 let col: *i64 = sys_mmap(IN * 8) as *i64 197 let col_i16: *i64 = sys_mmap(IN * 2) as *i64 198 var i: i64 = 0 199 while i < IN { col[i] = 1024 + (i - (i / 5) * 5) * 256; i = i + 1 } 200 var jj: i64 = 0 201 while jj < IN / 4 { 202 col_i16[jj] = q5s_pack4(col[jj * 4], col[jj * 4 + 1], col[jj * 4 + 2], col[jj * 4 + 3]) 203 jj = jj + 1 204 } 205 let sc_pre: *i64 = sys_mmap(n_blocks * 8 * 8) as *i64 206 var blk: i64 = 0 207 while blk < n_blocks { 208 var is_: i64 = 0 209 while is_ < 8 { 210 var s: i64 = 0 211 var l: i64 = 0 212 let base: i64 = blk * 256 + is_ * 32 213 while l < 32 { s = s + col[base + l]; l = l + 1 } 214 sc_pre[blk * 8 + is_] = s 215 is_ = is_ + 1 216 } 217 blk = blk + 1 218 } 219 220 let qpk: *i64 = sys_mmap(64) as *i64 221 let qhi: *i64 = sys_mmap(64) as *i64 222 let acc: *i64 = sys_mmap(32) as *i64 223 let sm: *i64 = sys_mmap(16) as *i64 224 225 let scalar_dot: i64 = q5k_dot_scalar(buf, w_off, n_blocks, col, sm) 226 let simd_dot: i64 = q5k_dot_simd(buf, w_off, n_blocks, col_i16, qpk, qhi, acc, sc_pre, sm) 227 228 let IT: i64 = 20000 229 let t0: i64 = nx_clock_monotonic_ns() 230 var s1: i64 = 0 231 var k: i64 = 0 232 while k < IT { s1 = s1 + q5k_dot_scalar(buf, w_off, n_blocks, col, sm); k = k + 1 } 233 let t1: i64 = nx_clock_monotonic_ns() 234 let scalar_ns: i64 = t1 - t0 235 let t2: i64 = nx_clock_monotonic_ns() 236 var s2: i64 = 0 237 k = 0 238 while k < IT { s2 = s2 + q5k_dot_simd(buf, w_off, n_blocks, col_i16, qpk, qhi, acc, sc_pre, sm); k = k + 1 } 239 let t3: i64 = nx_clock_monotonic_ns() 240 let simd_ns: i64 = t3 - t2 241 242 let ofd: i64 = sys_openat_wr("/tmp/q5k_simd.txt" as *u8, 0x1a4) 243 if ofd >= 0 { 244 q5s_emit(ofd, "scalar_dot" as *u8, 10, scalar_dot) 245 q5s_emit(ofd, "simd_dot" as *u8, 8, simd_dot) 246 q5s_emit(ofd, "scalar_ns" as *u8, 9, scalar_ns) 247 q5s_emit(ofd, "simd_ns" as *u8, 7, simd_ns) 248 if simd_ns > 0 { q5s_emit(ofd, "speedup_x100" as *u8, 12, scalar_ns * 100 / simd_ns) } 249 sys_close(ofd) 250 } 251 if simd_dot != scalar_dot { return 80 } 252 if simd_ns >= scalar_ns { return 81 } 253 return 0 254}