use anyhow::{anyhow, Result};
use rand::seq::SliceRandom;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
pub const NUM_CENTROIDS: usize = 256;
pub const DEFAULT_SUBVEC_DIM: usize = 8;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PQConfig {
pub dimension: usize,
pub num_subvectors: usize,
pub subvec_dim: usize,
pub num_centroids: usize,
pub kmeans_iterations: usize,
}
impl PQConfig {
pub fn for_dimension(dimension: usize) -> Self {
let subvec_dim = DEFAULT_SUBVEC_DIM;
let num_subvectors = dimension / subvec_dim;
assert!(
dimension % subvec_dim == 0,
"Dimension {} must be divisible by subvec_dim {}",
dimension,
subvec_dim
);
Self {
dimension,
num_subvectors,
subvec_dim,
num_centroids: NUM_CENTROIDS,
kmeans_iterations: 20,
}
}
pub fn minilm() -> Self {
Self::for_dimension(384)
}
pub fn clip() -> Self {
Self::for_dimension(768)
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ProductQuantizer {
pub config: PQConfig,
pub centroids: Vec<Vec<Vec<f32>>>,
pub trained: bool,
}
impl ProductQuantizer {
pub fn new(config: PQConfig) -> Self {
Self {
config,
centroids: Vec::new(),
trained: false,
}
}
pub fn train(config: PQConfig, training_vectors: &[Vec<f32>]) -> Result<Self> {
if training_vectors.is_empty() {
return Err(anyhow!("No training vectors provided"));
}
let first_dim = training_vectors[0].len();
if first_dim != config.dimension {
return Err(anyhow!(
"Training vector dimension {} doesn't match config {}",
first_dim,
config.dimension
));
}
let mut pq = Self::new(config);
pq.fit(training_vectors)?;
Ok(pq)
}
fn fit(&mut self, vectors: &[Vec<f32>]) -> Result<()> {
let n_vectors = vectors.len();
let n_subvectors = self.config.num_subvectors;
let subvec_dim = self.config.subvec_dim;
let n_centroids = self.config.num_centroids.min(n_vectors);
let iterations = self.config.kmeans_iterations;
tracing::info!(
"Training PQ: {} vectors, {} subvectors, {} centroids, {} iterations",
n_vectors,
n_subvectors,
n_centroids,
iterations
);
self.centroids = Vec::with_capacity(n_subvectors);
for subvec_idx in 0..n_subvectors {
let start = subvec_idx * subvec_dim;
let end = start + subvec_dim;
let subvectors: Vec<Vec<f32>> =
vectors.iter().map(|v| v[start..end].to_vec()).collect();
let centroids = self.kmeans(&subvectors, n_centroids, iterations)?;
self.centroids.push(centroids);
}
self.trained = true;
tracing::info!("PQ training complete");
Ok(())
}
fn kmeans(&self, vectors: &[Vec<f32>], k: usize, iterations: usize) -> Result<Vec<Vec<f32>>> {
let dim = vectors[0].len();
let n = vectors.len();
let mut rng = rand::thread_rng();
let mut indices: Vec<usize> = (0..n).collect();
indices.shuffle(&mut rng);
let mut centroids: Vec<Vec<f32>> = indices
.iter()
.take(k)
.map(|&i| vectors[i].clone())
.collect();
while centroids.len() < k {
let idx = indices[centroids.len() % n];
centroids.push(vectors[idx].clone());
}
let mut assignments = vec![0usize; n];
for _ in 0..iterations {
for (i, vec) in vectors.iter().enumerate() {
let mut best_centroid = 0;
let mut best_dist = f32::MAX;
for (c, centroid) in centroids.iter().enumerate() {
let dist = squared_l2_distance(vec, centroid);
if dist < best_dist {
best_dist = dist;
best_centroid = c;
}
}
assignments[i] = best_centroid;
}
let mut new_centroids: Vec<Vec<f32>> = vec![vec![0.0; dim]; k];
let mut counts = vec![0usize; k];
for (i, vec) in vectors.iter().enumerate() {
let c = assignments[i];
counts[c] += 1;
for (j, &v) in vec.iter().enumerate() {
new_centroids[c][j] += v;
}
}
for c in 0..k {
if counts[c] > 0 {
for j in 0..dim {
new_centroids[c][j] /= counts[c] as f32;
}
centroids[c] = new_centroids[c].clone();
}
}
}
Ok(centroids)
}
pub fn encode(&self, vector: &[f32]) -> Result<Vec<u8>> {
if !self.trained {
return Err(anyhow!("ProductQuantizer not trained"));
}
if vector.len() != self.config.dimension {
return Err(anyhow!(
"Vector dimension {} doesn't match config {}",
vector.len(),
self.config.dimension
));
}
let mut codes = Vec::with_capacity(self.config.num_subvectors);
let subvec_dim = self.config.subvec_dim;
for (subvec_idx, subspace_centroids) in self.centroids.iter().enumerate() {
let start = subvec_idx * subvec_dim;
let end = start + subvec_dim;
let subvector = &vector[start..end];
let mut best_centroid = 0u8;
let mut best_dist = f32::MAX;
for (c, centroid) in subspace_centroids.iter().enumerate() {
let dist = squared_l2_distance_slice(subvector, centroid);
if dist < best_dist {
best_dist = dist;
best_centroid = c as u8;
}
}
codes.push(best_centroid);
}
Ok(codes)
}
pub fn decode(&self, codes: &[u8]) -> Result<Vec<f32>> {
if !self.trained {
return Err(anyhow!("ProductQuantizer not trained"));
}
if codes.len() != self.config.num_subvectors {
return Err(anyhow!(
"Code length {} doesn't match num_subvectors {}",
codes.len(),
self.config.num_subvectors
));
}
let mut vector = Vec::with_capacity(self.config.dimension);
for (subvec_idx, &code) in codes.iter().enumerate() {
let subspace = &self.centroids[subvec_idx];
let code_idx = code as usize;
if code_idx >= subspace.len() {
return Err(anyhow!(
"PQ code {} out of bounds for subspace {} with {} centroids (data corruption?)",
code_idx,
subvec_idx,
subspace.len()
));
}
vector.extend_from_slice(&subspace[code_idx]);
}
Ok(vector)
}
pub fn asymmetric_distance(&self, query: &[f32], codes: &[u8]) -> Result<f32> {
if !self.trained {
return Err(anyhow!("ProductQuantizer not trained"));
}
let subvec_dim = self.config.subvec_dim;
let mut total_dist = 0.0f32;
for (subvec_idx, &code) in codes.iter().enumerate() {
let start = subvec_idx * subvec_dim;
let end = start + subvec_dim;
let query_subvec = &query[start..end];
let subspace = &self.centroids[subvec_idx];
let code_idx = code as usize;
if code_idx >= subspace.len() {
return Err(anyhow!(
"PQ code {} out of bounds for subspace {} with {} centroids (data corruption?)",
code_idx,
subvec_idx,
subspace.len()
));
}
total_dist += squared_l2_distance_slice(query_subvec, &subspace[code_idx]);
}
Ok(total_dist)
}
pub fn build_distance_table(&self, query: &[f32]) -> Result<Vec<Vec<f32>>> {
if !self.trained {
return Err(anyhow!("ProductQuantizer not trained"));
}
let subvec_dim = self.config.subvec_dim;
let n_centroids = self.config.num_centroids;
let mut table = Vec::with_capacity(self.config.num_subvectors);
for (subvec_idx, subspace_centroids) in self.centroids.iter().enumerate() {
let start = subvec_idx * subvec_dim;
let end = start + subvec_dim;
let query_subvec = &query[start..end];
let mut distances = Vec::with_capacity(n_centroids);
for centroid in subspace_centroids {
distances.push(squared_l2_distance_slice(query_subvec, centroid));
}
table.push(distances);
}
Ok(table)
}
#[inline]
pub fn distance_with_table(&self, table: &[Vec<f32>], codes: &[u8]) -> f32 {
let mut total = 0.0f32;
for (subvec_idx, &code) in codes.iter().enumerate() {
let code_idx = code as usize;
if subvec_idx >= table.len() || code_idx >= table[subvec_idx].len() {
return f32::MAX; }
total += table[subvec_idx][code_idx];
}
total
}
pub fn encode_batch(&self, vectors: &[Vec<f32>]) -> Result<Vec<Vec<u8>>> {
vectors.iter().map(|v| self.encode(v)).collect()
}
pub fn compressed_size(&self) -> usize {
self.config.num_subvectors }
pub fn original_size(&self) -> usize {
self.config.dimension * std::mem::size_of::<f32>()
}
pub fn compression_ratio(&self) -> f32 {
self.original_size() as f32 / self.compressed_size() as f32
}
}
#[inline]
fn squared_l2_distance(a: &[f32], b: &[f32]) -> f32 {
a.iter().zip(b.iter()).map(|(x, y)| (x - y).powi(2)).sum()
}
#[inline]
fn squared_l2_distance_slice(a: &[f32], b: &[f32]) -> f32 {
squared_l2_distance(a, b)
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CompressedVectorStore {
pub quantizer: ProductQuantizer,
pub codes: HashMap<u32, Vec<u8>>,
}
impl CompressedVectorStore {
pub fn new(quantizer: ProductQuantizer) -> Self {
Self {
quantizer,
codes: HashMap::new(),
}
}
pub fn train_and_create(config: PQConfig, training_vectors: &[Vec<f32>]) -> Result<Self> {
let quantizer = ProductQuantizer::train(config, training_vectors)?;
Ok(Self::new(quantizer))
}
pub fn add(&mut self, vector_id: u32, vector: &[f32]) -> Result<()> {
let codes = self.quantizer.encode(vector)?;
self.codes.insert(vector_id, codes);
Ok(())
}
pub fn get_codes(&self, vector_id: u32) -> Option<&Vec<u8>> {
self.codes.get(&vector_id)
}
pub fn decode(&self, vector_id: u32) -> Result<Vec<f32>> {
let codes = self
.codes
.get(&vector_id)
.ok_or_else(|| anyhow!("Vector {} not found", vector_id))?;
self.quantizer.decode(codes)
}
pub fn search(&self, query: &[f32], k: usize) -> Result<Vec<(u32, f32)>> {
let table = self.quantizer.build_distance_table(query)?;
let mut distances: Vec<(u32, f32)> = self
.codes
.iter()
.map(|(&id, codes)| (id, self.quantizer.distance_with_table(&table, codes)))
.collect();
distances.sort_by(|a, b| a.1.total_cmp(&b.1));
distances.truncate(k);
Ok(distances)
}
pub fn len(&self) -> usize {
self.codes.len()
}
pub fn is_empty(&self) -> bool {
self.codes.is_empty()
}
pub fn storage_bytes(&self) -> usize {
self.codes.len() * self.quantizer.compressed_size()
}
}
#[cfg(test)]
mod tests {
use super::*;
fn generate_random_vectors(n: usize, dim: usize) -> Vec<Vec<f32>> {
use rand::Rng;
let mut rng = rand::thread_rng();
(0..n)
.map(|_| (0..dim).map(|_| rng.gen::<f32>()).collect())
.collect()
}
#[test]
fn test_pq_encode_decode() {
let vectors = generate_random_vectors(1000, 384);
let config = PQConfig::minilm();
let pq = ProductQuantizer::train(config, &vectors).unwrap();
let original = &vectors[0];
let codes = pq.encode(original).unwrap();
let decoded = pq.decode(&codes).unwrap();
assert_eq!(codes.len(), 48); assert_eq!(decoded.len(), 384);
let mse: f32 = original
.iter()
.zip(decoded.iter())
.map(|(a, b)| (a - b).powi(2))
.sum::<f32>()
/ 384.0;
assert!(mse < 0.1, "MSE too high: {}", mse);
}
#[test]
fn test_compression_ratio() {
let config = PQConfig::minilm();
let pq = ProductQuantizer::new(config);
assert_eq!(pq.original_size(), 384 * 4); assert_eq!(pq.compressed_size(), 48); assert!((pq.compression_ratio() - 32.0).abs() < 0.01);
}
#[test]
fn test_distance_table() {
let vectors = generate_random_vectors(100, 384);
let config = PQConfig::minilm();
let pq = ProductQuantizer::train(config, &vectors).unwrap();
let query = &vectors[0];
let codes = pq.encode(&vectors[1]).unwrap();
let direct_dist = pq.asymmetric_distance(query, &codes).unwrap();
let table = pq.build_distance_table(query).unwrap();
let table_dist = pq.distance_with_table(&table, &codes);
assert!((direct_dist - table_dist).abs() < 1e-6);
}
#[test]
fn test_compressed_store_search() {
let vectors = generate_random_vectors(1000, 384);
let config = PQConfig::minilm();
let mut store = CompressedVectorStore::train_and_create(config, &vectors).unwrap();
for (i, v) in vectors.iter().enumerate() {
store.add(i as u32, v).unwrap();
}
let results = store.search(&vectors[0], 10).unwrap();
assert_eq!(results.len(), 10);
let query_in_top_results = results.iter().take(5).any(|(id, _)| *id == 0);
assert!(
query_in_top_results,
"Query vector not found in top 5 results"
);
}
}