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}