use std::io::{Read, Write};
use crate::error::{LaurusError, 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_fastscan_record_payload_size(m: u16) -> usize {
(m as usize).div_ceil(2)
}
pub(super) fn quantize_segment_pq_fastscan(
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, 16)?;
VectorQuantizer::from_pq_codebook(dim, cb.params, cb.codebook.clone())?
}
None => {
let mut quantizer = VectorQuantizer::new(
QuantizationMethod::ProductQuantizationFastScan { subvector_count },
dim,
);
quantizer.train(vectors)?;
quantizer
}
};
let (params, codebook_slice) = quantizer
.pq_state()
.expect("PQ FastScan quantizer trained successfully implies pq_state is set");
let params = *params;
if params.k != 16 {
return Err(LaurusError::InvalidOperation(format!(
"PQ FastScan must use k == 16, got k = {}",
params.k
)));
}
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<_>>()?;
for (i, c) in codes.iter().enumerate() {
for (j, &code) in c.iter().enumerate() {
if code >= 16 {
return Err(LaurusError::InvalidOperation(format!(
"FastScan code {code} (vector {i}, sub {j}) exceeds 4-bit range"
)));
}
}
}
Ok((params, codebook, codes))
}
#[inline]
fn pack_one_record(codes: &[u8]) -> Vec<u8> {
let m = codes.len();
let mut packed = vec![0u8; m.div_ceil(2)];
for (i, &c) in codes.iter().enumerate() {
let nibble = c & 0x0F;
if i % 2 == 0 {
packed[i / 2] |= nibble;
} else {
packed[i / 2] |= nibble << 4;
}
}
packed
}
#[inline]
fn unpack_one_record(packed: &[u8], m: usize) -> Vec<u8> {
let mut codes = vec![0u8; m];
for (i, slot) in codes.iter_mut().enumerate().take(m) {
let byte = packed[i / 2];
*slot = if i % 2 == 0 {
byte & 0x0F
} else {
(byte >> 4) & 0x0F
};
}
codes
}
pub(super) fn write_pq_fastscan_record<W: Write>(output: &mut W, codes: &[u8]) -> Result<()> {
let packed = pack_one_record(codes);
output.write_all(&packed)?;
Ok(())
}
pub(super) fn read_pq_fastscan_record<R: Read>(input: &mut R, params: PqParams) -> Result<Vec<u8>> {
let m = params.m as usize;
let packed_len = m.div_ceil(2);
let mut packed = vec![0u8; packed_len];
input.read_exact(&mut packed)?;
Ok(unpack_one_record(&packed, m))
}
pub(super) fn read_dequantized_pq_fastscan_vector<R: Read>(
input: &mut R,
params: PqParams,
codebook: &[f32],
) -> Result<Vec<f32>> {
let codes = read_pq_fastscan_record(input, params)?;
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 payload_size_is_ceil_m_over_two() {
assert_eq!(pq_fastscan_record_payload_size(0), 0);
assert_eq!(pq_fastscan_record_payload_size(1), 1);
assert_eq!(pq_fastscan_record_payload_size(2), 1);
assert_eq!(pq_fastscan_record_payload_size(3), 2);
assert_eq!(pq_fastscan_record_payload_size(8), 4);
assert_eq!(pq_fastscan_record_payload_size(16), 8);
}
#[test]
fn pack_unpack_round_trip_even_m() {
let codes = vec![0, 1, 2, 15, 8, 7, 4, 3];
let packed = pack_one_record(&codes);
assert_eq!(packed.len(), 4);
let unpacked = unpack_one_record(&packed, codes.len());
assert_eq!(unpacked, codes);
}
#[test]
fn pack_unpack_round_trip_odd_m() {
let codes = vec![5, 10, 15];
let packed = pack_one_record(&codes);
assert_eq!(packed.len(), 2);
let unpacked = unpack_one_record(&packed, codes.len());
assert_eq!(unpacked, codes);
}
#[test]
fn quantize_segment_pq_fastscan_returns_k_16_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_fastscan(&vectors, 4, 2, None).unwrap();
assert_eq!(params.m, 2);
assert_eq!(params.k, 16);
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);
for &code in c {
assert!(code < 16, "code {code} exceeds 4-bit range");
}
}
}
#[test]
fn shared_codebook_produces_byte_identical_fastscan_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_fastscan(&vectors, 32, 4, None).unwrap();
assert_eq!(trained_params.k, 16);
let shared = SharedPqCodebook {
params: trained_params,
codebook: trained_codebook.clone(),
};
let (shared_params, shared_codebook, shared_codes) =
quantize_segment_pq_fastscan(&vectors, 32, 4, Some(&shared)).unwrap();
assert_eq!(shared_params, trained_params);
assert_eq!(shared_codebook, trained_codebook);
assert_eq!(shared_codes, trained_codes);
}
#[test]
fn k256_codebook_is_rejected_by_the_fastscan_encode_path() {
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 = PqParams::new(2, 256, 2).unwrap();
let shared = SharedPqCodebook {
codebook: vec![0.0; params.codebook_len()],
params,
};
let err = quantize_segment_pq_fastscan(&vectors, 4, 2, Some(&shared)).unwrap_err();
assert!(
err.to_string().contains("k = 256"),
"the k mismatch must be named: {err}"
);
}
#[test]
fn write_then_read_pq_fastscan_round_trips_codes() {
let codes: Vec<Vec<u8>> = vec![vec![0, 1, 2, 3], vec![15, 14, 13, 12], vec![7, 8, 9, 10]];
let mut buf = Vec::new();
for c in &codes {
write_pq_fastscan_record(&mut buf, c).unwrap();
}
assert_eq!(buf.len(), codes.len() * pq_fastscan_record_payload_size(4));
let params = PqParams::new(4, 16, 1).unwrap();
let mut cursor = Cursor::new(&buf);
for original in &codes {
let recovered = read_pq_fastscan_record(&mut cursor, params).unwrap();
assert_eq!(&recovered, original);
}
}
}