code wiki / (root) / nx_q5k_dot_row_col.nx

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}