nx_q5k_dot_row_col.nx source
↩ module page · 149 lines · 6701 B
1// nx_q5k_dot_row_col.nx -- fused dequant + integer dot for Q5_K (the integer speed path for Q5_K weights).
2//
3// sd-server -> Nishi migration (task#21). Analogue of nx_q4k_dot_row_col for Q5_K, so attn_v/attn_output/
4// ffn_down (Q5_K in the real Q4_K_M model) can use the fast integer path too. Fuses the ggml-correct Q5_K
5// dequant (Q24 super-scales via _gguf_f16_to_q24, 4-group layout, qh 5th bit; value = d*sc*q5 - dmin*m in
6// Q24) with the dot against a Q10 activation, accumulating in Q34. Verified vs nx_q5_k_to_f32 + the ggml
7// golden dot on the real blk.0.attn_v block.
8// license_tier: ORIGINAL
9import "nx_syscalls.nx"
10import "nx_tier.nx"
11import "nx_le.nx"
12import "nx_strconv.nx"
13import "nx_gguf.nx"
14import "nx_gguf_load.nx"
15import "nx_q5_k_to_f32.nx"
16import "nx_f32.nx"
17import "nx_f32_cvt.nx"
18import "nx_f32_div.nx"
19
20// fused Q5_K dequant+dot: sum over values of (d*sc*q5 - dmin*m)[Q24] * col_q10[idx] -> Q34
21func nx_q5k_dot_row_col(buf: *u8, base_off: i64, n_blocks: i64, col_q10: *i64) -> i64 {
22 var dot: i64 = 0
23 var blk: i64 = 0
24 while blk < n_blocks {
25 let sb: i64 = base_off + blk * 176
26 let d_q24: i64 = _gguf_f16_to_q24(nx_le_read_u16(buf, sb))
27 let dmin_q24: i64 = _gguf_f16_to_q24(nx_le_read_u16(buf, sb + 2))
28 let scales_off: i64 = sb + 4
29 let qh_off: i64 = sb + 16
30 let qs_off: i64 = sb + 48
31 let col_base: i64 = blk * 256
32 var g: i64 = 0
33 while g < 4 {
34 let is0: i64 = g + g
35 let is1: i64 = is0 + 1
36 var sc0: i64 = 0
37 var m0: i64 = 0
38 var sc1: i64 = 0
39 var m1s: i64 = 0
40 if is0 < 4 {
41 sc0 = nx_le_read_u8(buf, scales_off + is0) & 0x3F
42 m0 = nx_le_read_u8(buf, scales_off + is0 + 4) & 0x3F
43 } else {
44 let k0: i64 = is0 - 4
45 sc0 = ((nx_le_read_u8(buf, scales_off + k0) >> 6) << 4) | (nx_le_read_u8(buf, scales_off + 8 + k0) & 0x0F)
46 m0 = ((nx_le_read_u8(buf, scales_off + 4 + k0) >> 6) << 4) | (nx_le_read_u8(buf, scales_off + 8 + k0) >> 4)
47 }
48 if is1 < 4 {
49 sc1 = nx_le_read_u8(buf, scales_off + is1) & 0x3F
50 m1s = nx_le_read_u8(buf, scales_off + is1 + 4) & 0x3F
51 } else {
52 let k1: i64 = is1 - 4
53 sc1 = ((nx_le_read_u8(buf, scales_off + k1) >> 6) << 4) | (nx_le_read_u8(buf, scales_off + 8 + k1) & 0x0F)
54 m1s = ((nx_le_read_u8(buf, scales_off + 4 + k1) >> 6) << 4) | (nx_le_read_u8(buf, scales_off + 8 + k1) >> 4)
55 }
56 let d1: i64 = d_q24 * sc0
57 let mm0: i64 = dmin_q24 * m0
58 let d2: i64 = d_q24 * sc1
59 let mm1: i64 = dmin_q24 * m1s
60 let grp: i64 = qs_off + g * 32
61 let u1: i64 = 1 << (g + g)
62 let u2: i64 = 1 << (g + g + 1)
63 var l: i64 = 0
64 while l < 32 {
65 let byte_v: i64 = nx_le_read_u8(buf, grp + l)
66 let qh_l: i64 = nx_le_read_u8(buf, qh_off + l)
67 var q_lo: i64 = byte_v & 0x0F
68 var q_hi: i64 = byte_v >> 4
69 if (qh_l & u1) != 0 { q_lo = q_lo + 16 }
70 if (qh_l & u2) != 0 { q_hi = q_hi + 16 }
71 let v_lo: i64 = d1 * q_lo - mm0
72 let v_hi: i64 = d2 * q_hi - mm1
73 dot = dot + v_lo * col_q10[col_base + is0 * 32 + l]
74 dot = dot + v_hi * col_q10[col_base + is1 * 32 + l]
75 l = l + 1
76 }
77 g = g + 1
78 }
79 blk = blk + 1
80 }
81 return dot
82}
83
84func kd_hexval(c: i64) -> i64 {
85 if c >= 48 { if c <= 57 { return c - 48 } }
86 if c >= 97 { if c <= 102 { return c - 87 } }
87 return 0
88}
89
90func main() -> i64 {
91 // embed the real blk.0.attn_v super-block 0 (Q5_K)
92 let hex: *u8 = "82046f15fdf2eceaf3b5dfa56565ffef507dce3a4f59475b485ad67adf51a09d1a2e66ea649f53795b3abb32898066cb6f8dff9f44e8405cfdca236f97e20bf51ecf4c0bae2380a2d10f0d4deeffa056063f91a559cae694f08a53bc6a77f6304d34b4759101794638b36e94abdbf312d9222dcbaddaab104d50b290a0c86f148387702e2f4264f00f2a7f86dbe20db7c8c80200f5a63245f45553b0faeb686c3cf85b1ac435f183c067e82e572ffe53" as *u8
93 let buf: *u8 = sys_mmap(176)
94 var i: i64 = 0
95 while i < 176 { buf[i] = ((kd_hexval((hex[i * 2]) as i64) << 4) | kd_hexval((hex[i * 2 + 1]) as i64)) as u8; i = i + 1 }
96
97 // activation col_q10[i] = ((i%5)+1) in Q10 (matches the golden dot's weighting)
98 let col: *i64 = sys_mmap(256 * 8) as *i64
99 i = 0
100 while i < 256 { col[i] = ((i - (i / 5) * 5) + 1) * 1024; i = i + 1 }
101
102 // integer fused dot -> Q34 -> f32
103 let dot_q34: i64 = nx_q5k_dot_row_col(buf, 0, 1, col)
104 let dot_q10: i64 = dot_q34 / 16777216 // Q34 -> Q10 (/2^24)
105 let int_f32: i64 = nx_q10_to_f32(dot_q10)
106
107 // f32 reference: nx_q5_k_to_f32 + f32 dot with the same (i%5)+1 weights
108 let yf: *i64 = sys_mmap(256 * 8) as *i64
109 nx_q5_k_to_f32(buf, 0, 256, yf)
110 var ref: i64 = 0
111 i = 0
112 while i < 256 { ref = nx_f32_add(ref, nx_f32_mul(yf[i], nx_i32_to_f32((i - (i / 5) * 5) + 1))); i = i + 1 }
113
114 let ofd: i64 = sys_openat_wr("/tmp/zimg_q5k_dot.txt" as *u8, 0x1a4)
115 if ofd >= 0 {
116 let line: *u8 = sys_mmap(64)
117 var lo: i64 = 0
118 let tag: *u8 = "int_f32_bits=" as *u8
119 var ti: i64 = 0
120 while tag[ti] != (0 as u8) { line[lo] = tag[ti]; lo = lo + 1; ti = ti + 1 }
121 let dec: *u8 = sys_mmap(32)
122 let nd: i64 = nx_strconv_format_i64(int_f32, dec)
123 var k: i64 = 0
124 while k < nd { line[lo] = dec[k]; lo = lo + 1; k = k + 1 }
125 line[lo] = 0x0A; lo = lo + 1
126 sys_write(ofd, line, lo)
127 let line2: *u8 = sys_mmap(64)
128 lo = 0
129 let tag2: *u8 = "ref_f32_bits=" as *u8
130 ti = 0
131 while tag2[ti] != (0 as u8) { line2[lo] = tag2[ti]; lo = lo + 1; ti = ti + 1 }
132 let nd2: i64 = nx_strconv_format_i64(ref, dec)
133 k = 0
134 while k < nd2 { line2[lo] = dec[k]; lo = lo + 1; k = k + 1 }
135 line2[lo] = 0x0A; lo = lo + 1
136 sys_write(ofd, line2, lo)
137 sys_close(ofd)
138 }
139
140 // gate 1: integer dot ~ f32 dot within 2%
141 let absref: i64 = ref & 0x7FFFFFFF
142 let tolf: i64 = nx_f32_div(nx_i32_to_f32(2), nx_i32_to_f32(100))
143 if (nx_f32_sub(int_f32, ref) & 0x7FFFFFFF) >= nx_f32_mul(tolf, absref) { return 80 }
144 // gate 2: f32 ref ~ ggml golden -2.357033 (bits 3222722976) within 1%
145 let gld: i64 = 3222722976
146 let tolg: i64 = nx_f32_div(nx_i32_to_f32(1), nx_i32_to_f32(100))
147 if (nx_f32_sub(ref, gld) & 0x7FFFFFFF) >= nx_f32_mul(tolg, gld & 0x7FFFFFFF) { return 81 }
148 return 0
149}