use serde::{Deserialize, Serialize};
use crate::error::{LaurusError, Result};
use crate::vector::core::vector::Vector;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
pub enum QuantizationMethod {
#[default]
Scalar8Bit,
ProductQuantization {
subvector_count: usize,
},
#[cfg(feature = "pq-fastscan")]
ProductQuantizationFastScan {
subvector_count: usize,
},
}
#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
pub struct ScalarQuantParams {
pub offset: f32,
pub scale: f32,
}
impl ScalarQuantParams {
pub fn train(vectors: &[Vector]) -> Result<Self> {
if vectors.is_empty() {
return Err(LaurusError::InvalidOperation(
"Cannot train scalar quantization on an empty vector set".to_string(),
));
}
Self::train_from_slices(vectors.iter().map(|v| v.data.as_slice()))
}
pub fn train_from_slices<'a>(slices: impl Iterator<Item = &'a [f32]>) -> Result<Self> {
let mut min_v = f32::INFINITY;
let mut max_v = f32::NEG_INFINITY;
let mut total_count: usize = 0;
for slice in slices {
for &x in slice {
if !x.is_finite() {
return Err(LaurusError::InvalidOperation(
"Training vectors contain NaN or infinite values".to_string(),
));
}
if x < min_v {
min_v = x;
}
if x > max_v {
max_v = x;
}
total_count += 1;
}
}
if total_count == 0 {
return Err(LaurusError::InvalidOperation(
"Training vectors are all empty (zero dimensions)".to_string(),
));
}
let range = max_v - min_v;
if range <= 0.0 {
return Ok(Self {
offset: min_v,
scale: 1.0,
});
}
Ok(Self {
offset: min_v,
scale: range / 255.0,
})
}
#[inline]
pub fn quantize_value(&self, v: f32) -> u8 {
let normalized = (v - self.offset) / self.scale;
normalized.round().clamp(0.0, 255.0) as u8
}
#[inline]
pub fn dequantize_value(&self, q: u8) -> f32 {
self.offset + self.scale * (q as f32)
}
pub fn quantize(&self, vector: &Vector) -> Vec<u8> {
vector
.data
.iter()
.map(|&v| self.quantize_value(v))
.collect()
}
pub fn quantize_slice(&self, data: &[f32]) -> Vec<u8> {
data.iter().map(|&v| self.quantize_value(v)).collect()
}
pub fn dequantize(&self, q: &[u8]) -> Vec<f32> {
q.iter().map(|&qi| self.dequantize_value(qi)).collect()
}
}
pub const PQ_KMEANS_ITERATIONS: usize = 25;
pub fn pq_train_codebook(
dimension: usize,
params: PqParams,
vectors: &[Vector],
) -> Result<Vec<f32>> {
if vectors.is_empty() {
return Err(LaurusError::InvalidOperation(
"Cannot train product quantization on an empty vector set".to_string(),
));
}
if dimension != params.original_dim() {
return Err(LaurusError::InvalidOperation(format!(
"PQ training dim mismatch: vectors imply {dimension}, params imply {}",
params.original_dim()
)));
}
for v in vectors {
if v.dimension() != dimension {
return Err(LaurusError::InvalidOperation(format!(
"PQ training input has mixed dimensions: expected {dimension}, got {}",
v.dimension()
)));
}
}
let k = params.k as usize;
let sub_dim = params.sub_dim as usize;
let n = vectors.len();
let effective_k = k.min(n.max(1));
let mut codebook = vec![0.0_f32; params.codebook_len()];
let train_sub = |(sub, slot): (usize, &mut [f32])| {
let mut sub_data: Vec<f32> = Vec::with_capacity(n * sub_dim);
for v in vectors {
sub_data.extend_from_slice(&v.data[sub * sub_dim..(sub + 1) * sub_dim]);
}
let seed: u64 =
0xCAFE_F00D_DEAD_BEEF_u64 ^ ((sub as u64).wrapping_mul(0x9E37_79B9_7F4A_7C15));
let centroids = kmeans_train(&sub_data, sub_dim, effective_k, PQ_KMEANS_ITERATIONS, seed);
let src_len = effective_k * sub_dim;
slot[..src_len].copy_from_slice(¢roids[..src_len]);
};
#[cfg(feature = "native")]
{
use rayon::prelude::*;
codebook
.par_chunks_mut(k * sub_dim)
.enumerate()
.for_each(train_sub);
Ok(codebook)
}
#[cfg(not(feature = "native"))]
{
codebook
.chunks_mut(k * sub_dim)
.enumerate()
.for_each(train_sub);
Ok(codebook)
}
}
pub fn pq_encode(vector: &[f32], params: PqParams, codebook: &[f32]) -> Vec<u8> {
debug_assert_eq!(codebook.len(), params.codebook_len());
let m = params.m as usize;
let k = params.k as usize;
let sub_dim = params.sub_dim as usize;
let mut codes = Vec::with_capacity(m);
for sub in 0..m {
let q_sub = &vector[sub * sub_dim..(sub + 1) * sub_dim];
let base = sub * k * sub_dim;
let mut best_k: u8 = 0;
let mut best_d = f32::INFINITY;
for ki in 0..k {
let c = &codebook[base + ki * sub_dim..base + (ki + 1) * sub_dim];
let sum = l2_squared_slice(q_sub, c);
if sum < best_d {
best_d = sum;
best_k = ki as u8;
}
}
codes.push(best_k);
}
codes
}
pub fn pq_decode(codes: &[u8], params: PqParams, codebook: &[f32]) -> Vec<f32> {
debug_assert_eq!(codes.len(), params.m as usize);
debug_assert_eq!(codebook.len(), params.codebook_len());
let m = params.m as usize;
let k = params.k as usize;
let sub_dim = params.sub_dim as usize;
let mut out = Vec::with_capacity(m * sub_dim);
for (sub, &code) in codes.iter().enumerate().take(m) {
let base = sub * k * sub_dim + code as usize * sub_dim;
out.extend_from_slice(&codebook[base..base + sub_dim]);
}
out
}
fn kmeans_train(data: &[f32], sub_dim: usize, k: usize, iters: usize, seed: u64) -> Vec<f32> {
let n = data.len() / sub_dim;
debug_assert_eq!(data.len(), n * sub_dim);
debug_assert!(n >= 1, "k-means needs at least 1 point");
let k = k.min(n.max(1));
let mut state = seed;
let mut centroids = vec![0.0_f32; k * sub_dim];
{
state = state
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1_442_695_040_888_963_407);
let first = ((state >> 32) as usize) % n;
centroids[..sub_dim].copy_from_slice(&data[first * sub_dim..(first + 1) * sub_dim]);
let mut min_d2: Vec<f32> = (0..n)
.map(|i| {
let p = &data[i * sub_dim..(i + 1) * sub_dim];
l2_squared_slice(p, ¢roids[..sub_dim])
})
.collect();
for c_idx in 1..k {
let total: f32 = min_d2.iter().sum();
let chosen = if total == 0.0 {
state = state
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1_442_695_040_888_963_407);
((state >> 32) as usize) % n
} else {
state = state
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1_442_695_040_888_963_407);
let r = ((state >> 32) as f32 / u32::MAX as f32) * total;
let mut acc = 0.0_f32;
let mut pick = n - 1;
for (i, &d) in min_d2.iter().enumerate() {
acc += d;
if acc >= r {
pick = i;
break;
}
}
pick
};
centroids[c_idx * sub_dim..(c_idx + 1) * sub_dim]
.copy_from_slice(&data[chosen * sub_dim..(chosen + 1) * sub_dim]);
let new_c = ¢roids[c_idx * sub_dim..(c_idx + 1) * sub_dim];
for i in 0..n {
let p = &data[i * sub_dim..(i + 1) * sub_dim];
let d = l2_squared_slice(p, new_c);
if d < min_d2[i] {
min_d2[i] = d;
}
}
}
}
let mut sums = vec![0.0_f32; k * sub_dim];
let mut counts = vec![0u32; k];
for _ in 0..iters {
sums.iter_mut().for_each(|x| *x = 0.0);
counts.iter_mut().for_each(|c| *c = 0);
for i in 0..n {
let p = &data[i * sub_dim..(i + 1) * sub_dim];
let mut best_j = 0usize;
let mut best_d = l2_squared_slice(p, ¢roids[..sub_dim]);
for j in 1..k {
let c = ¢roids[j * sub_dim..(j + 1) * sub_dim];
let d = l2_squared_slice(p, c);
if d < best_d {
best_d = d;
best_j = j;
}
}
counts[best_j] += 1;
let sum_base = best_j * sub_dim;
for d in 0..sub_dim {
sums[sum_base + d] += p[d];
}
}
for j in 0..k {
if counts[j] > 0 {
let inv = 1.0 / counts[j] as f32;
for d in 0..sub_dim {
centroids[j * sub_dim + d] = sums[j * sub_dim + d] * inv;
}
}
}
}
centroids
}
#[inline]
fn l2_squared_slice(a: &[f32], b: &[f32]) -> f32 {
use wide::f32x8;
debug_assert_eq!(a.len(), b.len());
let mut acc = f32x8::ZERO;
let chunks_a = a.chunks_exact(8);
let chunks_b = b.chunks_exact(8);
let rem_a = chunks_a.remainder();
let rem_b = chunks_b.remainder();
for (ca, cb) in chunks_a.zip(chunks_b) {
let diff = f32x8::from(ca) - f32x8::from(cb);
acc += diff * diff;
}
let mut sum: f32 = acc.reduce_add();
for (x, y) in rem_a.iter().zip(rem_b.iter()) {
let d = x - y;
sum += d * d;
}
sum
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub struct PqParams {
pub m: u16,
pub k: u16,
pub sub_dim: u16,
}
impl PqParams {
pub fn new(m: u16, k: u16, sub_dim: u16) -> Result<Self> {
if m == 0 || k == 0 || sub_dim == 0 {
return Err(LaurusError::InvalidOperation(format!(
"PqParams components must be > 0 (got m={m}, k={k}, sub_dim={sub_dim})"
)));
}
if !matches!(k, 16 | 256) {
return Err(LaurusError::InvalidOperation(format!(
"PqParams::k must be one of {{16, 256}} (got {k}); 256 is the \
8-bit PQ variant (Issue #481 Stage 3), 16 is the FastScan \
4-bit variant (Issue #651 / #692)"
)));
}
Ok(Self { m, k, sub_dim })
}
pub fn from_dim_and_m(dim: usize, m: usize) -> Result<Self> {
Self::from_dim_and_m_k(dim, m, 256)
}
pub fn from_dim_and_m_k(dim: usize, m: usize, k: u16) -> Result<Self> {
if m == 0 {
return Err(LaurusError::InvalidOperation(
"Product quantization subvector_count must be > 0".to_string(),
));
}
if !dim.is_multiple_of(m) {
return Err(LaurusError::InvalidOperation(format!(
"Product quantization subvector_count {m} must divide vector \
dimension {dim} (got {dim} % {m} = {})",
dim % m
)));
}
let sub_dim = dim / m;
Self::new(
u16::try_from(m).map_err(|_| {
LaurusError::InvalidOperation(format!(
"Product quantization subvector_count {m} exceeds u16::MAX"
))
})?,
k,
u16::try_from(sub_dim).map_err(|_| {
LaurusError::InvalidOperation(format!(
"Product quantization sub_dim {sub_dim} exceeds u16::MAX"
))
})?,
)
}
#[inline]
pub fn codebook_len(&self) -> usize {
self.m as usize * self.k as usize * self.sub_dim as usize
}
#[inline]
pub fn codebook_byte_size(&self) -> usize {
self.codebook_len() * 4
}
#[inline]
pub fn original_dim(&self) -> usize {
self.m as usize * self.sub_dim as usize
}
}
#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
#[repr(C)]
pub struct QuantizedVectorMeta {
pub sum_q: u32,
pub norm_q: f32,
}
impl QuantizedVectorMeta {
pub fn from_quantized(q: &[u8], params: &ScalarQuantParams) -> Self {
let mut sum_q: u32 = 0;
let mut norm_sq: f32 = 0.0;
for &qi in q {
sum_q += qi as u32;
let dq = params.dequantize_value(qi);
norm_sq += dq * dq;
}
Self {
sum_q,
norm_q: norm_sq.sqrt(),
}
}
pub const SERIALIZED_SIZE: usize = 8;
}
#[derive(Debug, Clone)]
pub struct VectorQuantizer {
method: QuantizationMethod,
dimension: usize,
state: QuantizerState,
}
#[derive(Debug, Clone)]
enum QuantizerState {
Untrained,
Scalar8Bit(ScalarQuantParams),
ProductQuantization {
params: PqParams,
codebook: Vec<f32>,
},
}
impl VectorQuantizer {
pub fn new(method: QuantizationMethod, dimension: usize) -> Self {
Self {
method,
dimension,
state: QuantizerState::Untrained,
}
}
pub fn from_params(
method: QuantizationMethod,
dimension: usize,
params: ScalarQuantParams,
) -> Result<Self> {
match method {
QuantizationMethod::Scalar8Bit => Ok(Self {
method,
dimension,
state: QuantizerState::Scalar8Bit(params),
}),
QuantizationMethod::ProductQuantization { .. } => Err(LaurusError::InvalidOperation(
"Use VectorQuantizer::from_pq_codebook for ProductQuantization; \
from_params is Scalar8Bit-only"
.to_string(),
)),
#[cfg(feature = "pq-fastscan")]
QuantizationMethod::ProductQuantizationFastScan { .. } => {
Err(LaurusError::InvalidOperation(
"Use VectorQuantizer::from_pq_codebook for ProductQuantizationFastScan; \
from_params is Scalar8Bit-only"
.to_string(),
))
}
}
}
pub fn from_pq_codebook(
dimension: usize,
params: PqParams,
codebook: Vec<f32>,
) -> Result<Self> {
if params.original_dim() != dimension {
return Err(LaurusError::InvalidOperation(format!(
"PQ codebook dim mismatch: params imply {}, quantizer wants {dimension}",
params.original_dim()
)));
}
if codebook.len() != params.codebook_len() {
return Err(LaurusError::InvalidOperation(format!(
"PQ codebook length {} does not match params (m={}, k={}, sub_dim={} -> {})",
codebook.len(),
params.m,
params.k,
params.sub_dim,
params.codebook_len()
)));
}
Ok(Self {
method: QuantizationMethod::ProductQuantization {
subvector_count: params.m as usize,
},
dimension,
state: QuantizerState::ProductQuantization { params, codebook },
})
}
pub fn train(&mut self, vectors: &[Vector]) -> Result<()> {
match self.method {
QuantizationMethod::Scalar8Bit => {
let params = ScalarQuantParams::train(vectors)?;
self.state = QuantizerState::Scalar8Bit(params);
Ok(())
}
QuantizationMethod::ProductQuantization { subvector_count } => {
let params = PqParams::from_dim_and_m(self.dimension, subvector_count)?;
let codebook = pq_train_codebook(self.dimension, params, vectors)?;
self.state = QuantizerState::ProductQuantization { params, codebook };
Ok(())
}
#[cfg(feature = "pq-fastscan")]
QuantizationMethod::ProductQuantizationFastScan { subvector_count } => {
let params = PqParams::from_dim_and_m_k(self.dimension, subvector_count, 16)?;
let codebook = pq_train_codebook(self.dimension, params, vectors)?;
self.state = QuantizerState::ProductQuantization { params, codebook };
Ok(())
}
}
}
pub fn quantize(&self, vector: &Vector) -> Result<(Vec<u8>, QuantizedVectorMeta)> {
if vector.dimension() != self.dimension {
return Err(LaurusError::InvalidOperation(format!(
"Vector dimension mismatch: expected {}, got {}",
self.dimension,
vector.dimension()
)));
}
match &self.state {
QuantizerState::Untrained => Err(LaurusError::InvalidOperation(
"Quantizer must be trained before quantizing vectors".to_string(),
)),
QuantizerState::Scalar8Bit(params) => {
let q = params.quantize(vector);
let meta = QuantizedVectorMeta::from_quantized(&q, params);
Ok((q, meta))
}
QuantizerState::ProductQuantization { params, codebook } => {
let codes = pq_encode(&vector.data, *params, codebook);
Ok((
codes,
QuantizedVectorMeta {
sum_q: 0,
norm_q: 0.0,
},
))
}
}
}
pub fn params(&self) -> Option<&ScalarQuantParams> {
match &self.state {
QuantizerState::Scalar8Bit(p) => Some(p),
_ => None,
}
}
pub fn pq_state(&self) -> Option<(&PqParams, &[f32])> {
match &self.state {
QuantizerState::ProductQuantization { params, codebook } => {
Some((params, codebook.as_slice()))
}
_ => None,
}
}
pub fn method(&self) -> QuantizationMethod {
self.method
}
pub fn dimension(&self) -> usize {
self.dimension
}
pub fn is_trained(&self) -> bool {
!matches!(self.state, QuantizerState::Untrained)
}
pub fn compression_ratio(&self) -> f32 {
match self.method {
QuantizationMethod::Scalar8Bit => 4.0,
QuantizationMethod::ProductQuantization { subvector_count } => {
if subvector_count == 0 {
1.0
} else {
(self.dimension * 4) as f32 / subvector_count as f32
}
}
#[cfg(feature = "pq-fastscan")]
QuantizationMethod::ProductQuantizationFastScan { subvector_count } => {
if subvector_count == 0 {
1.0
} else {
(self.dimension * 4) as f32 / (subvector_count as f32 / 2.0)
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn vec_of(values: &[f32]) -> Vector {
Vector::new(values.to_vec())
}
#[test]
fn train_picks_min_and_max_across_all_vectors() {
let vectors = vec![
vec_of(&[0.0, 1.0, 2.0]),
vec_of(&[-1.5, 3.0, 0.5]),
vec_of(&[5.0, -2.0, 1.0]),
];
let params = ScalarQuantParams::train(&vectors).unwrap();
assert_eq!(params.offset, -2.0); assert!((params.scale - (5.0 - -2.0) / 255.0).abs() < 1e-7);
}
#[test]
fn train_on_empty_set_is_invalid_operation() {
let err = ScalarQuantParams::train(&[]).unwrap_err();
assert!(matches!(err, LaurusError::InvalidOperation(_)));
}
#[test]
fn train_on_constant_data_uses_unit_scale() {
let vectors = vec![vec_of(&[3.5, 3.5, 3.5]); 4];
let params = ScalarQuantParams::train(&vectors).unwrap();
assert_eq!(params.offset, 3.5);
assert_eq!(params.scale, 1.0);
for &v in &[3.5, 3.5, 3.5] {
assert_eq!(params.quantize_value(v), 0);
}
}
#[test]
fn train_rejects_non_finite_values() {
let vectors = vec![vec_of(&[1.0, f32::NAN, 2.0])];
let err = ScalarQuantParams::train(&vectors).unwrap_err();
assert!(matches!(err, LaurusError::InvalidOperation(_)));
let vectors = vec![vec_of(&[1.0, f32::INFINITY, 2.0])];
let err = ScalarQuantParams::train(&vectors).unwrap_err();
assert!(matches!(err, LaurusError::InvalidOperation(_)));
}
#[test]
fn train_from_slices_matches_train_on_equivalent_input() {
let vectors = vec![
vec_of(&[0.0, 1.0, 2.0]),
vec_of(&[-1.5, 3.0, 0.5]),
vec_of(&[5.0, -2.0, 1.0]),
];
let via_vectors = ScalarQuantParams::train(&vectors).unwrap();
let slices: Vec<&[f32]> = vectors.iter().map(|v| v.data.as_slice()).collect();
let via_slices = ScalarQuantParams::train_from_slices(slices.into_iter()).unwrap();
assert_eq!(via_vectors, via_slices);
}
#[test]
fn train_from_slices_rejects_empty_iterator() {
let err = ScalarQuantParams::train_from_slices(std::iter::empty()).unwrap_err();
assert!(matches!(err, LaurusError::InvalidOperation(_)));
let err = ScalarQuantParams::train_from_slices([[].as_slice()].into_iter()).unwrap_err();
assert!(matches!(err, LaurusError::InvalidOperation(_)));
}
#[test]
fn quantize_value_roundtrips_within_scale() {
let params = ScalarQuantParams {
offset: -1.0,
scale: 2.0 / 255.0, };
for v in [-1.0_f32, -0.5, 0.0, 0.25, 0.99] {
let q = params.quantize_value(v);
let dq = params.dequantize_value(q);
assert!(
(v - dq).abs() <= params.scale,
"v = {v}, dq = {dq}, scale = {}",
params.scale
);
}
}
#[test]
fn quantize_value_saturates_outside_training_range() {
let params = ScalarQuantParams {
offset: 0.0,
scale: 1.0 / 255.0, };
assert_eq!(params.quantize_value(-5.0), 0);
assert_eq!(params.quantize_value(5.0), 255);
}
#[test]
fn quantize_vector_roundtrip_within_scale() {
let vectors = vec![
vec_of(&[-1.0, -0.5, 0.0, 0.5, 1.0]),
vec_of(&[0.1, -0.2, 0.3, -0.4, 0.5]),
];
let params = ScalarQuantParams::train(&vectors).unwrap();
for v in &vectors {
let q = params.quantize(v);
let dq = params.dequantize(&q);
for (orig, recovered) in v.data.iter().zip(dq.iter()) {
assert!(
(orig - recovered).abs() <= params.scale,
"orig = {orig}, recovered = {recovered}, scale = {}",
params.scale
);
}
}
}
#[test]
fn meta_sum_q_matches_summed_quantized_bytes() {
let params = ScalarQuantParams {
offset: 0.0,
scale: 1.0 / 255.0,
};
let q: Vec<u8> = vec![0, 64, 128, 200, 255];
let meta = QuantizedVectorMeta::from_quantized(&q, ¶ms);
let expected_sum: u32 = q.iter().map(|&x| x as u32).sum();
assert_eq!(meta.sum_q, expected_sum);
}
#[test]
fn meta_norm_q_matches_dequantized_norm() {
let params = ScalarQuantParams {
offset: -1.0,
scale: 2.0 / 255.0,
};
let q: Vec<u8> = vec![0, 64, 128, 192, 255];
let meta = QuantizedVectorMeta::from_quantized(&q, ¶ms);
let dq = params.dequantize(&q);
let expected: f32 = dq.iter().map(|x| x * x).sum::<f32>().sqrt();
assert!((meta.norm_q - expected).abs() < 1e-5);
}
#[test]
fn quantizer_requires_training_before_quantize() {
let q = VectorQuantizer::new(QuantizationMethod::Scalar8Bit, 3);
let err = q.quantize(&vec_of(&[1.0, 2.0, 3.0])).unwrap_err();
assert!(matches!(err, LaurusError::InvalidOperation(_)));
}
#[test]
fn quantizer_train_then_quantize_returns_data_and_meta() {
let mut q = VectorQuantizer::new(QuantizationMethod::Scalar8Bit, 3);
let training = vec![vec_of(&[-1.0, 0.0, 1.0]), vec_of(&[-0.5, 0.5, 0.25])];
q.train(&training).unwrap();
assert!(q.is_trained());
let (bytes, meta) = q.quantize(&vec_of(&[0.0, 0.5, -0.25])).unwrap();
assert_eq!(bytes.len(), 3);
let expected_sum: u32 = bytes.iter().map(|&x| x as u32).sum();
assert_eq!(meta.sum_q, expected_sum);
assert!(meta.norm_q.is_finite() && meta.norm_q >= 0.0);
}
#[test]
fn quantizer_rejects_dimension_mismatch() {
let mut q = VectorQuantizer::new(QuantizationMethod::Scalar8Bit, 3);
q.train(&[vec_of(&[1.0, 2.0, 3.0])]).unwrap();
let err = q.quantize(&vec_of(&[1.0, 2.0])).unwrap_err();
assert!(matches!(err, LaurusError::InvalidOperation(_)));
}
#[test]
fn product_quantization_train_succeeds_on_valid_inputs() {
let dim = 8usize;
let m = 4usize;
let mut q = VectorQuantizer::new(
QuantizationMethod::ProductQuantization { subvector_count: m },
dim,
);
let training: Vec<Vector> = (0..32)
.map(|i| {
let v: Vec<f32> = (0..dim).map(|j| (i + j) as f32 * 0.1).collect();
vec_of(&v)
})
.collect();
q.train(&training).unwrap();
assert!(q.is_trained());
let (pq_params, codebook) = q.pq_state().expect("PQ trained");
assert_eq!(pq_params.m as usize, m);
assert_eq!(pq_params.k, 256);
assert_eq!(pq_params.sub_dim as usize, dim / m);
assert_eq!(codebook.len(), pq_params.codebook_len());
assert!(q.params().is_none());
}
#[test]
fn product_quantization_encode_decode_roundtrips_to_codebook() {
let dim = 4usize;
let m = 2usize;
let mut q = VectorQuantizer::new(
QuantizationMethod::ProductQuantization { subvector_count: m },
dim,
);
let training = vec![
vec_of(&[10.0, 10.0, 20.0, 20.0]),
vec_of(&[-10.0, -10.0, -20.0, -20.0]),
];
q.train(&training).unwrap();
let (codes, _meta) = q
.quantize(&vec_of(&[10.5, 10.5, 20.5, 20.5]))
.expect("quantize");
assert_eq!(codes.len(), m);
let (params, cb) = q.pq_state().unwrap();
let decoded = pq_decode(&codes, *params, cb);
assert_eq!(decoded.len(), dim);
}
#[test]
fn product_quantization_from_pq_codebook_roundtrips() {
let params = PqParams::new(4, 256, 2).unwrap();
let codebook = vec![0.0_f32; params.codebook_len()];
let q = VectorQuantizer::from_pq_codebook(8, params, codebook.clone()).unwrap();
assert!(q.is_trained());
assert_eq!(q.pq_state().unwrap().0, ¶ms);
assert_eq!(q.pq_state().unwrap().1.len(), codebook.len());
}
#[test]
fn product_quantization_from_pq_codebook_rejects_size_mismatch() {
let params = PqParams::new(4, 256, 2).unwrap();
let bad = vec![0.0_f32; params.codebook_len() - 1];
let err = VectorQuantizer::from_pq_codebook(8, params, bad).unwrap_err();
assert!(matches!(err, LaurusError::InvalidOperation(_)));
}
#[test]
fn pq_params_validate_divides_dim() {
let err = PqParams::from_dim_and_m(10, 3).unwrap_err();
assert!(matches!(err, LaurusError::InvalidOperation(_)));
let ok = PqParams::from_dim_and_m(12, 3).unwrap();
assert_eq!(ok.original_dim(), 12);
assert_eq!(ok.sub_dim, 4);
}
#[test]
fn from_params_roundtrips_for_scalar8bit() {
let params = ScalarQuantParams {
offset: -2.0,
scale: 4.0 / 255.0,
};
let q = VectorQuantizer::from_params(QuantizationMethod::Scalar8Bit, 5, params).unwrap();
assert_eq!(q.method(), QuantizationMethod::Scalar8Bit);
assert_eq!(q.dimension(), 5);
assert_eq!(q.params(), Some(¶ms));
}
#[test]
fn default_method_is_scalar_8bit() {
let m: QuantizationMethod = Default::default();
assert_eq!(m, QuantizationMethod::Scalar8Bit);
}
#[test]
fn compression_ratio_scalar_is_4x() {
let q = VectorQuantizer::new(QuantizationMethod::Scalar8Bit, 128);
assert_eq!(q.compression_ratio(), 4.0);
}
#[test]
fn meta_serialized_size_is_eight() {
assert_eq!(QuantizedVectorMeta::SERIALIZED_SIZE, 8);
}
#[test]
fn pq_train_codebook_is_byte_identical_across_runs() {
let dim = 16usize;
let m = 4usize;
let params = PqParams::from_dim_and_m(dim, m).unwrap();
let mut state: u64 = 0x5EED_1234_ABCD_9876;
let mut next = move || {
state = state
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1_442_695_040_888_963_407);
((state >> 33) as f32 / u32::MAX as f32) * 2.0 - 1.0
};
let vectors: Vec<Vector> = (0..512)
.map(|_| Vector::new((0..dim).map(|_| next()).collect()))
.collect();
let a = pq_train_codebook(dim, params, &vectors).expect("train a");
let b = pq_train_codebook(dim, params, &vectors).expect("train b");
assert_eq!(a.len(), b.len());
let a_bits: Vec<u32> = a.iter().map(|f| f.to_bits()).collect();
let b_bits: Vec<u32> = b.iter().map(|f| f.to_bits()).collect();
assert_eq!(
a_bits, b_bits,
"two pq_train_codebook runs must be byte-identical"
);
}
}