use crate::{Result, VQuantError};
#[derive(Debug, Clone, Copy)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct MatryoshkaQuantizer {
lo: f32,
hi: f32,
max_bits: u8,
}
impl MatryoshkaQuantizer {
pub fn new(lo: f32, hi: f32, max_bits: u8) -> Result<Self> {
if !(1..=8).contains(&max_bits) {
return Err(VQuantError::InvalidConfig {
field: "max_bits",
reason: "must be in 1..=8",
});
}
if !lo.is_finite() || !hi.is_finite() || hi <= lo {
return Err(VQuantError::InvalidConfig {
field: "range",
reason: "require finite lo < hi",
});
}
Ok(Self { lo, hi, max_bits })
}
pub fn fit_minmax(values: &[f32], max_bits: u8) -> Result<Self> {
let (lo, hi) = finite_min_max(values)?;
Self::new(lo, hi, max_bits)
}
pub fn fit(values: &[f32], max_bits: u8, precisions: &[u8], weights: &[f32]) -> Result<Self> {
if precisions.is_empty() || precisions.len() != weights.len() {
return Err(VQuantError::InvalidConfig {
field: "precisions/weights",
reason: "non-empty and equal length required",
});
}
for &r in precisions {
if !(1..=max_bits).contains(&r) {
return Err(VQuantError::InvalidConfig {
field: "precisions",
reason: "each precision must be in 1..=max_bits",
});
}
}
let (lo0, hi0) = finite_min_max(values)?;
let mut sorted: Vec<f32> = values.iter().copied().filter(|v| v.is_finite()).collect();
sorted.sort_by(|a, b| a.partial_cmp(b).unwrap());
const GAMMAS: [f32; 7] = [0.0, 0.0025, 0.005, 0.01, 0.02, 0.05, 0.1];
let mut best: Option<(f32, Self)> = None;
for &gamma in &GAMMAS {
let (lo, hi) = if gamma == 0.0 {
(lo0, hi0)
} else {
(quantile(&sorted, gamma), quantile(&sorted, 1.0 - gamma))
};
let Ok(q) = Self::new(lo, hi, max_bits) else {
continue;
};
let err = q.weighted_mse(values, precisions, weights);
if best.as_ref().map_or(true, |(b, _)| err < *b) {
best = Some((err, q));
}
}
best.map(|(_, q)| q).ok_or(VQuantError::InvalidConfig {
field: "values",
reason: "no valid clip range found",
})
}
pub fn max_bits(&self) -> u8 {
self.max_bits
}
pub fn quantize(&self, value: f32) -> u8 {
let levels = ((1u16 << self.max_bits) - 1) as f32; let t = ((value - self.lo) / (self.hi - self.lo)).clamp(0.0, 1.0);
(t * levels).round() as u8
}
pub fn slice(&self, code: u8, bits: u8) -> u8 {
let r = bits.clamp(1, self.max_bits);
let shift = self.max_bits - r;
let max_r = (1u16 << r) - 1;
let sliced = ((code as f32) / ((1u16 << shift) as f32)).round() as u16;
sliced.min(max_r) as u8
}
pub fn dequantize(&self, code_r: u8, bits: u8) -> f32 {
let r = bits.clamp(1, self.max_bits);
let shift = self.max_bits - r;
let levels = ((1u16 << self.max_bits) - 1) as f32; let grid_value = (code_r as u16 * (1u16 << shift)) as f32; self.lo + (grid_value / levels) * (self.hi - self.lo)
}
pub fn reconstruct(&self, value: f32, bits: u8) -> f32 {
let code = self.quantize(value);
let sliced = self.slice(code, bits);
self.dequantize(sliced, bits)
}
pub fn weighted_mse(&self, values: &[f32], precisions: &[u8], weights: &[f32]) -> f32 {
let mut total = 0.0f32;
let mut wsum = 0.0f32;
for (&r, &w) in precisions.iter().zip(weights.iter()) {
let mut se = 0.0f32;
let mut n = 0usize;
for &v in values.iter().filter(|v| v.is_finite()) {
let d = v - self.reconstruct(v, r);
se += d * d;
n += 1;
}
if n > 0 {
total += w * (se / n as f32);
wsum += w;
}
}
if wsum > 0.0 {
total / wsum
} else {
0.0
}
}
}
fn finite_min_max(values: &[f32]) -> Result<(f32, f32)> {
let mut lo = f32::INFINITY;
let mut hi = f32::NEG_INFINITY;
for &v in values {
if v.is_finite() {
lo = lo.min(v);
hi = hi.max(v);
}
}
if !lo.is_finite() || !hi.is_finite() {
return Err(VQuantError::InvalidConfig {
field: "values",
reason: "no finite values",
});
}
if hi <= lo {
hi = lo + 1.0; }
Ok((lo, hi))
}
fn quantile(sorted: &[f32], q: f32) -> f32 {
if sorted.is_empty() {
return 0.0;
}
let pos = q.clamp(0.0, 1.0) * (sorted.len() - 1) as f32;
let lo = pos.floor() as usize;
let hi = pos.ceil() as usize;
let frac = pos - lo as f32;
sorted[lo] * (1.0 - frac) + sorted[hi] * frac
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn slice_matches_eq6() {
let q = MatryoshkaQuantizer::new(0.0, 1.0, 8).unwrap();
assert_eq!(q.slice(200, 2), 3);
assert_eq!(q.slice(95, 4), 6);
assert_eq!(q.slice(123, 8), 123);
}
#[test]
fn reconstruction_error_decreases_with_bits() {
let values: Vec<f32> = (0..256).map(|i| i as f32 / 255.0).collect();
let q = MatryoshkaQuantizer::fit_minmax(&values, 8).unwrap();
let e8 = q.weighted_mse(&values, &[8], &[1.0]);
let e4 = q.weighted_mse(&values, &[4], &[1.0]);
let e2 = q.weighted_mse(&values, &[2], &[1.0]);
assert!(e8 <= e4, "8-bit error {e8} should be <= 4-bit {e4}");
assert!(e4 <= e2, "4-bit error {e4} should be <= 2-bit {e2}");
}
#[test]
fn fit_beats_minmax_at_low_bits() {
let mut values: Vec<f32> = Vec::new();
for i in 0..200 {
let x = ((i as f32) * 0.0173).sin();
values.push(x);
}
values.extend_from_slice(&[40.0, -38.0, 45.0]);
let precisions = [8u8, 4, 2];
let weights = [0.1f32, 0.1, 1.0]; let fitted = MatryoshkaQuantizer::fit(&values, 8, &precisions, &weights).unwrap();
let baseline = MatryoshkaQuantizer::fit_minmax(&values, 8).unwrap();
let core: Vec<f32> = values.iter().copied().filter(|v| v.abs() <= 1.0).collect();
let e_fit = fitted.weighted_mse(&core, &[2], &[1.0]);
let e_base = baseline.weighted_mse(&core, &[2], &[1.0]);
assert!(
e_fit < e_base,
"fit int2 core error {e_fit} should beat min/max {e_base}"
);
}
#[test]
fn rejects_bad_config() {
assert!(MatryoshkaQuantizer::new(0.0, 1.0, 0).is_err());
assert!(MatryoshkaQuantizer::new(0.0, 1.0, 9).is_err());
assert!(MatryoshkaQuantizer::new(1.0, 1.0, 8).is_err());
assert!(MatryoshkaQuantizer::fit_minmax(&[], 8).is_err());
}
use proptest::prelude::*;
proptest! {
#[test]
fn slice_stays_in_b_bit_range(code in any::<u8>(), bits in 1u8..=8) {
let q = MatryoshkaQuantizer::new(0.0, 1.0, 8).unwrap();
let sliced = q.slice(code, bits);
let max_code = (1u16 << bits) - 1;
prop_assert!(
u16::from(sliced) <= max_code,
"slice({code}, {bits}) = {sliced} exceeds {max_code}"
);
}
#[test]
fn slice_at_full_precision_is_identity(code in any::<u8>()) {
let q = MatryoshkaQuantizer::new(0.0, 1.0, 8).unwrap();
prop_assert_eq!(q.slice(code, 8), code);
}
#[test]
fn slice_preserves_code_order(a in any::<u8>(), b in any::<u8>(), bits in 1u8..=8) {
let q = MatryoshkaQuantizer::new(0.0, 1.0, 8).unwrap();
let (lo, hi) = if a <= b { (a, b) } else { (b, a) };
prop_assert!(
q.slice(lo, bits) <= q.slice(hi, bits),
"slice not monotone at {bits} bits: slice({lo})={}, slice({hi})={}",
q.slice(lo, bits),
q.slice(hi, bits)
);
}
}
}