use half::f16;
#[cfg_attr(not(feature = "parallel"), allow(unused_imports))]
use crate::par::{IndexedParallelIterator, ParallelIterator, ParallelSlice, ParallelSliceMut};
#[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);
pub fn dequantize_q4_0_block(block: &BlockQ4_0) -> [f32; 32] {
let d = f16::from_bits(block.d).to_f32();
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
}
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);
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);
}
}
pub fn dequantize_q4_0_matrix(src: &[u8], m: usize, k: usize, out: &mut [f32]) {
debug_assert_eq!(
k % 32,
0,
"dequantize_q4_0_matrix: k must be a multiple of 32"
);
let row_bytes = (k / 32) * size_of::<BlockQ4_0>();
debug_assert_eq!(
src.len(),
m * row_bytes,
"dequantize_q4_0_matrix: src length mismatch"
);
debug_assert_eq!(
out.len(),
m * k,
"dequantize_q4_0_matrix: out length mismatch"
);
out.par_chunks_mut(k)
.zip(src.par_chunks(row_bytes))
.for_each(|(dst_row, src_row)| dequantize_q4_0_row(src_row, dst_row));
}
pub fn vec_dot_q4_0_f32_scalar(block: &BlockQ4_0, y: &[f32]) -> f32 {
debug_assert_eq!(y.len(), 32);
let d = f16::from_bits(block.d).to_f32();
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::from_bits(block.d).to_f32();
let m = f16::from_bits(block.m).to_f32();
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);
}
}
pub fn dequantize_q4_1_matrix(src: &[u8], m: usize, k: usize, out: &mut [f32]) {
debug_assert_eq!(
k % 32,
0,
"dequantize_q4_1_matrix: k must be a multiple of 32"
);
let row_bytes = (k / 32) * size_of::<BlockQ4_1>();
debug_assert_eq!(
src.len(),
m * row_bytes,
"dequantize_q4_1_matrix: src length mismatch"
);
debug_assert_eq!(
out.len(),
m * k,
"dequantize_q4_1_matrix: out length mismatch"
);
out.par_chunks_mut(k)
.zip(src.par_chunks(row_bytes))
.for_each(|(dst_row, src_row)| dequantize_q4_1_row(src_row, dst_row));
}
pub fn vec_dot_q4_1_f32(block: &BlockQ4_1, y: &[f32]) -> f32 {
debug_assert_eq!(y.len(), 32);
let d = f16::from_bits(block.d).to_f32();
let m = f16::from_bits(block.m).to_f32();
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::from_bits(block.delta).to_f32();
let mut out = [0.0f32; 32];
for (o, &q) in out.iter_mut().zip(block.quants.iter()) {
*o = q as f32 * d;
}
out
}
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);
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);
}
}
pub fn dequantize_q8_0_matrix(src: &[u8], m: usize, k: usize, out: &mut [f32]) {
debug_assert_eq!(
k % 32,
0,
"dequantize_q8_0_matrix: k must be a multiple of 32"
);
let row_bytes = (k / 32) * size_of::<BlockQ8_0>();
debug_assert_eq!(
src.len(),
m * row_bytes,
"dequantize_q8_0_matrix: src length mismatch"
);
debug_assert_eq!(
out.len(),
m * k,
"dequantize_q8_0_matrix: out length mismatch"
);
out.par_chunks_mut(k)
.zip(src.par_chunks(row_bytes))
.for_each(|(dst_row, src_row)| dequantize_q8_0_row(src_row, dst_row));
}
pub fn dequantize_q4_k_m_matrix(src: &[u8], m: usize, k: usize, out: &mut [f32]) {
debug_assert_eq!(
k % 256,
0,
"dequantize_q4_k_m_matrix: k must be a multiple of 256"
);
let row_bytes = (k / 256) * size_of::<BlockQ4KM>();
debug_assert_eq!(
src.len(),
m * row_bytes,
"dequantize_q4_k_m_matrix: src length mismatch"
);
debug_assert_eq!(
out.len(),
m * k,
"dequantize_q4_k_m_matrix: out length mismatch"
);
out.par_chunks_mut(k)
.zip(src.par_chunks(row_bytes))
.for_each(|(dst_row, src_row)| dequantize_q4_k_m_row(src_row, dst_row));
}
pub fn dequantize_q6_k_matrix(src: &[u8], m: usize, k: usize, out: &mut [f32]) {
debug_assert_eq!(
k % 256,
0,
"dequantize_q6_k_matrix: k must be a multiple of 256"
);
let row_bytes = (k / 256) * size_of::<BlockQ6K>();
debug_assert_eq!(
src.len(),
m * row_bytes,
"dequantize_q6_k_matrix: src length mismatch"
);
debug_assert_eq!(
out.len(),
m * k,
"dequantize_q6_k_matrix: out length mismatch"
);
out.par_chunks_mut(k)
.zip(src.par_chunks(row_bytes))
.for_each(|(dst_row, src_row)| dequantize_q6_k_row(src_row, dst_row));
}
pub fn vec_dot_q8_0_f32_scalar(block: &BlockQ8_0, y: &[f32]) -> f32 {
debug_assert_eq!(y.len(), 32);
let d = f16::from_bits(block.delta).to_f32();
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::from_bits(block.d).to_f32();
let dmin = f16::from_bits(block.dmin).to_f32();
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::from_bits(block.d).to_f32();
let dmin = f16::from_bits(block.dmin).to_f32();
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::from_bits(block.d).to_f32();
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::from_bits(block.d).to_f32();
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 {
vec_dot_q6_k_f32_scalar(block, y)
}
pub fn dequantize_q5_k_block(block: &BlockQ5K) -> [f32; 256] {
let d = f16::from_bits(block.d).to_f32();
let dmin = f16::from_bits(block.dmin).to_f32();
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::from_bits(block.d).to_f32();
let dmin = f16::from_bits(block.dmin).to_f32();
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 {
vec_dot_q5_k_f32_scalar(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::*;
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_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);
}
}