use half::f16;
use rayon::prelude::*;
use super::common::{make_qkx2_quants, make_qkx3_quants, make_qp_quants, nearest_int};
pub const QK_K: usize = 256;
pub const K_SCALE_SIZE: usize = 12;
pub const BLOCK_BYTES: usize = 2 + 2 + K_SCALE_SIZE + QK_K / 8 + QK_K / 2;
#[inline]
fn get_scale_min_k4(j: usize, q: &[u8]) -> (u8, u8) {
if j < 4 {
(q[j] & 63, q[j + 4] & 63)
} else {
let d = (q[j + 4] & 0xF) | ((q[j - 4] >> 6) << 4);
let m = (q[j + 4] >> 4) | ((q[j] >> 6) << 4);
(d, m)
}
}
#[inline]
fn pack_scale_min_k4(j: usize, scales: &mut [u8; K_SCALE_SIZE], ls: u8, lm: u8) {
if j < 4 {
scales[j] = ls;
scales[j + 4] = lm;
} else {
scales[j + 4] = (ls & 0xF) | ((lm & 0xF) << 4);
scales[j - 4] |= (ls >> 4) << 6;
scales[j] |= (lm >> 4) << 6;
}
}
fn finalize_block(
x: &[f32],
d_f16: f16,
dmin_f16: f16,
scales: &[u8; K_SCALE_SIZE],
l_init: &[u8; QK_K],
out: &mut Vec<u8>,
) {
let d_back = f16::to_f32(d_f16);
let dmin_back = f16::to_f32(dmin_f16);
let mut l: [u8; QK_K] = *l_init;
for j in 0..QK_K / 32 {
let (sc, m) = get_scale_min_k4(j, scales);
let d = d_back * (sc as f32);
if d == 0.0 {
continue;
}
let dm = dmin_back * (m as f32);
for ii in 0..32 {
let li = nearest_int((x[32 * j + ii] + dm) / d);
let clamped = li.max(0).min(31);
l[32 * j + ii] = clamped as u8;
}
}
out.extend_from_slice(&d_f16.to_le_bytes());
out.extend_from_slice(&dmin_f16.to_le_bytes());
out.extend_from_slice(scales);
let qh_pos = out.len();
out.extend_from_slice(&[0u8; QK_K / 8]);
let ql_pos = out.len();
out.extend_from_slice(&[0u8; QK_K / 2]);
let mut m1: u8 = 1;
let mut m2: u8 = 2;
for chunk in 0..4 {
let n = chunk * 64;
let ql_off = ql_pos + chunk * 32;
for j in 0..32 {
let mut l1 = l[n + j] as i32;
if l1 > 15 {
l1 -= 16;
out[qh_pos + j] |= m1;
}
let mut l2 = l[n + j + 32] as i32;
if l2 > 15 {
l2 -= 16;
out[qh_pos + j] |= m2;
}
out[ql_off + j] = (l1 as u8) | ((l2 as u8) << 4);
}
m1 <<= 2;
m2 <<= 2;
}
}
fn quantize_row_q5_k_ref(x: &[f32], out: &mut Vec<u8>) {
debug_assert!(x.len() % QK_K == 0);
let nb = x.len() / QK_K;
let mut mins = [0.0f32; QK_K / 32];
let mut scales_arr = [0.0f32; QK_K / 32];
let mut weights = [0.0f32; 32];
let mut l_buf = [0u8; 32];
let mut laux = [0u8; 32];
for i in 0..nb {
let xb = &x[i * QK_K..(i + 1) * QK_K];
let mut max_scale: f32 = 0.0;
let mut max_min: f32 = 0.0;
let mut l_full = [0u8; QK_K];
for j in 0..QK_K / 32 {
let mut sum_x2 = 0.0f32;
for l in 0..32 {
sum_x2 = xb[32 * j + l].mul_add(xb[32 * j + l], sum_x2);
}
let av_x = (sum_x2 / 32.0).sqrt();
for l in 0..32 {
weights[l] = av_x + xb[32 * j + l].abs();
}
let mut the_min = 0.0f32;
let scale = make_qkx2_quants(
32,
31,
&xb[32 * j..32 * j + 32],
&weights,
&mut l_full[32 * j..32 * j + 32],
&mut the_min,
&mut laux,
-0.5,
0.1,
15,
false,
);
scales_arr[j] = scale;
mins[j] = the_min;
let _ = &mut l_buf;
if scale > max_scale {
max_scale = scale;
}
if the_min > max_min {
max_min = the_min;
}
}
let inv_scale = if max_scale > 0.0 {
63.0 / max_scale
} else {
0.0
};
let inv_min = if max_min > 0.0 { 63.0 / max_min } else { 0.0 };
let mut scales_packed = [0u8; K_SCALE_SIZE];
for j in 0..QK_K / 32 {
let ls = (nearest_int(inv_scale * scales_arr[j]) as u8).min(63);
let lm = (nearest_int(inv_min * mins[j]) as u8).min(63);
pack_scale_min_k4(j, &mut scales_packed, ls, lm);
}
let d_f16 = f16::from_f32(max_scale / 63.0);
let dmin_f16 = f16::from_f32(max_min / 63.0);
finalize_block(xb, d_f16, dmin_f16, &scales_packed, &l_full, out);
}
}
fn quantize_row_q5_k_impl(x: &[f32], quant_weights: &[f32], out: &mut Vec<u8>) {
debug_assert!(x.len() % QK_K == 0);
let nb = x.len() / QK_K;
let mut mins = [0.0f32; QK_K / 32];
let mut scales_arr = [0.0f32; QK_K / 32];
let mut sw = [0.0f32; QK_K / 32];
let mut ls_arr = [0u8; QK_K / 32];
let mut lm_arr = [0u8; QK_K / 32];
let mut weights = [0.0f32; 32];
let mut l_full = [0u8; QK_K];
let mut laux = [0u8; 32];
for i in 0..nb {
let xb = &x[i * QK_K..(i + 1) * QK_K];
let qw_block = &quant_weights[i * QK_K..(i + 1) * QK_K];
let mut sum_x2 = 0.0f32;
for l in 0..QK_K {
sum_x2 = xb[l].mul_add(xb[l], sum_x2);
}
let sigma2 = 2.0 * sum_x2 / (QK_K as f32);
let _av_x = sigma2.sqrt();
for j in 0..QK_K / 32 {
let qw = &qw_block[32 * j..32 * j + 32];
for l in 0..32 {
weights[l] = qw[l] * (sigma2 + xb[32 * j + l] * xb[32 * j + l]).sqrt();
}
let mut sumw = 0.0f32;
for l in 0..32 {
sumw += weights[l];
}
sw[j] = sumw;
let mut the_min = 0.0f32;
let scale = make_qkx3_quants(
32,
31,
&xb[32 * j..32 * j + 32],
Some(&weights),
&mut l_full[32 * j..32 * j + 32],
&mut the_min,
&mut laux,
-0.9,
0.05,
36,
false,
);
scales_arr[j] = scale;
mins[j] = the_min;
}
let d_block = make_qp_quants(QK_K / 32, 63, &scales_arr, &mut ls_arr, &sw);
let m_block = make_qp_quants(QK_K / 32, 63, &mins, &mut lm_arr, &sw);
let mut scales_packed = [0u8; K_SCALE_SIZE];
for j in 0..QK_K / 32 {
let ls = ls_arr[j].min(63);
let lm = lm_arr[j].min(63);
pack_scale_min_k4(j, &mut scales_packed, ls, lm);
}
let d_f16 = f16::from_f32(d_block);
let dmin_f16 = f16::from_f32(m_block);
finalize_block(xb, d_f16, dmin_f16, &scales_packed, &l_full, out);
}
}
pub fn quantize(src: &[f32], n_per_row: usize, imatrix: Option<&[f32]>) -> Vec<u8> {
assert!(
n_per_row % QK_K == 0,
"n_per_row {} not multiple of QK_K {}",
n_per_row,
QK_K
);
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 row_blocks = n_per_row / QK_K;
let row_bytes = row_blocks * 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_k_ref(row_x, &mut tmp),
Some(qw) => quantize_row_q5_k_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_k_512_noim_input.bin");
let expected = read_bytes("q5_k_512_noim_expected.bin");
let got = quantize(&input, 512, None);
assert_eq!(got.len(), expected.len(), "Q5_K noim length mismatch");
assert_eq!(got, expected, "Q5_K noim byte-cmp failed");
}
#[test]
fn byte_cmp_im() {
let input = read_f32s("q5_k_512_im_input.bin");
let expected = read_bytes("q5_k_512_im_expected.bin");
let imatrix = make_imatrix(512, 2);
let got = quantize(&input, 512, Some(&imatrix));
assert_eq!(got.len(), expected.len(), "Q5_K im length mismatch");
assert_eq!(got, expected, "Q5_K im byte-cmp failed");
}
}