#![warn(missing_docs, clippy::missing_docs_in_private_items)]
#[cfg_attr(not(feature = "parallel"), allow(unused_imports))]
use crate::par::{IndexedParallelIterator, ParallelIterator, ParallelSlice, ParallelSliceMut};
#[inline(always)]
pub fn f16_to_f32(bits: u16) -> f32 {
#[cfg(target_arch = "aarch64")]
unsafe {
let out: f32;
std::arch::asm!(
"fcvt {0:s}, {1:h}",
out(vreg) out,
in(vreg) bits,
options(pure, nomem, nostack)
);
out
}
#[cfg(not(target_arch = "aarch64"))]
{
let sign = (u32::from(bits) & 0x8000) << 16;
let exp = (u32::from(bits) >> 10) & 0x1F;
let mant = u32::from(bits) & 0x03FF;
let rest = match exp {
0 if mant == 0 => 0,
0 => {
let shift = mant.leading_zeros() - 21;
((113 - shift) << 23) | ((mant << (shift + 13)) & 0x007F_FFFF)
}
0x1F if mant == 0 => 0x7F80_0000,
0x1F => 0x7FC0_0000 | (mant << 13),
_ => ((exp + 112) << 23) | (mant << 13),
};
f32::from_bits(sign | rest)
}
}
#[inline(always)]
pub fn bf16_to_f32(bits: u16) -> f32 {
let quiet = if (bits & 0x7FFF) > 0x7F80 {
0x0040_0000
} else {
0
};
f32::from_bits((u32::from(bits) << 16) | quiet)
}
#[inline(always)]
pub(crate) fn f32_to_f16(v: f32) -> u16 {
let b = v.to_bits();
let sign = ((b >> 16) & 0x8000) as u16;
let exp = ((b >> 23) & 0xFF) as i32;
let mant = b & 0x007F_FFFF;
if exp == 0xFF {
return if mant == 0 {
sign | 0x7C00
} else {
sign | 0x7E00 | ((mant >> 13) as u16 & 0x03FF)
};
}
let e = exp - 127 + 15;
if e >= 0x1F {
return sign | 0x7C00; }
if e <= 0 {
if e < -10 {
return sign;
}
let m = mant | 0x0080_0000;
let shift = (14 - e) as u32; let round = ((m >> (shift - 1)) & 1)
& (((m & ((1 << (shift - 1)) - 1)) != 0) as u32 | ((m >> shift) & 1));
return sign | ((m >> shift) + round) as u16;
}
let lsb = (mant >> 13) & 1;
let rounded = mant + 0x0FFF + lsb;
let carry = rounded >> 23;
let e = e as u32 + carry;
if e >= 0x1F {
return sign | 0x7C00;
}
sign | ((e << 10) as u16) | (((rounded >> 13) & 0x03FF) as u16)
}
#[repr(C, packed)]
#[derive(Debug, Clone, Copy)]
pub struct BlockQ4_0 {
pub d: u16,
pub qs: [u8; 16],
}
const _: () = assert!(size_of::<BlockQ4_0>() == 18);
#[repr(C, packed)]
#[derive(Debug, Clone, Copy)]
pub struct BlockQ4_1 {
pub d: u16,
pub m: u16,
pub qs: [u8; 16],
}
const _: () = assert!(size_of::<BlockQ4_1>() == 20);
#[repr(C, packed)]
#[derive(Debug, Clone, Copy)]
pub struct BlockQ8_0 {
pub delta: u16,
pub quants: [i8; 32],
}
const _: () = assert!(size_of::<BlockQ8_0>() == 34);
#[repr(C, packed)]
#[derive(Debug, Clone, Copy)]
pub struct BlockQ4KM {
pub d: u16,
pub dmin: u16,
pub scales: [u8; 12],
pub qs: [u8; 128],
}
const _: () = assert!(size_of::<BlockQ4KM>() == 144);
#[repr(C, packed)]
#[derive(Debug, Clone, Copy)]
pub struct BlockQ6K {
pub ql: [u8; 128],
pub qh: [u8; 64],
pub scales: [i8; 16],
pub d: u16,
}
const _: () = assert!(size_of::<BlockQ6K>() == 210);
#[repr(C, packed)]
#[derive(Debug, Clone, Copy)]
pub struct BlockQ5K {
pub d: u16,
pub dmin: u16,
pub scales: [u8; 12],
pub qh: [u8; 32],
pub qs: [u8; 128],
}
const _: () = assert!(size_of::<BlockQ5K>() == 176);
macro_rules! dequantize_matrix {
($name:ident, $elems:expr, $block:ty, $row:path, $($summary:expr),+ $(,)?) => {
$(#[doc = $summary])+
pub fn $name(src: &[u8], m: usize, k: usize, out: &mut [f32]) {
debug_assert_eq!(
k % $elems,
0,
concat!(
stringify!($name),
": k must be a multiple of ",
stringify!($elems)
)
);
let row_bytes = (k / $elems) * size_of::<$block>();
debug_assert_eq!(
src.len(),
m * row_bytes,
concat!(stringify!($name), ": src length mismatch")
);
debug_assert_eq!(
out.len(),
m * k,
concat!(stringify!($name), ": out length mismatch")
);
let num_threads = crate::par::current_num_threads().max(1);
let rows_per_chunk = (m / num_threads).max(1);
let dst_chunk_len = rows_per_chunk * k;
let src_chunk_len = rows_per_chunk * row_bytes;
out.par_chunks_mut(dst_chunk_len)
.zip(src.par_chunks(src_chunk_len))
.for_each(|(dst_chunk, src_chunk)| {
for (dst_row, src_row) in dst_chunk.chunks_mut(k).zip(src_chunk.chunks(row_bytes)) {
$row(src_row, dst_row);
}
});
}
};
}
pub fn dequantize_q4_0_block(block: &BlockQ4_0) -> [f32; 32] {
let d = f16_to_f32(block.d);
let mut out = [0.0f32; 32];
for i in 0..16 {
let byte = block.qs[i];
let lo = (byte & 0xF) as i32 - 8;
let hi = (byte >> 4) as i32 - 8;
out[i] = lo as f32 * d;
out[i + 16] = hi as f32 * d;
}
out
}
#[inline]
pub fn dequantize_q4_0_row(src: &[u8], dst: &mut [f32]) {
let block_size = size_of::<BlockQ4_0>();
let n_blocks = src.len() / block_size;
debug_assert_eq!(src.len() % block_size, 0);
debug_assert_eq!(dst.len(), n_blocks * 32);
#[cfg(target_arch = "aarch64")]
unsafe {
use std::arch::aarch64::*;
let mask_lo = vdupq_n_u8(0x0F);
let offset_8 = vdupq_n_s8(0x8);
for i in 0..n_blocks {
let block_ptr = src.as_ptr().add(i * block_size) as *const BlockQ4_0;
let block = &*block_ptr;
let d_val = f16_to_f32(block.d);
let d_vec = vdupq_n_f32(d_val);
let v = vld1q_u8(block.qs.as_ptr());
let v_lo = vsubq_s8(vreinterpretq_s8_u8(vandq_u8(v, mask_lo)), offset_8);
let v_hi = vsubq_s8(vreinterpretq_s8_u8(vshrq_n_u8::<4>(v)), offset_8);
let v_lo_s16_lo = vmovl_s8(vget_low_s8(v_lo));
let v_lo_s16_hi = vmovl_high_s8(v_lo);
let v_lo_0 = vmulq_f32(vcvtq_f32_s32(vmovl_s16(vget_low_s16(v_lo_s16_lo))), d_vec);
let v_lo_1 = vmulq_f32(vcvtq_f32_s32(vmovl_high_s16(v_lo_s16_lo)), d_vec);
let v_lo_2 = vmulq_f32(vcvtq_f32_s32(vmovl_s16(vget_low_s16(v_lo_s16_hi))), d_vec);
let v_lo_3 = vmulq_f32(vcvtq_f32_s32(vmovl_high_s16(v_lo_s16_hi)), d_vec);
let v_hi_s16_lo = vmovl_s8(vget_low_s8(v_hi));
let v_hi_s16_hi = vmovl_high_s8(v_hi);
let v_hi_0 = vmulq_f32(vcvtq_f32_s32(vmovl_s16(vget_low_s16(v_hi_s16_lo))), d_vec);
let v_hi_1 = vmulq_f32(vcvtq_f32_s32(vmovl_high_s16(v_hi_s16_lo)), d_vec);
let v_hi_2 = vmulq_f32(vcvtq_f32_s32(vmovl_s16(vget_low_s16(v_hi_s16_hi))), d_vec);
let v_hi_3 = vmulq_f32(vcvtq_f32_s32(vmovl_high_s16(v_hi_s16_hi)), d_vec);
let out_ptr = dst.as_mut_ptr().add(i * 32);
vst1q_f32(out_ptr, v_lo_0);
vst1q_f32(out_ptr.add(4), v_lo_1);
vst1q_f32(out_ptr.add(8), v_lo_2);
vst1q_f32(out_ptr.add(12), v_lo_3);
vst1q_f32(out_ptr.add(16), v_hi_0);
vst1q_f32(out_ptr.add(20), v_hi_1);
vst1q_f32(out_ptr.add(24), v_hi_2);
vst1q_f32(out_ptr.add(28), v_hi_3);
}
}
#[cfg(not(target_arch = "aarch64"))]
for i in 0..n_blocks {
let block_bytes = &src[i * block_size..(i + 1) * block_size];
let block = unsafe { &*(block_bytes.as_ptr() as *const BlockQ4_0) };
let values = dequantize_q4_0_block(block);
dst[i * 32..(i + 1) * 32].copy_from_slice(&values);
}
}
dequantize_matrix!(
dequantize_q4_0_matrix,
32,
BlockQ4_0,
dequantize_q4_0_row,
"Dequantize a Q4_0 matrix of shape `[m, k]` (row-major) to `out`."
);
pub fn vec_dot_q4_0_f32_scalar(block: &BlockQ4_0, y: &[f32]) -> f32 {
debug_assert_eq!(y.len(), 32);
let d = f16_to_f32(block.d);
let mut sum = 0.0f32;
for i in 0..16 {
let byte = block.qs[i];
let lo = (byte & 0xF) as i32 - 8;
let hi = (byte >> 4) as i32 - 8;
sum += lo as f32 * y[i];
sum += hi as f32 * y[i + 16];
}
sum * d
}
pub fn dequantize_q4_1_block(block: &BlockQ4_1) -> [f32; 32] {
let d = f16_to_f32(block.d);
let m = f16_to_f32(block.m);
let mut out = [0.0f32; 32];
for i in 0..16 {
let byte = block.qs[i];
let lo = (byte & 0xF) as i32;
let hi = (byte >> 4) as i32;
out[i] = lo as f32 * d + m;
out[i + 16] = hi as f32 * d + m;
}
out
}
pub fn dequantize_q4_1_row(src: &[u8], dst: &mut [f32]) {
let block_size = size_of::<BlockQ4_1>();
let n_blocks = src.len() / block_size;
debug_assert_eq!(src.len() % block_size, 0);
debug_assert_eq!(dst.len(), n_blocks * 32);
for i in 0..n_blocks {
let block_bytes = &src[i * block_size..(i + 1) * block_size];
let block = unsafe { &*(block_bytes.as_ptr() as *const BlockQ4_1) };
let values = dequantize_q4_1_block(block);
dst[i * 32..(i + 1) * 32].copy_from_slice(&values);
}
}
dequantize_matrix!(
dequantize_q4_1_matrix,
32,
BlockQ4_1,
dequantize_q4_1_row,
"Dequantize a Q4_1 matrix of shape `[m, k]` (row-major) to `out`."
);
pub fn vec_dot_q4_1_f32(block: &BlockQ4_1, y: &[f32]) -> f32 {
debug_assert_eq!(y.len(), 32);
let d = f16_to_f32(block.d);
let m = f16_to_f32(block.m);
let mut qsum = 0.0f32;
let mut ysum = 0.0f32;
for i in 0..16 {
let byte = block.qs[i];
let lo = (byte & 0xF) as i32;
let hi = (byte >> 4) as i32;
qsum += lo as f32 * y[i];
qsum += hi as f32 * y[i + 16];
ysum += y[i] + y[i + 16];
}
qsum * d + m * ysum
}
pub fn dequantize_q8_0_block(block: &BlockQ8_0) -> [f32; 32] {
let d = f16_to_f32(block.delta);
let mut out = [0.0f32; 32];
for (o, &q) in out.iter_mut().zip(block.quants.iter()) {
*o = q as f32 * d;
}
out
}
#[inline]
pub fn dequantize_q8_0_row(src: &[u8], dst: &mut [f32]) {
let block_size = size_of::<BlockQ8_0>();
let n_blocks = src.len() / block_size;
debug_assert_eq!(src.len() % block_size, 0);
debug_assert_eq!(dst.len(), n_blocks * 32);
#[cfg(target_arch = "aarch64")]
unsafe {
use std::arch::aarch64::*;
for i in 0..n_blocks {
let block = &*(src.as_ptr().add(i * block_size) as *const BlockQ8_0);
let d_val = f16_to_f32(block.delta);
let d_vec = vdupq_n_f32(d_val);
let q0 = vld1q_s8(block.quants.as_ptr());
let q1 = vld1q_s8(block.quants.as_ptr().add(16));
let q0_s16_lo = vmovl_s8(vget_low_s8(q0));
let q0_s16_hi = vmovl_high_s8(q0);
let q1_s16_lo = vmovl_s8(vget_low_s8(q1));
let q1_s16_hi = vmovl_high_s8(q1);
let f0 = vmulq_f32(vcvtq_f32_s32(vmovl_s16(vget_low_s16(q0_s16_lo))), d_vec);
let f1 = vmulq_f32(vcvtq_f32_s32(vmovl_high_s16(q0_s16_lo)), d_vec);
let f2 = vmulq_f32(vcvtq_f32_s32(vmovl_s16(vget_low_s16(q0_s16_hi))), d_vec);
let f3 = vmulq_f32(vcvtq_f32_s32(vmovl_high_s16(q0_s16_hi)), d_vec);
let f4 = vmulq_f32(vcvtq_f32_s32(vmovl_s16(vget_low_s16(q1_s16_lo))), d_vec);
let f5 = vmulq_f32(vcvtq_f32_s32(vmovl_high_s16(q1_s16_lo)), d_vec);
let f6 = vmulq_f32(vcvtq_f32_s32(vmovl_s16(vget_low_s16(q1_s16_hi))), d_vec);
let f7 = vmulq_f32(vcvtq_f32_s32(vmovl_high_s16(q1_s16_hi)), d_vec);
let out_ptr = dst.as_mut_ptr().add(i * 32);
vst1q_f32(out_ptr, f0);
vst1q_f32(out_ptr.add(4), f1);
vst1q_f32(out_ptr.add(8), f2);
vst1q_f32(out_ptr.add(12), f3);
vst1q_f32(out_ptr.add(16), f4);
vst1q_f32(out_ptr.add(20), f5);
vst1q_f32(out_ptr.add(24), f6);
vst1q_f32(out_ptr.add(28), f7);
}
}
#[cfg(not(target_arch = "aarch64"))]
for i in 0..n_blocks {
let block_bytes = &src[i * block_size..(i + 1) * block_size];
let block = unsafe { &*(block_bytes.as_ptr() as *const BlockQ8_0) };
let values = dequantize_q8_0_block(block);
dst[i * 32..(i + 1) * 32].copy_from_slice(&values);
}
}
dequantize_matrix!(
dequantize_q8_0_matrix,
32,
BlockQ8_0,
dequantize_q8_0_row,
"Dequantize a Q8_0 matrix of shape `[m, k]` (row-major) to `out`."
);
dequantize_matrix!(
dequantize_q4_k_m_matrix,
256,
BlockQ4KM,
dequantize_q4_k_m_row,
"Dequantize a Q4_K matrix of shape `[m, k]` (row-major) to `out`.",
"",
"Superblocks are 256 wide, so `k` must be a multiple of 256 (not 32).",
);
dequantize_matrix!(
dequantize_q5_k_matrix,
256,
BlockQ5K,
dequantize_q5_k_row,
"Dequantize a Q5_K matrix of shape `[m, k]` (row-major) to `out`.",
"",
"Superblocks are 256 wide, so `k` must be a multiple of 256 (not 32).",
"",
"Exists for the BLAS prefill route, which dequantizes the weight and SGEMMs",
"rather than running an int8 kernel: Q5_K is the one shipped K-quant with no",
"int8 GEMM, so before this it fell through to the per-token GEMV and prefill",
"collapsed (measured 4-5 tok/s against Q4_K_M's 228 on the same model and host).",
);
dequantize_matrix!(
dequantize_q6_k_matrix,
256,
BlockQ6K,
dequantize_q6_k_row,
"Dequantize a Q6_K matrix of shape `[m, k]` (row-major) to `out`.",
"",
"Superblocks are 256 wide, so `k` must be a multiple of 256 (not 32).",
);
pub fn vec_dot_q8_0_f32_scalar(block: &BlockQ8_0, y: &[f32]) -> f32 {
debug_assert_eq!(y.len(), 32);
let d = f16_to_f32(block.delta);
let sum: f32 = block
.quants
.iter()
.zip(y.iter())
.map(|(&q, &y)| q as f32 * y)
.sum();
sum * d
}
pub(crate) fn decode_q4km_scales(scales: &[u8; 12]) -> ([u8; 8], [u8; 8]) {
let mut sc = [0u8; 8];
let mut mn = [0u8; 8];
for j in 0..4 {
sc[j] = scales[j] & 63;
mn[j] = scales[j + 4] & 63;
}
for j in 4..8 {
sc[j] = (scales[j + 4] & 0xF) | ((scales[j - 4] >> 6) << 4);
mn[j] = (scales[j + 4] >> 4) | ((scales[j] >> 6) << 4);
}
(sc, mn)
}
pub fn dequantize_q4_k_m_block(block: &BlockQ4KM) -> [f32; 256] {
let d = f16_to_f32(block.d);
let dmin = f16_to_f32(block.dmin);
let (sc, mn) = decode_q4km_scales(&block.scales);
let mut out = [0.0f32; 256];
let qs = &block.qs;
for j in 0..8 {
let sc_val = d * sc[j] as f32;
let mn_val = dmin * mn[j] as f32;
let _ = (sc_val, mn_val); }
let mut qi = 0; let mut yi = 0;
for j in 0..4 {
let d_sc1 = d * sc[j * 2] as f32;
let d_mn1 = dmin * mn[j * 2] as f32;
let d_sc2 = d * sc[j * 2 + 1] as f32;
let d_mn2 = dmin * mn[j * 2 + 1] as f32;
for l in 0..32 {
out[yi + l] = d_sc1 * (qs[qi + l] & 0xF) as f32 - d_mn1;
out[yi + l + 32] = d_sc2 * (qs[qi + l] >> 4) as f32 - d_mn2;
}
qi += 32;
yi += 64;
}
out
}
pub fn dequantize_q4_k_m_row(src: &[u8], dst: &mut [f32]) {
let block_size = size_of::<BlockQ4KM>();
let n_blocks = src.len() / block_size;
debug_assert_eq!(src.len() % block_size, 0);
debug_assert_eq!(dst.len(), n_blocks * 256);
for i in 0..n_blocks {
let block_bytes = &src[i * block_size..(i + 1) * block_size];
let block = unsafe { &*(block_bytes.as_ptr() as *const BlockQ4KM) };
let values = dequantize_q4_k_m_block(block);
dst[i * 256..(i + 1) * 256].copy_from_slice(&values);
}
}
pub fn vec_dot_q4_k_m_f32_scalar(block: &BlockQ4KM, y: &[f32]) -> f32 {
debug_assert_eq!(y.len(), 256);
let d = f16_to_f32(block.d);
let dmin = f16_to_f32(block.dmin);
let (sc, mn) = decode_q4km_scales(&block.scales);
let qs = &block.qs;
let mut sumf = 0.0f32;
let mut qi = 0usize;
let mut yi = 0usize;
for j in 0..4 {
let sc1 = sc[j * 2] as f32;
let mn1 = mn[j * 2] as f32;
let sc2 = sc[j * 2 + 1] as f32;
let mn2 = mn[j * 2 + 1] as f32;
let mut sum1 = 0.0f32;
let mut sum2 = 0.0f32;
let mut sum_mn1 = 0.0f32;
let mut sum_mn2 = 0.0f32;
for l in 0..32 {
sum1 += (qs[qi + l] & 0xF) as f32 * y[yi + l];
sum2 += (qs[qi + l] >> 4) as f32 * y[yi + l + 32];
sum_mn1 += y[yi + l];
sum_mn2 += y[yi + l + 32];
}
sumf += d * (sc1 * sum1 + sc2 * sum2) - dmin * (mn1 * sum_mn1 + mn2 * sum_mn2);
qi += 32;
yi += 64;
}
sumf
}
pub fn dequantize_q6_k_block(block: &BlockQ6K) -> [f32; 256] {
let d = f16_to_f32(block.d);
let ql = &block.ql;
let qh = &block.qh;
let sc = &block.scales;
let mut out = [0.0f32; 256];
let mut ql_off = 0usize;
let mut qh_off = 0usize;
let mut sc_off = 0usize;
let mut y_off = 0usize;
for _n in 0..2 {
for l in 0..32 {
let is = l / 16;
let q1 = ((ql[ql_off + l] & 0xF) | ((qh[qh_off + l] & 3) << 4)) as i8 - 32;
let q2 = ((ql[ql_off + l + 32] & 0xF) | (((qh[qh_off + l] >> 2) & 3) << 4)) as i8 - 32;
let q3 = ((ql[ql_off + l] >> 4) | (((qh[qh_off + l] >> 4) & 3) << 4)) as i8 - 32;
let q4 = ((ql[ql_off + l + 32] >> 4) | (((qh[qh_off + l] >> 6) & 3) << 4)) as i8 - 32;
out[y_off + l] = d * sc[sc_off + is] as f32 * q1 as f32;
out[y_off + l + 32] = d * sc[sc_off + is + 2] as f32 * q2 as f32;
out[y_off + l + 64] = d * sc[sc_off + is + 4] as f32 * q3 as f32;
out[y_off + l + 96] = d * sc[sc_off + is + 6] as f32 * q4 as f32;
}
y_off += 128;
ql_off += 64;
qh_off += 32;
sc_off += 8;
}
out
}
pub fn dequantize_q6_k_row(src: &[u8], dst: &mut [f32]) {
let block_size = size_of::<BlockQ6K>();
let n_blocks = src.len() / block_size;
debug_assert_eq!(src.len() % block_size, 0);
debug_assert_eq!(dst.len(), n_blocks * 256);
for i in 0..n_blocks {
let block_bytes = &src[i * block_size..(i + 1) * block_size];
let block = unsafe { &*(block_bytes.as_ptr() as *const BlockQ6K) };
let values = dequantize_q6_k_block(block);
dst[i * 256..(i + 1) * 256].copy_from_slice(&values);
}
}
pub fn vec_dot_q6_k_f32_scalar(block: &BlockQ6K, y: &[f32]) -> f32 {
debug_assert_eq!(y.len(), 256);
let d = f16_to_f32(block.d);
let ql = &block.ql;
let qh = &block.qh;
let sc = &block.scales;
let mut sumf = 0.0f32;
let mut ql_off = 0usize;
let mut qh_off = 0usize;
let mut sc_off = 0usize;
let mut y_off = 0usize;
for _n in 0..2 {
for l in 0..32 {
let is = l / 16;
let q1 = ((ql[ql_off + l] & 0xF) | ((qh[qh_off + l] & 3) << 4)) as i8 - 32;
let q2 = ((ql[ql_off + l + 32] & 0xF) | (((qh[qh_off + l] >> 2) & 3) << 4)) as i8 - 32;
let q3 = ((ql[ql_off + l] >> 4) | (((qh[qh_off + l] >> 4) & 3) << 4)) as i8 - 32;
let q4 = ((ql[ql_off + l + 32] >> 4) | (((qh[qh_off + l] >> 6) & 3) << 4)) as i8 - 32;
sumf += sc[sc_off + is] as f32 * q1 as f32 * y[y_off + l];
sumf += sc[sc_off + is + 2] as f32 * q2 as f32 * y[y_off + l + 32];
sumf += sc[sc_off + is + 4] as f32 * q3 as f32 * y[y_off + l + 64];
sumf += sc[sc_off + is + 6] as f32 * q4 as f32 * y[y_off + l + 96];
}
y_off += 128;
ql_off += 64;
qh_off += 32;
sc_off += 8;
}
sumf * d
}
pub fn vec_dot_q6_k_f32(block: &BlockQ6K, y: &[f32]) -> f32 {
crate::backend::simd::vec_dot_q6_k_f32(block, y)
}
pub fn dequantize_q5_k_block(block: &BlockQ5K) -> [f32; 256] {
let d = f16_to_f32(block.d);
let dmin = f16_to_f32(block.dmin);
let (sc, mn) = decode_q4km_scales(&block.scales);
let ql = &block.qs; let qh = &block.qh;
let mut out = [0.0f32; 256];
let mut qi = 0usize; let mut yi = 0usize; let mut u1: u8 = 1;
let mut u2: u8 = 2;
for j in 0..4 {
let d1 = d * sc[j * 2] as f32;
let m1 = dmin * mn[j * 2] as f32;
let d2 = d * sc[j * 2 + 1] as f32;
let m2 = dmin * mn[j * 2 + 1] as f32;
for l in 0..32 {
let hi = if qh[l] & u1 != 0 { 16.0 } else { 0.0 };
out[yi + l] = d1 * ((ql[qi + l] & 0xF) as f32 + hi) - m1;
}
for l in 0..32 {
let hi = if qh[l] & u2 != 0 { 16.0 } else { 0.0 };
out[yi + l + 32] = d2 * ((ql[qi + l] >> 4) as f32 + hi) - m2;
}
qi += 32;
yi += 64;
u1 <<= 2;
u2 <<= 2;
}
out
}
pub fn dequantize_q5_k_row(src: &[u8], dst: &mut [f32]) {
let block_size = size_of::<BlockQ5K>();
let n_blocks = src.len() / block_size;
debug_assert_eq!(src.len() % block_size, 0);
debug_assert_eq!(dst.len(), n_blocks * 256);
for i in 0..n_blocks {
let block_bytes = &src[i * block_size..(i + 1) * block_size];
let block = unsafe { &*(block_bytes.as_ptr() as *const BlockQ5K) };
let values = dequantize_q5_k_block(block);
dst[i * 256..(i + 1) * 256].copy_from_slice(&values);
}
}
pub fn vec_dot_q5_k_f32_scalar(block: &BlockQ5K, y: &[f32]) -> f32 {
debug_assert_eq!(y.len(), 256);
let d = f16_to_f32(block.d);
let dmin = f16_to_f32(block.dmin);
let (sc, mn) = decode_q4km_scales(&block.scales);
let ql = &block.qs;
let qh = &block.qh;
let mut sumf = 0.0f32;
let mut qi = 0usize;
let mut yi = 0usize;
let mut u1: u8 = 1;
let mut u2: u8 = 2;
for j in 0..4 {
let sc1 = sc[j * 2] as f32;
let mn1 = mn[j * 2] as f32;
let sc2 = sc[j * 2 + 1] as f32;
let mn2 = mn[j * 2 + 1] as f32;
let mut sum1 = 0.0f32;
let mut sum2 = 0.0f32;
let mut sum_mn1 = 0.0f32;
let mut sum_mn2 = 0.0f32;
for l in 0..32 {
let hi1 = if qh[l] & u1 != 0 { 16.0 } else { 0.0 };
let hi2 = if qh[l] & u2 != 0 { 16.0 } else { 0.0 };
let q1 = (ql[qi + l] & 0xF) as f32 + hi1;
let q2 = (ql[qi + l] >> 4) as f32 + hi2;
sum1 += q1 * y[yi + l];
sum2 += q2 * y[yi + l + 32];
sum_mn1 += y[yi + l];
sum_mn2 += y[yi + l + 32];
}
sumf += d * (sc1 * sum1 + sc2 * sum2) - dmin * (mn1 * sum_mn1 + mn2 * sum_mn2);
qi += 32;
yi += 64;
u1 <<= 2;
u2 <<= 2;
}
sumf
}
pub fn vec_dot_q5_k_f32(block: &BlockQ5K, y: &[f32]) -> f32 {
crate::backend::simd::vec_dot_q5_k_f32(block, y)
}
pub fn vec_dot_q4_0_f32(block: &BlockQ4_0, y: &[f32]) -> f32 {
crate::backend::simd::vec_dot_q4_0_f32(block, y)
}
pub fn vec_dot_q8_0_f32(block: &BlockQ8_0, y: &[f32]) -> f32 {
crate::backend::simd::vec_dot_q8_0_f32(block, y)
}
pub fn vec_dot_q4_k_m_f32(block: &BlockQ4KM, y: &[f32]) -> f32 {
crate::backend::simd::vec_dot_q4_k_m_f32(block, y)
}
#[cfg(test)]
mod tests {
use super::*;
use half::f16;
#[test]
fn f32_to_f16_matches_half_crate() {
for bits in 0..=u16::MAX {
let v = f16_to_f32(bits);
assert_eq!(
f32_to_f16(v),
f16::from_f32(v).to_bits(),
"round-trip of half {bits:#06x} ({v})"
);
}
for &nan_bits in &[
0x7FC0_0000u32, 0xFFC0_0000, 0x7F80_0001, 0x7F80_1000, 0xFF80_0001, ] {
let v = f32::from_bits(nan_bits);
assert_eq!(
f32_to_f16(v),
f16::from_f32(v).to_bits(),
"narrowing f32 NaN {nan_bits:#010x}"
);
}
let cases: &[f32] = &[
0.0,
-0.0,
1.0,
-1.0,
f32::INFINITY,
f32::NEG_INFINITY,
65504.0, 65519.0, 65520.0, -65520.0,
1.0e30, 6.103_516e-5, 6.097_555e-5, 5.960_464_5e-8, 2.980_232_2e-8, 2.980_232_5e-8, 1.0e-10, -1.0e-10,
1.000_976_6, 1.000_488_3, 1.001_464_8, 2048.0, 2049.0,
4098.0,
];
for &v in cases {
assert_eq!(
f32_to_f16(v),
f16::from_f32(v).to_bits(),
"boundary case {v:e}"
);
}
let mut st = 0x1234_5678u32;
for _ in 0..200_000 {
st = st.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
let v = f32::from_bits(st);
assert_eq!(
f32_to_f16(v),
f16::from_f32(v).to_bits(),
"random f32 {:#010x} ({v:e})",
st
);
}
}
#[test]
fn bf16_widen_matches_half_crate_exhaustively() {
for bits in 0..=u16::MAX {
let want = half::bf16::from_bits(bits).to_f32();
let got = bf16_to_f32(bits);
if want.is_nan() {
assert!(got.is_nan(), "bits {bits:#06x}: want NaN, got {got}");
assert_eq!(
got.to_bits(),
want.to_bits(),
"bits {bits:#06x}: NaN payload/sign differs"
);
} else {
assert_eq!(
got.to_bits(),
want.to_bits(),
"bits {bits:#06x}: want {want}, got {got}"
);
}
}
}
#[test]
fn f16_widen_matches_half_crate_exhaustively() {
for bits in 0..=u16::MAX {
let want = half::f16::from_bits(bits).to_f32();
let got = f16_to_f32(bits);
if want.is_nan() {
assert!(got.is_nan(), "bits {bits:#06x}: want NaN, got {got}");
assert_eq!(
got.to_bits() & 0x807F_FFFF,
want.to_bits() & 0x807F_FFFF,
"bits {bits:#06x}: NaN payload/sign differs"
);
} else {
assert_eq!(
got.to_bits(),
want.to_bits(),
"bits {bits:#06x}: want {want}, got {got}"
);
}
}
}
fn make_q8_0_block(scale: f32, quants: [i8; 32]) -> BlockQ8_0 {
BlockQ8_0 {
delta: f16::from_f32(scale).to_bits(),
quants,
}
}
#[test]
fn test_dequantize_q4_1_matches_ggml_formula() {
let mut qs = [0u8; 16];
for (i, qsi) in qs.iter_mut().enumerate() {
*qsi = (i as u8) | (((15 - i) as u8) << 4);
}
let d = 0.25f32;
let m = -1.5f32;
let block = BlockQ4_1 {
d: f16::from_f32(d).to_bits(),
m: f16::from_f32(m).to_bits(),
qs,
};
let out = dequantize_q4_1_block(&block);
let d = f16::from_f32(d).to_f32();
let m = f16::from_f32(m).to_f32();
for i in 0..16 {
let want_lo = i as f32 * d + m;
let want_hi = (15 - i) as f32 * d + m;
assert!(
(out[i] - want_lo).abs() < 1e-5,
"lo[{i}]: got {} want {want_lo}",
out[i]
);
assert!(
(out[i + 16] - want_hi).abs() < 1e-5,
"hi[{i}]: got {} want {want_hi}",
out[i + 16]
);
}
}
#[test]
fn test_vec_dot_q4_1_matches_dequantize() {
let mut st = 0x9e37_79b9u64;
let mut lcg = || {
st = st.wrapping_mul(6364136223846793005).wrapping_add(1);
((st >> 40) as f32 / (1u64 << 24) as f32) * 2.0 - 1.0
};
for trial in 0..8 {
let mut qs = [0u8; 16];
for qsi in qs.iter_mut() {
*qsi = ((lcg() + 1.0) * 127.0) as u8;
}
let block = BlockQ4_1 {
d: f16::from_f32(0.1 + trial as f32 * 0.05).to_bits(),
m: f16::from_f32(lcg()).to_bits(),
qs,
};
let y: Vec<f32> = (0..32).map(|_| lcg()).collect();
let want: f32 = dequantize_q4_1_block(&block)
.iter()
.zip(&y)
.map(|(a, b)| a * b)
.sum();
let got = vec_dot_q4_1_f32(&block, &y);
assert!(
(got - want).abs() <= 1e-4 * (1.0 + want.abs()),
"trial {trial}: got {got} want {want}"
);
}
}
#[test]
fn test_dequantize_q4_1_zero_scale_keeps_min() {
let block = BlockQ4_1 {
d: f16::from_f32(0.0).to_bits(),
m: f16::from_f32(2.5).to_bits(),
qs: [0xAB; 16],
};
let out = dequantize_q4_1_block(&block);
assert!(out.iter().all(|v| (v - 2.5).abs() < 1e-5), "{out:?}");
}
#[test]
fn test_dequantize_q4_1_row_matches_block() {
let mut bytes = Vec::new();
for b in 0..3u16 {
bytes.extend_from_slice(&f16::from_f32(0.5).to_bits().to_le_bytes());
bytes.extend_from_slice(&f16::from_f32(-0.25).to_bits().to_le_bytes());
bytes.extend_from_slice(&[b as u8 | 0x30; 16]);
}
let mut dst = vec![0.0f32; 96];
dequantize_q4_1_row(&bytes, &mut dst);
for b in 0..3usize {
let block = BlockQ4_1 {
d: f16::from_f32(0.5).to_bits(),
m: f16::from_f32(-0.25).to_bits(),
qs: [b as u8 | 0x30; 16],
};
let want = dequantize_q4_1_block(&block);
assert_eq!(&dst[b * 32..(b + 1) * 32], &want[..], "block {b}");
}
}
#[test]
fn test_dequantize_q4_0_simple() {
let block = BlockQ4_0 {
d: f16::from_f32(1.0).to_bits(),
qs: [0x88; 16], };
let out = dequantize_q4_0_block(&block);
for (i, &v) in out.iter().enumerate() {
assert!(v.abs() < 1e-3, "expected 0.0 at {i}, got {v}");
}
}
#[test]
fn test_dequantize_q4_0_varied() {
let mut qs = [0u8; 16];
for (i, qsi) in qs.iter_mut().enumerate() {
*qsi = (i as u8) | (15 << 4);
}
let block = BlockQ4_0 {
d: f16::from_f32(0.5).to_bits(),
qs,
};
let out = dequantize_q4_0_block(&block);
for (i, &v) in out.iter().enumerate().take(16) {
let expected = (i as f32 - 8.0) * 0.5;
assert!(
(v - expected).abs() < 1e-3,
"lo[{i}]: got {v}, expected {expected}"
);
}
for (i, &v) in out.iter().enumerate().skip(16) {
assert!((v - 3.5).abs() < 1e-3, "hi[{i}]: got {v}, expected 3.5");
}
}
#[test]
fn test_vec_dot_q4_0_matches_dequantize() {
let mut qs = [0u8; 16];
for (i, qsi) in qs.iter_mut().enumerate() {
*qsi = ((i % 13) as u8) | (((i % 7) as u8) << 4);
}
let block = BlockQ4_0 {
d: f16::from_f32(0.3).to_bits(),
qs,
};
let y: Vec<f32> = (0..32).map(|i| (i as f32 - 16.0) * 0.1).collect();
let dequantized = dequantize_q4_0_block(&block);
let expected: f32 = dequantized.iter().zip(y.iter()).map(|(a, b)| a * b).sum();
let got = vec_dot_q4_0_f32(&block, &y);
assert!(
(got - expected).abs() < 1e-3,
"vec_dot Q4_0 mismatch: got {got}, expected {expected}"
);
}
#[test]
fn test_dequantize_q8_0_simple() {
let block = make_q8_0_block(0.5, {
let mut q = [0i8; 32];
for (i, qi) in q.iter_mut().enumerate() {
*qi = i as i8;
}
q
});
let out = dequantize_q8_0_block(&block);
for (i, &v) in out.iter().enumerate() {
let expected = i as f32 * 0.5;
assert!(
(v - expected).abs() < 1e-3,
"mismatch at {i}: got {v}, expected {expected}"
);
}
}
#[test]
fn test_dequantize_q8_0_row() {
let block1 = make_q8_0_block(1.0, {
let mut q = [0i8; 32];
for (i, qi) in q.iter_mut().enumerate() {
*qi = (i as i8) - 16;
}
q
});
let block2 = make_q8_0_block(0.25, [1i8; 32]);
let mut src = vec![0u8; 68];
unsafe {
std::ptr::copy_nonoverlapping(&block1 as *const _ as *const u8, src.as_mut_ptr(), 34);
std::ptr::copy_nonoverlapping(
&block2 as *const _ as *const u8,
src.as_mut_ptr().add(34),
34,
);
}
let mut dst = vec![0.0f32; 64];
dequantize_q8_0_row(&src, &mut dst);
for (i, &v) in dst.iter().enumerate().take(32) {
let expected = (i as f32 - 16.0) * 1.0;
assert!(
(v - expected).abs() < 1e-3,
"block1[{i}]: got {v}, expected {expected}"
);
}
for i in 0..32 {
let expected = 1.0 * 0.25;
assert!(
(dst[32 + i] - expected).abs() < 1e-3,
"block2[{i}]: got {}, expected {expected}",
dst[32 + i]
);
}
}
#[test]
fn test_dequantize_q5_k_matrix_matches_row() {
let m = 128;
let k = 512; let blocks_per_row = k / 256;
let row_bytes = blocks_per_row * size_of::<BlockQ5K>();
let mut st = 0x5EED_1234u32;
let mut byte = || {
st = st.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
(st >> 24) as u8
};
let mut src = vec![0u8; m * row_bytes];
src.iter_mut().for_each(|b| *b = byte());
(0..m)
.flat_map(|row| (0..blocks_per_row).map(move |b| (row, b)))
.for_each(|(row, b)| {
let off = row * row_bytes + b * size_of::<BlockQ5K>();
let d = f16::from_f32(0.05 + (row % 7) as f32 * 0.01).to_bits();
let dmin = f16::from_f32(0.02 + (b % 3) as f32 * 0.01).to_bits();
src[off..off + 2].copy_from_slice(&d.to_le_bytes());
src[off + 2..off + 4].copy_from_slice(&dmin.to_le_bytes());
});
let mut via_matrix = vec![0.0f32; m * k];
dequantize_q5_k_matrix(&src, m, k, &mut via_matrix);
let mut via_rows = vec![0.0f32; m * k];
via_rows
.chunks_mut(k)
.zip(src.chunks(row_bytes))
.for_each(|(dst, s)| dequantize_q5_k_row(s, dst));
assert_eq!(
via_matrix, via_rows,
"matrix and row dequantization diverged"
);
}
#[test]
fn test_dequantize_q4_0_matrix_matches_row() {
let m = 128; let k = 64; let blocks_per_row = k / 32;
let row_bytes = blocks_per_row * size_of::<BlockQ4_0>();
let mut src = vec![0u8; m * row_bytes];
for row in 0..m {
for b in 0..blocks_per_row {
let block = BlockQ4_0 {
d: f16::from_f32(0.1 + (row as f32) * 0.01).to_bits(),
qs: {
let mut qs = [0u8; 16];
for (i, q) in qs.iter_mut().enumerate() {
*q = ((row + b * 7 + i * 3) as u8).wrapping_mul(17);
}
qs
},
};
let offset = row * row_bytes + b * size_of::<BlockQ4_0>();
unsafe {
std::ptr::copy_nonoverlapping(
&block as *const _ as *const u8,
src.as_mut_ptr().add(offset),
size_of::<BlockQ4_0>(),
);
}
}
}
let mut matrix_out = vec![0.0f32; m * k];
dequantize_q4_0_matrix(&src, m, k, &mut matrix_out);
let mut expected = vec![0.0f32; m * k];
for row in 0..m {
let src_row = &src[row * row_bytes..(row + 1) * row_bytes];
let dst_row = &mut expected[row * k..(row + 1) * k];
dequantize_q4_0_row(src_row, dst_row);
}
assert_eq!(matrix_out, expected);
}
#[test]
fn test_dequantize_q8_0_matrix_matches_row() {
let m = 96;
let k = 96; let blocks_per_row = k / 32;
let row_bytes = blocks_per_row * size_of::<BlockQ8_0>();
let mut src = vec![0u8; m * row_bytes];
for row in 0..m {
for b in 0..blocks_per_row {
let block = make_q8_0_block(0.05 * (1 + row) as f32 + 0.001 * b as f32, {
let mut q = [0i8; 32];
for (i, slot) in q.iter_mut().enumerate() {
*slot = ((row + b + i) as i8).wrapping_mul(5).wrapping_sub(64);
}
q
});
let offset = row * row_bytes + b * size_of::<BlockQ8_0>();
unsafe {
std::ptr::copy_nonoverlapping(
&block as *const _ as *const u8,
src.as_mut_ptr().add(offset),
size_of::<BlockQ8_0>(),
);
}
}
}
let mut matrix_out = vec![0.0f32; m * k];
dequantize_q8_0_matrix(&src, m, k, &mut matrix_out);
let mut expected = vec![0.0f32; m * k];
for row in 0..m {
let src_row = &src[row * row_bytes..(row + 1) * row_bytes];
let dst_row = &mut expected[row * k..(row + 1) * k];
dequantize_q8_0_row(src_row, dst_row);
}
assert_eq!(matrix_out, expected);
}
#[test]
fn test_vec_dot_q8_0() {
let block = make_q8_0_block(0.1, {
let mut q = [0i8; 32];
for (i, qi) in q.iter_mut().enumerate() {
*qi = (i as i8) * 2 - 31;
}
q
});
let y: Vec<f32> = (0..32).map(|i| i as f32 * 0.5).collect();
let dequantized = dequantize_q8_0_block(&block);
let expected: f32 = dequantized.iter().zip(y.iter()).map(|(a, b)| a * b).sum();
let got = vec_dot_q8_0_f32(&block, &y);
assert!(
(got - expected).abs() < 1e-3,
"vec_dot mismatch: got {got}, expected {expected}"
);
}
#[test]
fn test_dequantize_q4_k_m_basic() {
let mut block = BlockQ4KM {
d: f16::from_f32(1.0).to_bits(),
dmin: f16::from_f32(0.0).to_bits(), scales: [0u8; 12],
qs: [0u8; 128],
};
for i in 0..4 {
block.scales[i] = 1; }
for i in 4..8 {
block.scales[i] = 0; }
for i in 8..12 {
block.scales[i] = 0x01; }
for b in block.qs.iter_mut() {
*b = 0x33; }
let out = dequantize_q4_k_m_block(&block);
for (i, &v) in out.iter().enumerate() {
assert!(
(v - 3.0).abs() < 1e-3,
"mismatch at {i}: got {v}, expected 3.0"
);
}
}
#[test]
fn kquant_matrix_dequant_matches_row_dequant() {
let mut st = 0x2468_1357u64;
let next = |st: &mut u64| {
*st = st.wrapping_mul(6364136223846793005).wrapping_add(1);
(*st >> 33) as u8
};
let (m, k) = (3usize, 512usize);
let nb = k / 256;
let mut src = Vec::new();
for _ in 0..m * nb {
let blk = BlockQ4KM {
d: half::f16::from_f32(0.03).to_bits(),
dmin: half::f16::from_f32(0.01).to_bits(),
scales: std::array::from_fn(|_| next(&mut st)),
qs: std::array::from_fn(|_| next(&mut st)),
};
src.extend_from_slice(unsafe {
std::slice::from_raw_parts((&raw const blk) as *const u8, size_of::<BlockQ4KM>())
});
}
let mut got = vec![0.0f32; m * k];
dequantize_q4_k_m_matrix(&src, m, k, &mut got);
let row_bytes = nb * size_of::<BlockQ4KM>();
for i in 0..m {
let mut want = vec![0.0f32; k];
dequantize_q4_k_m_row(&src[i * row_bytes..(i + 1) * row_bytes], &mut want);
assert_eq!(&got[i * k..(i + 1) * k], &want[..], "Q4_K row {i} mismatch");
}
let mut src = Vec::new();
for _ in 0..m * nb {
let blk = BlockQ6K {
ql: std::array::from_fn(|_| next(&mut st)),
qh: std::array::from_fn(|_| next(&mut st)),
scales: std::array::from_fn(|_| next(&mut st) as i8),
d: half::f16::from_f32(0.02).to_bits(),
};
src.extend_from_slice(unsafe {
std::slice::from_raw_parts((&raw const blk) as *const u8, size_of::<BlockQ6K>())
});
}
let mut got = vec![0.0f32; m * k];
dequantize_q6_k_matrix(&src, m, k, &mut got);
let row_bytes = nb * size_of::<BlockQ6K>();
for i in 0..m {
let mut want = vec![0.0f32; k];
dequantize_q6_k_row(&src[i * row_bytes..(i + 1) * row_bytes], &mut want);
assert_eq!(&got[i * k..(i + 1) * k], &want[..], "Q6_K row {i} mismatch");
}
}
#[test]
fn test_vec_dot_q4km_matches_dequantize() {
let mut block = BlockQ4KM {
d: f16::from_f32(0.5).to_bits(),
dmin: f16::from_f32(0.1).to_bits(),
scales: [0u8; 12],
qs: [0u8; 128],
};
for i in 0..4 {
block.scales[i] = 2;
}
for i in 4..8 {
block.scales[i] = 1;
}
for i in 8..12 {
block.scales[i] = 0x21; }
for (i, b) in block.qs.iter_mut().enumerate() {
*b = ((i % 7) as u8) | (((i % 11) as u8) << 4);
}
let y: Vec<f32> = (0..256).map(|i| (i as f32 - 128.0) * 0.01).collect();
let dequantized = dequantize_q4_k_m_block(&block);
let expected: f32 = dequantized.iter().zip(y.iter()).map(|(a, b)| a * b).sum();
let got = vec_dot_q4_k_m_f32(&block, &y);
assert!(
(got - expected).abs() < 1e-2,
"vec_dot mismatch: got {got}, expected {expected}"
);
}
#[test]
fn test_dequantize_q6_k_basic() {
let mut block = BlockQ6K {
ql: [0u8; 128],
qh: [0u8; 64],
scales: [1i8; 16],
d: f16::from_f32(1.0).to_bits(),
};
for b in block.ql.iter_mut() {
*b = 0x00;
}
for b in block.qh.iter_mut() {
*b = 0xAA; }
let out = dequantize_q6_k_block(&block);
for (i, &v) in out.iter().enumerate() {
assert!(v.abs() < 1e-5, "expected ~0.0 at {i}, got {v}");
}
}
#[test]
fn test_vec_dot_q6_k_matches_dequantize() {
let mut block = BlockQ6K {
ql: [0u8; 128],
qh: [0u8; 64],
scales: [0i8; 16],
d: f16::from_f32(0.5).to_bits(),
};
for (i, s) in block.scales.iter_mut().enumerate() {
*s = (i as i8 % 5) + 1;
}
for (i, b) in block.ql.iter_mut().enumerate() {
*b = ((i % 13) as u8) | (((i % 9) as u8) << 4);
}
for (i, b) in block.qh.iter_mut().enumerate() {
*b = (i % 256) as u8;
}
let y: Vec<f32> = (0..256).map(|i| (i as f32 - 128.0) * 0.01).collect();
let dequantized = dequantize_q6_k_block(&block);
let expected: f32 = dequantized.iter().zip(y.iter()).map(|(a, b)| a * b).sum();
let got = vec_dot_q6_k_f32(&block, &y);
assert!(
(got - expected).abs() < 1e-2,
"vec_dot Q6_K mismatch: got {got}, expected {expected}"
);
}
#[test]
fn test_q5_k_block_is_176_bytes() {
assert_eq!(size_of::<BlockQ5K>(), 176);
}
#[test]
fn test_dequantize_q5_k_all_zero_quants() {
let mut block = BlockQ5K {
d: f16::from_f32(1.0).to_bits(),
dmin: f16::from_f32(0.0).to_bits(),
scales: [0u8; 12],
qh: [0u8; 32],
qs: [0u8; 128],
};
for s in block.scales.iter_mut().take(4) {
*s = 1;
}
let out = dequantize_q5_k_block(&block);
for (i, &v) in out.iter().enumerate() {
assert!(v.abs() < 1e-5, "expected ~0.0 at {i}, got {v}");
}
}
#[test]
fn test_dequantize_q5_k_high_bit_extends_range() {
let mut block = BlockQ5K {
d: f16::from_f32(1.0).to_bits(),
dmin: f16::from_f32(0.0).to_bits(),
scales: [0u8; 12],
qh: [0u8; 32],
qs: [0u8; 128],
};
block.scales[0] = 1; block.qs[0] = 0x0F; block.qh[0] = 0x01; let out = dequantize_q5_k_block(&block);
assert!(
(out[0] - 31.0).abs() < 1e-4,
"expected 31.0 (15 + 16) at index 0, got {}",
out[0]
);
assert!(
out[1].abs() < 1e-4,
"expected 0.0 at index 1, got {}",
out[1]
);
}
#[test]
fn test_vec_dot_q5_k_matches_dequantize() {
let mut block = BlockQ5K {
d: f16::from_f32(0.5).to_bits(),
dmin: f16::from_f32(0.125).to_bits(),
scales: [0u8; 12],
qh: [0u8; 32],
qs: [0u8; 128],
};
for (i, s) in block.scales.iter_mut().enumerate() {
*s = ((i * 7 + 3) % 64) as u8;
}
for (i, b) in block.qs.iter_mut().enumerate() {
*b = ((i % 7) as u8) | (((i % 11) as u8) << 4);
}
for (i, b) in block.qh.iter_mut().enumerate() {
*b = ((i * 37) % 256) as u8;
}
let y: Vec<f32> = (0..256).map(|i| (i as f32 - 128.0) * 0.01).collect();
let dequantized = dequantize_q5_k_block(&block);
let expected: f32 = dequantized.iter().zip(y.iter()).map(|(a, b)| a * b).sum();
let got = vec_dot_q5_k_f32(&block, &y);
assert!(
(got - expected).abs() < 1e-2,
"vec_dot Q5_K mismatch: got {got}, expected {expected}"
);
}
#[test]
fn test_dequantize_q5_k_row_multiple_blocks() {
let mut bytes = vec![0u8; 2 * size_of::<BlockQ5K>()];
bytes[0..2].copy_from_slice(&f16::from_f32(1.0).to_bits().to_le_bytes());
let b1 = size_of::<BlockQ5K>();
bytes[b1..b1 + 2].copy_from_slice(&f16::from_f32(2.0).to_bits().to_le_bytes());
bytes[4] = 1;
bytes[b1 + 4] = 1;
bytes[b1 + 4 + 12 + 32] = 0x03; let mut dst = vec![0.0f32; 512];
dequantize_q5_k_row(&bytes, &mut dst);
assert!(
(dst[256] - 6.0).abs() < 1e-3,
"expected 6.0 at block-1 value 0, got {}",
dst[256]
);
}
#[test]
fn test_decode_q4km_scales_roundtrip() {
let mut scales = [0u8; 12];
scales[0] = 5;
scales[1] = 10;
scales[2] = 15;
scales[3] = 20;
scales[4] = 1;
scales[5] = 2;
scales[6] = 3;
scales[7] = 4;
scales[8] = 0;
scales[9] = 0;
scales[10] = 0;
scales[11] = 0;
let (sc, mn) = decode_q4km_scales(&scales);
assert_eq!(sc[0], 5);
assert_eq!(sc[1], 10);
assert_eq!(sc[2], 15);
assert_eq!(sc[3], 20);
assert_eq!(mn[0], 1);
assert_eq!(mn[1], 2);
assert_eq!(mn[2], 3);
assert_eq!(mn[3], 4);
}
}