use std::io::{Read, Write};
use crate::error::Result;
use crate::vector::core::quantization::{PqParams, QuantizationMethod, VectorQuantizer, pq_decode};
use crate::vector::core::vector::Vector;
use crate::vector::index::pq_codebook::SharedPqCodebook;
#[inline]
#[cfg(test)]
pub(super) const fn pq_record_payload_size(m: u16) -> usize {
m as usize
}
pub(super) fn quantize_segment_pq(
vectors: &[Vector],
dim: usize,
subvector_count: usize,
shared: Option<&SharedPqCodebook>,
) -> Result<(PqParams, Vec<f32>, Vec<Vec<u8>>)> {
let quantizer = match shared {
Some(cb) => {
cb.validate_for(dim, subvector_count, 256)?;
VectorQuantizer::from_pq_codebook(dim, cb.params, cb.codebook.clone())?
}
None => {
let mut quantizer = VectorQuantizer::new(
QuantizationMethod::ProductQuantization { subvector_count },
dim,
);
quantizer.train(vectors)?;
quantizer
}
};
let (params, codebook_slice) = quantizer
.pq_state()
.expect("PQ quantizer is trained or built from an existing codebook, so pq_state is set");
let params = *params;
let codebook: Vec<f32> = codebook_slice.to_vec();
let codes: Vec<Vec<u8>> = vectors
.iter()
.map(|v| quantizer.quantize(v).map(|(c, _)| c))
.collect::<Result<_>>()?;
Ok((params, codebook, codes))
}
pub(super) fn write_pq_record<W: Write>(output: &mut W, codes: &[u8]) -> Result<()> {
output.write_all(codes)?;
Ok(())
}
pub(super) fn read_dequantized_pq_vector<R: Read>(
input: &mut R,
params: PqParams,
codebook: &[f32],
) -> Result<Vec<f32>> {
let mut codes = vec![0u8; params.m as usize];
input.read_exact(&mut codes)?;
Ok(pq_decode(&codes, params, codebook))
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::Cursor;
fn vec_of(values: &[f32]) -> Vector {
Vector::new(values.to_vec())
}
#[test]
fn quantize_segment_pq_returns_params_codebook_codes() {
let vectors = vec![
vec_of(&[10.0, 10.0, 20.0, 20.0]),
vec_of(&[-10.0, -10.0, -20.0, -20.0]),
vec_of(&[10.1, 10.1, 20.1, 20.1]),
];
let (params, codebook, codes) = quantize_segment_pq(&vectors, 4, 2, None).unwrap();
assert_eq!(params.m, 2);
assert_eq!(params.k, 256);
assert_eq!(params.sub_dim, 2);
assert_eq!(codebook.len(), params.codebook_len());
assert_eq!(codes.len(), 3);
for c in &codes {
assert_eq!(c.len(), 2);
}
}
#[test]
fn write_then_read_pq_roundtrips_to_codebook() {
let vectors = vec![
vec_of(&[10.0, 10.0, 20.0, 20.0]),
vec_of(&[-10.0, -10.0, -20.0, -20.0]),
];
let (params, codebook, codes) = quantize_segment_pq(&vectors, 4, 2, None).unwrap();
let mut buf = Vec::new();
for c in &codes {
write_pq_record(&mut buf, c).unwrap();
}
assert_eq!(buf.len(), codes.len() * pq_record_payload_size(params.m));
let mut cursor = Cursor::new(&buf);
for (i, original) in vectors.iter().enumerate() {
let recovered = read_dequantized_pq_vector(&mut cursor, params, &codebook).unwrap();
assert_eq!(recovered.len(), 4);
assert_eq!(
recovered[0].signum(),
original.data[0].signum(),
"vector {i} sign flip on first coord"
);
}
}
#[test]
fn payload_size_equals_m() {
assert_eq!(pq_record_payload_size(0), 0);
assert_eq!(pq_record_payload_size(16), 16);
}
#[test]
fn shared_codebook_produces_byte_identical_codes_to_fresh_training() {
let mut state: u64 = 0x1234_5678_9ABC_DEF0;
let vectors: Vec<Vector> = (0..300)
.map(|_| {
let data: Vec<f32> = (0..32)
.map(|_| {
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
})
.collect();
Vector::new(data)
})
.collect();
let (trained_params, trained_codebook, trained_codes) =
quantize_segment_pq(&vectors, 32, 4, None).unwrap();
let shared = crate::vector::index::pq_codebook::SharedPqCodebook {
params: trained_params,
codebook: trained_codebook.clone(),
};
let (shared_params, shared_codebook, shared_codes) =
quantize_segment_pq(&vectors, 32, 4, Some(&shared)).unwrap();
assert_eq!(shared_params, trained_params);
assert_eq!(shared_codebook, trained_codebook);
assert_eq!(shared_codes, trained_codes);
}
}