use crate::partitioning::kmeans::KMeansEuclidean;
use crate::simd::l2_distance_squared;
use crate::RetrieveError;
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ProductQuantizer {
dimension: usize,
num_codebooks: usize,
codebook_size: usize,
subvector_dim: usize,
codebooks: Vec<f32>,
}
impl ProductQuantizer {
pub fn new(
dimension: usize,
num_codebooks: usize,
codebook_size: usize,
) -> Result<Self, RetrieveError> {
if dimension == 0 || num_codebooks == 0 || codebook_size == 0 {
return Err(RetrieveError::InvalidParameter(
"all parameters must be greater than 0".into(),
));
}
if !dimension.is_multiple_of(num_codebooks) {
return Err(RetrieveError::InvalidParameter(
"dimension must be divisible by num_codebooks".into(),
));
}
Ok(Self {
dimension,
num_codebooks,
codebook_size,
subvector_dim: dimension / num_codebooks,
codebooks: Vec::new(),
})
}
pub fn fit(&mut self, vectors: &[f32], num_vectors: usize) -> Result<(), RetrieveError> {
self.fit_with_seed(vectors, num_vectors, None)
}
pub(crate) fn fit_with_seed(
&mut self,
vectors: &[f32],
num_vectors: usize,
seed: Option<u64>,
) -> Result<(), RetrieveError> {
self.fit_with_seed_and_max_iter(vectors, num_vectors, seed, 100)
}
pub(crate) fn fit_with_seed_and_max_iter(
&mut self,
vectors: &[f32],
num_vectors: usize,
seed: Option<u64>,
max_iter: usize,
) -> Result<(), RetrieveError> {
let mut flat_codebooks =
Vec::with_capacity(self.num_codebooks * self.codebook_size * self.subvector_dim);
let mut actual_codebook_size = self.codebook_size;
for codebook_idx in 0..self.num_codebooks {
let start_dim = codebook_idx * self.subvector_dim;
let end_dim = (codebook_idx + 1) * self.subvector_dim;
let mut flat = Vec::with_capacity(num_vectors * self.subvector_dim);
for i in 0..num_vectors {
let vec = get_vector(vectors, self.dimension, i);
flat.extend_from_slice(&vec[start_dim..end_dim]);
}
let mut kmeans = KMeansEuclidean::new(self.subvector_dim, self.codebook_size)?
.with_max_iter(max_iter);
if let Some(seed) = seed {
kmeans = kmeans.with_seed(seed.wrapping_add(codebook_idx as u64));
}
kmeans.fit(&flat, num_vectors)?;
let centroids = kmeans.centroids();
if codebook_idx == 0 {
actual_codebook_size = centroids.len();
}
for codeword in centroids {
flat_codebooks.extend_from_slice(codeword);
}
}
self.codebook_size = actual_codebook_size;
self.codebooks = flat_codebooks;
Ok(())
}
#[inline]
fn get_codeword(&self, codebook_idx: usize, code: usize) -> &[f32] {
let offset = (codebook_idx * self.codebook_size + code) * self.subvector_dim;
&self.codebooks[offset..offset + self.subvector_dim]
}
pub fn quantize(&self, vector: &[f32]) -> Vec<u8> {
let mut codes = Vec::with_capacity(self.num_codebooks);
for codebook_idx in 0..self.num_codebooks {
let start_dim = codebook_idx * self.subvector_dim;
let end_dim = (codebook_idx + 1) * self.subvector_dim;
let subvector = &vector[start_dim..end_dim];
let mut best_code = 0u8;
let mut best_dist = f32::INFINITY;
for code in 0..self.codebook_size {
let codeword = self.get_codeword(codebook_idx, code);
let dist = l2_distance_squared(subvector, codeword);
if dist < best_dist {
best_dist = dist;
best_code = code.min(255) as u8;
}
}
codes.push(best_code);
}
codes
}
pub fn approximate_distance(&self, query: &[f32], codes: &[u8]) -> f32 {
let mut total_dist = 0.0;
for (codebook_idx, &code) in codes.iter().enumerate() {
let start_dim = codebook_idx * self.subvector_dim;
let end_dim = (codebook_idx + 1) * self.subvector_dim;
let query_subvector = &query[start_dim..end_dim];
let codeword = self.get_codeword(codebook_idx, code as usize);
total_dist += l2_distance_squared(query_subvector, codeword);
}
total_dist
}
pub fn compute_adc_table(&self, query: &[f32]) -> Result<Vec<f32>, RetrieveError> {
let mut table = Vec::with_capacity(self.num_codebooks * self.codebook_size);
self.compute_adc_table_into(query, &mut table)?;
Ok(table)
}
pub fn compute_adc_table_into(
&self,
query: &[f32],
table: &mut Vec<f32>,
) -> Result<(), RetrieveError> {
if query.len() != self.dimension {
return Err(RetrieveError::DimensionMismatch {
query_dim: query.len(),
doc_dim: self.dimension,
});
}
let table_len = self.num_codebooks * self.codebook_size;
table.clear();
if self.subvector_dim == 1 {
table.resize(table_len, 0.0);
for (codebook_idx, (&q, table_chunk)) in query
.iter()
.take(self.num_codebooks)
.zip(table.chunks_exact_mut(self.codebook_size))
.enumerate()
{
let codebook_offset = codebook_idx * self.codebook_size;
let codebook =
&self.codebooks[codebook_offset..codebook_offset + self.codebook_size];
for (&codeword, out) in codebook.iter().zip(table_chunk) {
let diff = q - codeword;
*out = diff * diff;
}
}
return Ok(());
}
table.reserve(table_len);
if self.subvector_dim <= 8 {
for codebook_idx in 0..self.num_codebooks {
let start_dim = codebook_idx * self.subvector_dim;
let end_dim = start_dim + self.subvector_dim;
let query_subvector = &query[start_dim..end_dim];
let codebook_offset = codebook_idx * self.codebook_size * self.subvector_dim;
let codebook_end = codebook_offset + self.codebook_size * self.subvector_dim;
let codebook = &self.codebooks[codebook_offset..codebook_end];
for codeword in codebook.chunks_exact(self.subvector_dim) {
let mut dist = 0.0f32;
for dim in 0..self.subvector_dim {
let diff = query_subvector[dim] - codeword[dim];
dist += diff * diff;
}
table.push(dist);
}
}
return Ok(());
}
for codebook_idx in 0..self.num_codebooks {
let start_dim = codebook_idx * self.subvector_dim;
let end_dim = (codebook_idx + 1) * self.subvector_dim;
let query_subvector = &query[start_dim..end_dim];
for code in 0..self.codebook_size {
let codeword = self.get_codeword(codebook_idx, code);
let dist = l2_distance_squared(query_subvector, codeword);
table.push(dist);
}
}
Ok(())
}
#[inline(always)]
pub fn distance_with_table(&self, table: &[f32], codes: &[u8]) -> f32 {
let mut total_dist = 0.0;
for (codebook_idx, &code) in codes.iter().enumerate() {
let idx = codebook_idx * self.codebook_size + code as usize;
total_dist += table[idx];
}
total_dist
}
pub fn reconstruct(&self, codes: &[u8]) -> Vec<f32> {
let mut result = Vec::with_capacity(self.dimension);
for (m, &code) in codes.iter().enumerate() {
result.extend_from_slice(self.get_codeword(m, code as usize));
}
result
}
pub fn num_codebooks(&self) -> usize {
self.num_codebooks
}
pub fn subvector_dim(&self) -> usize {
self.subvector_dim
}
pub fn codebook_size(&self) -> usize {
self.codebook_size
}
pub(crate) fn owned_bytes(&self) -> usize {
self.codebooks.capacity() * std::mem::size_of::<f32>()
}
}
#[inline]
fn get_vector(vectors: &[f32], dimension: usize, idx: usize) -> &[f32] {
let start = idx * dimension;
let end = start + dimension;
&vectors[start..end]
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn adc_table_specializes_one_dimensional_subvectors() {
let pq = ProductQuantizer {
dimension: 2,
num_codebooks: 2,
codebook_size: 3,
subvector_dim: 1,
codebooks: vec![0.0, 1.0, 3.0, -1.0, 2.0, 4.0],
};
let mut table = Vec::new();
pq.compute_adc_table_into(&[2.0, 1.0], &mut table)
.expect("valid query");
assert_eq!(table, vec![4.0, 1.0, 1.0, 4.0, 1.0, 9.0]);
}
#[test]
fn adc_table_specializes_small_multidimensional_subvectors() {
let pq = ProductQuantizer {
dimension: 6,
num_codebooks: 2,
codebook_size: 2,
subvector_dim: 3,
codebooks: vec![
0.0, 0.0, 0.0, 1.0, 1.0, 1.0, 2.0, 2.0, 2.0, 3.0, 3.0, 3.0,
],
};
let mut table = Vec::new();
pq.compute_adc_table_into(&[1.0, 2.0, 3.0, 4.0, 5.0, 6.0], &mut table)
.expect("valid query");
assert_eq!(table, vec![14.0, 5.0, 29.0, 14.0]);
}
}