use std::io::{Read, Write};
use crate::error::Result;
use crate::vector::core::quantization::{
QuantizationMethod, QuantizedVectorMeta, ScalarQuantParams, VectorQuantizer,
};
use crate::vector::core::vector::Vector;
pub(super) type QuantizedRecord = (Vec<u8>, QuantizedVectorMeta);
#[inline]
pub(super) const fn quantized_record_payload_size(dim: usize) -> usize {
dim + QuantizedVectorMeta::SERIALIZED_SIZE
}
pub(super) fn quantize_segment(
vectors: &[Vector],
dim: usize,
) -> Result<(ScalarQuantParams, Vec<QuantizedRecord>)> {
let mut quantizer = VectorQuantizer::new(QuantizationMethod::Scalar8Bit, dim);
quantizer.train(vectors)?;
let params = *quantizer
.params()
.expect("quantizer trained successfully implies params are set");
let records: Vec<QuantizedRecord> = vectors
.iter()
.map(|v| quantizer.quantize(v))
.collect::<Result<_>>()?;
Ok((params, records))
}
pub(super) fn write_quantized_record<W: Write>(
output: &mut W,
int8_data: &[u8],
meta: QuantizedVectorMeta,
) -> Result<()> {
output.write_all(int8_data)?;
output.write_all(&meta.sum_q.to_le_bytes())?;
output.write_all(&meta.norm_q.to_le_bytes())?;
Ok(())
}
pub(super) fn read_dequantized_vector<R: Read>(
input: &mut R,
dim: usize,
params: &ScalarQuantParams,
) -> Result<Vec<f32>> {
let mut int8_buf = vec![0u8; dim];
input.read_exact(&mut int8_buf)?;
let mut meta_buf = [0u8; QuantizedVectorMeta::SERIALIZED_SIZE];
input.read_exact(&mut meta_buf)?;
Ok(int8_buf
.iter()
.map(|&b| params.dequantize_value(b))
.collect())
}
#[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_returns_params_and_records() {
let vectors = vec![
vec_of(&[-1.0, 0.0, 1.0]),
vec_of(&[-0.5, 0.5, 0.25]),
vec_of(&[0.1, -0.4, 0.9]),
];
let (params, records) = quantize_segment(&vectors, 3).unwrap();
assert!(params.scale > 0.0);
assert_eq!(records.len(), 3);
for (i, (q, meta)) in records.iter().enumerate() {
assert_eq!(q.len(), 3, "vector {i}");
let expected = QuantizedVectorMeta::from_quantized(q, ¶ms);
assert_eq!(meta.sum_q, expected.sum_q);
assert!((meta.norm_q - expected.norm_q).abs() < 1e-5);
}
}
#[test]
fn write_then_read_dequantized_roundtrips_within_scale() {
let dim = 8;
let vectors = vec![
vec_of(&[-1.0, -0.7, -0.3, 0.0, 0.2, 0.5, 0.8, 1.0]),
vec_of(&[0.1, 0.2, 0.3, 0.4, -0.4, -0.3, -0.2, -0.1]),
];
let (params, records) = quantize_segment(&vectors, dim).unwrap();
let mut buf = Vec::new();
for (q, meta) in &records {
write_quantized_record(&mut buf, q, *meta).unwrap();
}
assert_eq!(
buf.len(),
records.len() * quantized_record_payload_size(dim)
);
let mut cursor = Cursor::new(&buf);
for (i, original) in vectors.iter().enumerate() {
let recovered = read_dequantized_vector(&mut cursor, dim, ¶ms).unwrap();
for (j, (orig, rec)) in original.data.iter().zip(recovered.iter()).enumerate() {
assert!(
(orig - rec).abs() <= params.scale + 1e-6,
"vector {i} dim {j}: orig = {orig}, rec = {rec}, scale = {}",
params.scale
);
}
}
}
#[test]
fn payload_size_is_dim_plus_eight() {
assert_eq!(quantized_record_payload_size(0), 8);
assert_eq!(quantized_record_payload_size(128), 136);
}
}