use half::f16;
use rayon::prelude::*;
use super::common::make_qkx3_quants;
pub const QK5_1: usize = 32;
pub const BLOCK_BYTES: usize = 2 + 2 + 4 + QK5_1 / 2;
fn quantize_row_q5_1_ref(x: &[f32], out: &mut Vec<u8>) {
let qk = QK5_1;
debug_assert!(x.len() % qk == 0);
let nb = x.len() / qk;
for i in 0..nb {
let block = &x[i * qk..(i + 1) * qk];
let mut min = f32::MAX;
let mut max = f32::MIN;
for &v in block {
if v < min {
min = v;
}
if v > max {
max = v;
}
}
let d = (max - min) / 31.0; let id = if d != 0.0 { 1.0 / d } else { 0.0 };
let d_f16 = f16::from_f32(d);
let m_f16 = f16::from_f32(min);
out.extend_from_slice(&d_f16.to_le_bytes());
out.extend_from_slice(&m_f16.to_le_bytes());
let mut qh: u32 = 0;
let qh_pos = out.len();
out.extend_from_slice(&[0u8; 4]);
let qs_pos = out.len();
out.extend_from_slice(&[0u8; QK5_1 / 2]);
for j in 0..qk / 2 {
let x0 = (block[j] - min) * id;
let x1 = (block[qk / 2 + j] - min) * id;
let xi0 = (x0 + 0.5) as u8;
let xi1 = (x1 + 0.5) as u8;
out[qs_pos + j] = (xi0 & 0x0F) | ((xi1 & 0x0F) << 4);
qh |= (((xi0 & 0x10) as u32) >> 4) << j;
qh |= (((xi1 & 0x10) as u32) >> 4) << (j + qk / 2);
}
let qh_bytes = qh.to_le_bytes();
out[qh_pos..qh_pos + 4].copy_from_slice(&qh_bytes);
}
}
fn quantize_row_q5_1_impl(x: &[f32], quant_weights: &[f32], out: &mut Vec<u8>) {
debug_assert_eq!(QK5_1, 32);
let n_per_row = x.len();
let mut sum_x2 = 0.0f32;
for j in 0..n_per_row {
sum_x2 += x[j] * x[j];
}
let sigma2 = sum_x2 / (n_per_row as f32);
let nb = n_per_row / QK5_1;
let mut weight = [0.0f32; QK5_1];
let mut l_buf = [0u8; QK5_1];
let mut laux = [0u8; QK5_1];
for ib in 0..nb {
let xb = &x[QK5_1 * ib..QK5_1 * (ib + 1)];
let qw = &quant_weights[QK5_1 * ib..QK5_1 * (ib + 1)];
for j in 0..QK5_1 {
weight[j] = qw[j] * (sigma2 + xb[j] * xb[j]).sqrt();
}
let mut the_min = 0.0f32;
let d = make_qkx3_quants(
QK5_1,
31,
xb,
Some(&weight),
&mut l_buf,
&mut the_min,
&mut laux,
-0.9,
0.05,
36,
false,
);
let d_f16 = f16::from_f32(d);
let m_f16 = f16::from_f32(-the_min);
out.extend_from_slice(&d_f16.to_le_bytes());
out.extend_from_slice(&m_f16.to_le_bytes());
let mut qh: u32 = 0;
let qh_pos = out.len();
out.extend_from_slice(&[0u8; 4]);
let qs_pos = out.len();
out.extend_from_slice(&[0u8; QK5_1 / 2]);
for j in 0..16 {
let xi0 = l_buf[j];
let xi1 = l_buf[j + 16];
out[qs_pos + j] = (xi0 & 0x0F) | ((xi1 & 0x0F) << 4);
qh |= (((xi0 & 0x10) as u32) >> 4) << j;
qh |= (((xi1 & 0x10) as u32) >> 4) << (j + 16);
}
let qh_bytes = qh.to_le_bytes();
out[qh_pos..qh_pos + 4].copy_from_slice(&qh_bytes);
}
}
pub fn quantize(src: &[f32], n_per_row: usize, imatrix: Option<&[f32]>) -> Vec<u8> {
assert!(
n_per_row % QK5_1 == 0,
"n_per_row {} not multiple of QK5_1 {}",
n_per_row,
QK5_1
);
assert!(
src.len() % n_per_row == 0,
"src len {} not multiple of n_per_row {}",
src.len(),
n_per_row
);
if let Some(im) = imatrix {
assert_eq!(
im.len(),
n_per_row,
"imatrix length {} must equal n_per_row {}",
im.len(),
n_per_row
);
}
let nrow = src.len() / n_per_row;
let blocks_per_row = n_per_row / QK5_1;
let row_bytes = blocks_per_row * BLOCK_BYTES;
let mut out = vec![0u8; nrow * row_bytes];
out.par_chunks_exact_mut(row_bytes)
.enumerate()
.for_each(|(row, dst)| {
let row_x = &src[row * n_per_row..(row + 1) * n_per_row];
let mut tmp = Vec::with_capacity(row_bytes);
match imatrix {
None => quantize_row_q5_1_ref(row_x, &mut tmp),
Some(qw) => quantize_row_q5_1_impl(row_x, qw, &mut tmp),
}
debug_assert_eq!(tmp.len(), row_bytes);
dst.copy_from_slice(&tmp);
});
out
}
#[cfg(test)]
mod tests {
use super::*;
use std::fs;
use std::path::PathBuf;
fn fixture_path(name: &str) -> PathBuf {
let manifest =
std::env::var("CARGO_MANIFEST_DIR").expect("CARGO_MANIFEST_DIR not set by cargo test");
PathBuf::from(manifest)
.join("tests/fixtures/ggml_quants")
.join(name)
}
fn read_f32s(name: &str) -> Vec<f32> {
let bytes = fs::read(fixture_path(name)).expect("read fixture");
assert!(
bytes.len() % 4 == 0,
"fixture {} not a multiple of 4 bytes",
name
);
let mut out = Vec::with_capacity(bytes.len() / 4);
for chunk in bytes.chunks_exact(4) {
out.push(f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]));
}
out
}
fn read_bytes(name: &str) -> Vec<u8> {
fs::read(fixture_path(name)).expect("read fixture")
}
fn make_imatrix(n: usize, seed: u32) -> Vec<f32> {
let mut state = seed;
(0..n)
.map(|_| {
state = state.wrapping_add(0x6D2B79F5);
let mut t = state;
t = (t ^ (t >> 15)).wrapping_mul(t | 1);
t ^= t.wrapping_add((t ^ (t >> 7)).wrapping_mul(t | 61));
let u = t ^ (t >> 14);
let v = (u as f32 / u32::MAX as f32) * 2.0 - 1.0;
v.abs() + 1e-3
})
.collect()
}
#[test]
fn byte_cmp_noim() {
let input = read_f32s("q5_1_64_noim_input.bin");
let expected = read_bytes("q5_1_64_noim_expected.bin");
let got = quantize(&input, 64, None);
assert_eq!(got.len(), expected.len(), "Q5_1 noim length mismatch");
assert_eq!(got, expected, "Q5_1 noim byte-cmp failed");
}
#[test]
fn byte_cmp_im() {
let input = read_f32s("q5_1_64_im_input.bin");
let expected = read_bytes("q5_1_64_im_expected.bin");
let imatrix = make_imatrix(64, 2);
let got = quantize(&input, 64, Some(&imatrix));
assert_eq!(got.len(), expected.len(), "Q5_1 im length mismatch");
assert_eq!(got, expected, "Q5_1 im byte-cmp failed");
}
}