use crate::quantized::k_quants::{
BlockIQ4nl, BlockQ4_0, BlockQ8_0, GgmlType, KVALUES_IQ4NL, QK4_0, QK4_NL, QK8_0,
};
use half::f16;
quant_format! {
name: BlockQ4_0Demo,
dtype: Q4_0,
block_elems: QK4_0,
byte_size: 18,
vec_dot: BlockQ8_0,
fields: {
d: f16,
qs: [u8; QK4_0 / 2],
},
decode: |xs, ys| {
let k = ys.len();
let qk = Self::BLCK_SIZE;
debug_assert!(k.is_multiple_of(qk), "dequantize_row_q4_0: {k} is not divisible by {qk}");
let nb = k / qk;
for i in 0..nb {
let d = xs[i].d.to_f32();
for j in 0..(qk / 2) {
let x0 = (xs[i].qs[j] & 0x0F) as i16 - 8;
let x1 = (xs[i].qs[j] >> 4) as i16 - 8;
ys[i * qk + j] = (x0 as f32) * d;
ys[i * qk + j + qk / 2] = (x1 as f32) * d;
}
}
},
}
quant_format! {
name: BlockQ8_0Demo,
dtype: Q8_0,
block_elems: QK8_0,
byte_size: 34,
vec_dot: BlockQ8_0,
fields: {
d: f16,
qs: [i8; QK8_0],
},
decode: |xs, ys| {
let k = ys.len();
debug_assert!(k.is_multiple_of(QK8_0), "dequantize_row_q8_0: {k} is not divisible by {QK8_0}");
let nb = k / QK8_0;
for i in 0..nb {
let d = xs[i].d.to_f32();
for j in 0..QK8_0 {
ys[i * QK8_0 + j] = xs[i].qs[j] as f32 * d;
}
}
},
}
quant_format! {
name: BlockIQ4nlDemo,
dtype: IQ4_NL,
block_elems: QK4_NL,
byte_size: 18,
vec_dot: BlockQ8_0,
fields: {
d: f16,
qs: [u8; QK4_NL / 2],
},
decode: |xs, ys| {
let k = ys.len();
debug_assert!(k.is_multiple_of(QK4_NL), "dequantize_row_iq4_nl: {k} is not divisible by {QK4_NL}");
let nb = k / QK4_NL;
for i in 0..nb {
let d = xs[i].d.to_f32();
for j in 0..QK4_NL / 2 {
let q = xs[i].qs[j];
ys[i * QK4_NL + j] = d * KVALUES_IQ4NL[(q & 0x0f) as usize] as f32;
ys[i * QK4_NL + j + QK4_NL / 2] = d * KVALUES_IQ4NL[(q >> 4) as usize] as f32;
}
}
},
}
fn sample_bytes(n: usize) -> Vec<u8> {
let mut state: u32 = 0x9E37_79B9;
(0..n)
.map(|_| {
state ^= state << 13;
state ^= state >> 17;
state ^= state << 5;
(state & 0xff) as u8
})
.collect()
}
fn as_blocks<T>(bytes: &[u8]) -> &[T] {
let sz = std::mem::size_of::<T>();
assert_eq!(bytes.len() % sz, 0);
assert_eq!(bytes.as_ptr() as usize % std::mem::align_of::<T>(), 0);
unsafe { std::slice::from_raw_parts(bytes.as_ptr() as *const T, bytes.len() / sz) }
}
fn aligned_sample(nbytes: usize) -> Vec<u8> {
let words = nbytes.div_ceil(16);
let mut backing: Vec<u128> = vec![0u128; words];
let raw =
unsafe { std::slice::from_raw_parts_mut(backing.as_mut_ptr() as *mut u8, words * 16) };
let src = sample_bytes(nbytes);
raw[..nbytes].copy_from_slice(&src);
let ptr = backing.as_ptr() as *const u8;
std::mem::forget(backing);
unsafe { Vec::from_raw_parts(ptr as *mut u8, nbytes, words * 16) }
}
fn assert_dequant_eq<Real, Demo>(n_blocks: usize, elems_per_block: usize)
where
Real: GgmlType,
Demo: GgmlType,
{
assert_eq!(
std::mem::size_of::<Demo>(),
std::mem::size_of::<Real>(),
"macro struct size differs from hand-written struct"
);
let nbytes = n_blocks * std::mem::size_of::<Real>();
let buf = aligned_sample(nbytes);
let real: &[Real] = as_blocks(&buf);
let demo: &[Demo] = as_blocks(&buf);
let n = n_blocks * elems_per_block;
let mut out_real = vec![0f32; n];
let mut out_demo = vec![0f32; n];
Real::to_float(real, &mut out_real);
Demo::to_float(demo, &mut out_demo);
for (i, (r, d)) in out_real.iter().zip(out_demo.iter()).enumerate() {
assert_eq!(
r.to_bits(),
d.to_bits(),
"dequant mismatch at element {i}: real={r} demo={d}"
);
}
}
#[test]
fn q4_0_macro_matches_handwritten() {
assert_dequant_eq::<BlockQ4_0, BlockQ4_0Demo>(8, QK4_0);
}
#[test]
fn q8_0_macro_matches_handwritten() {
assert_dequant_eq::<BlockQ8_0, BlockQ8_0Demo>(8, QK8_0);
}
#[test]
fn iq4nl_macro_matches_handwritten() {
assert_dequant_eq::<BlockIQ4nl, BlockIQ4nlDemo>(8, QK4_NL);
}
#[test]
fn iq4nl_macro_vec_dot_matches_reference() {
let n_blocks = 4usize;
let xbytes = aligned_sample(n_blocks * std::mem::size_of::<BlockIQ4nlDemo>());
let ybytes = aligned_sample(n_blocks * std::mem::size_of::<BlockQ8_0>());
let xs: &[BlockIQ4nlDemo] = as_blocks(&xbytes);
let ys: &[BlockQ8_0] = as_blocks(&ybytes);
let n = n_blocks * QK4_NL;
let got = BlockIQ4nlDemo::vec_dot(n, xs, ys);
let mut xf = vec![0f32; n];
BlockIQ4nlDemo::to_float(xs, &mut xf);
let mut want = 0f32;
for (b, y) in ys.iter().enumerate() {
let dy = y.d.to_f32();
for j in 0..QK8_0 {
want += xf[b * QK8_0 + j] * (y.qs[j] as f32 * dy);
}
}
assert_eq!(
got.to_bits(),
want.to_bits(),
"vec_dot mismatch: {got} vs {want}"
);
}