use anyhow::{Context, Result};
use serde::{Deserialize, Serialize};
use tracing::{info, warn};
use crate::error::TurboPropError;
const DEFAULT_KMEANS_ITERATIONS: usize = 10;
const DEFAULT_SUBVECTOR_SIZE: usize = 8;
const DEFAULT_CODEBOOK_SIZE: usize = 256;
const DEFAULT_QUANTIZATION_BITS: u8 = 8;
const DEFAULT_CLUSTERING_THRESHOLD: f32 = 0.9;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CompressionConfig {
pub algorithm: CompressionAlgorithm,
pub quantization_bits: u8,
pub enable_delta_compression: bool,
pub clustering_threshold: f32,
pub kmeans_iterations: usize,
pub subvector_size: usize,
pub codebook_size: usize,
}
impl Default for CompressionConfig {
fn default() -> Self {
Self {
algorithm: CompressionAlgorithm::ScalarQuantization,
quantization_bits: DEFAULT_QUANTIZATION_BITS,
enable_delta_compression: true,
clustering_threshold: DEFAULT_CLUSTERING_THRESHOLD,
kmeans_iterations: DEFAULT_KMEANS_ITERATIONS,
subvector_size: DEFAULT_SUBVECTOR_SIZE,
codebook_size: DEFAULT_CODEBOOK_SIZE,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum CompressionAlgorithm {
None,
ScalarQuantization,
ProductQuantization,
LearnedCompression,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CompressedVector {
pub data: Vec<u8>,
pub metadata: CompressionMetadata,
pub original_length: usize,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CompressionMetadata {
pub algorithm: CompressionAlgorithm,
pub quantization_params: QuantizationParams,
pub codebook: Option<Vec<Vec<Vec<f32>>>>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct QuantizationParams {
pub min_value: f32,
pub max_value: f32,
pub num_levels: u32,
pub scale: f32,
}
#[derive(Debug, Clone, Default)]
pub struct CompressionStats {
pub original_size_bytes: usize,
pub compressed_size_bytes: usize,
pub compression_ratio: f64,
pub vectors_processed: usize,
pub avg_compression_time_us: f64,
}
impl CompressionStats {
pub fn compression_percentage(&self) -> f64 {
if self.original_size_bytes == 0 {
0.0
} else {
100.0 * (1.0 - (self.compressed_size_bytes as f64 / self.original_size_bytes as f64))
}
}
}
pub struct VectorCompressor {
config: CompressionConfig,
stats: CompressionStats,
}
impl VectorCompressor {
pub fn new(config: CompressionConfig) -> Self {
Self {
config,
stats: CompressionStats::default(),
}
}
pub fn compress_batch(&mut self, vectors: &[Vec<f32>]) -> Result<Vec<CompressedVector>> {
self.log_compression_start(vectors);
let start_time = std::time::Instant::now();
let compressed = self.apply_compression_algorithm(vectors)?;
let elapsed = start_time.elapsed();
self.update_stats(vectors, &compressed, elapsed);
self.log_compression_complete(elapsed);
Ok(compressed)
}
fn log_compression_start(&self, vectors: &[Vec<f32>]) {
info!(
"Compressing {} vectors using {:?}",
vectors.len(),
self.config.algorithm
);
}
fn apply_compression_algorithm(
&mut self,
vectors: &[Vec<f32>],
) -> Result<Vec<CompressedVector>> {
match self.config.algorithm {
CompressionAlgorithm::None => self.compress_none(vectors),
CompressionAlgorithm::ScalarQuantization => self.compress_scalar_quantization(vectors),
CompressionAlgorithm::ProductQuantization => {
self.compress_product_quantization(vectors)
}
CompressionAlgorithm::LearnedCompression => self.compress_learned(vectors),
}
}
fn log_compression_complete(&self, elapsed: std::time::Duration) {
info!(
"Compression completed: {:.1}% size reduction, {:.2}ms total",
self.stats.compression_percentage(),
elapsed.as_secs_f64() * 1000.0
);
}
pub fn decompress_batch(&self, compressed: &[CompressedVector]) -> Result<Vec<Vec<f32>>> {
info!("Decompressing {} vectors", compressed.len());
let mut decompressed = Vec::with_capacity(compressed.len());
for compressed_vector in compressed {
let vector = self.decompress_single(compressed_vector)?;
decompressed.push(vector);
}
info!("Decompression completed: {} vectors", decompressed.len());
Ok(decompressed)
}
fn compress_none(&mut self, vectors: &[Vec<f32>]) -> Result<Vec<CompressedVector>> {
vectors
.iter()
.map(|vector| {
let data = vector.iter().flat_map(|&f| f.to_le_bytes()).collect();
Ok(CompressedVector {
data,
metadata: CompressionMetadata {
algorithm: CompressionAlgorithm::None,
quantization_params: QuantizationParams {
min_value: 0.0,
max_value: 0.0,
num_levels: 0,
scale: 1.0,
},
codebook: None,
},
original_length: vector.len(),
})
})
.collect()
}
fn compress_scalar_quantization(
&mut self,
vectors: &[Vec<f32>],
) -> Result<Vec<CompressedVector>> {
let (global_min, global_max) = self.calculate_global_range(vectors);
let num_levels = (1u32 << self.config.quantization_bits) - 1;
let scale = (global_max - global_min) / num_levels as f32;
let quantization_params = QuantizationParams {
min_value: global_min,
max_value: global_max,
num_levels,
scale,
};
vectors
.iter()
.map(|vector| {
let quantized_data = vector
.iter()
.map(|&value| {
let normalized = (value - global_min) / scale;
normalized.round().min(num_levels as f32).max(0.0) as u8
})
.collect();
Ok(CompressedVector {
data: quantized_data,
metadata: CompressionMetadata {
algorithm: CompressionAlgorithm::ScalarQuantization,
quantization_params: quantization_params.clone(),
codebook: None,
},
original_length: vector.len(),
})
})
.collect()
}
fn compress_product_quantization(
&mut self,
vectors: &[Vec<f32>],
) -> Result<Vec<CompressedVector>> {
if vectors.is_empty() {
return Ok(Vec::new());
}
let vector_dim = vectors[0].len();
let subvector_size = self.config.subvector_size;
let num_subvectors = vector_dim.div_ceil(subvector_size);
let codebook_size = self.config.codebook_size;
info!(
"Building product quantization codebook: {} subvectors of size {}",
num_subvectors, subvector_size
);
let codebook =
self.build_product_quantization_codebook(vectors, subvector_size, codebook_size)?;
let compressed = vectors
.iter()
.map(|vector| {
let mut codes = Vec::with_capacity(codebook.len());
for (i, centroid) in codebook.iter().enumerate() {
let start = i * subvector_size;
let end = (start + subvector_size).min(vector.len());
let subvector = &vector[start..end];
let best_code = self.find_closest_centroid(subvector, centroid);
codes.push(best_code);
}
CompressedVector {
data: codes,
metadata: CompressionMetadata {
algorithm: CompressionAlgorithm::ProductQuantization,
quantization_params: QuantizationParams {
min_value: 0.0,
max_value: 0.0,
num_levels: codebook_size as u32,
scale: 1.0,
},
codebook: Some(codebook.clone()),
},
original_length: vector.len(),
}
})
.collect::<Vec<_>>();
Ok(compressed)
}
fn compress_learned(&mut self, vectors: &[Vec<f32>]) -> Result<Vec<CompressedVector>> {
warn!("Learned compression not fully implemented, using scalar quantization");
self.compress_scalar_quantization(vectors)
}
fn decompress_single(&self, compressed: &CompressedVector) -> Result<Vec<f32>> {
match compressed.metadata.algorithm {
CompressionAlgorithm::None => self.decompress_none(compressed),
CompressionAlgorithm::ScalarQuantization => {
self.decompress_scalar_quantization(compressed)
}
CompressionAlgorithm::ProductQuantization => {
self.decompress_product_quantization(compressed)
}
CompressionAlgorithm::LearnedCompression => self.decompress_learned(compressed),
}
}
fn decompress_none(&self, compressed: &CompressedVector) -> Result<Vec<f32>> {
let floats = compressed
.data
.chunks_exact(4)
.map(|bytes| f32::from_le_bytes([bytes[0], bytes[1], bytes[2], bytes[3]]))
.collect();
Ok(floats)
}
fn decompress_scalar_quantization(&self, compressed: &CompressedVector) -> Result<Vec<f32>> {
let params = &compressed.metadata.quantization_params;
let decompressed = compressed
.data
.iter()
.map(|&quantized| params.min_value + (quantized as f32 * params.scale))
.collect();
Ok(decompressed)
}
fn decompress_product_quantization(&self, compressed: &CompressedVector) -> Result<Vec<f32>> {
let codebook = compressed
.metadata
.codebook
.as_ref()
.context("Product quantization codebook missing")?;
let mut decompressed = Vec::new();
for (subvector_idx, &code) in compressed.data.iter().enumerate() {
if let Some(subspace_codebook) = codebook.get(subvector_idx) {
if let Some(centroid) = subspace_codebook.get(code as usize) {
decompressed.extend_from_slice(centroid);
} else {
return Err(TurboPropError::other(format!(
"Invalid code {} for subvector {}: code exceeds codebook size",
code, subvector_idx
))
.into());
}
} else {
return Err(TurboPropError::other(format!(
"Missing codebook for subvector {}: compression data is corrupted",
subvector_idx
))
.into());
}
}
decompressed.truncate(compressed.original_length);
Ok(decompressed)
}
fn decompress_learned(&self, compressed: &CompressedVector) -> Result<Vec<f32>> {
self.decompress_scalar_quantization(compressed)
}
fn calculate_global_range(&self, vectors: &[Vec<f32>]) -> (f32, f32) {
let mut global_min = f32::INFINITY;
let mut global_max = f32::NEG_INFINITY;
for vector in vectors {
for &value in vector {
global_min = global_min.min(value);
global_max = global_max.max(value);
}
}
(global_min, global_max)
}
fn build_product_quantization_codebook(
&self,
vectors: &[Vec<f32>],
subvector_size: usize,
codebook_size: usize,
) -> Result<Vec<Vec<Vec<f32>>>> {
let vector_dim = vectors[0].len();
let num_subvectors = vector_dim.div_ceil(subvector_size);
let mut codebook = Vec::with_capacity(num_subvectors);
for i in 0..num_subvectors {
let start = i * subvector_size;
let end = (start + subvector_size).min(vector_dim);
let actual_size = end - start;
let subvectors: Vec<Vec<f32>> = vectors
.iter()
.map(|v| {
let mut subvec = v[start..end].to_vec();
subvec.resize(actual_size, 0.0);
subvec
})
.collect();
let centroids =
self.kmeans_centroids(&subvectors, codebook_size.min(subvectors.len()))?;
codebook.push(centroids);
}
Ok(codebook)
}
fn kmeans_centroids(&self, vectors: &[Vec<f32>], k: usize) -> Result<Vec<Vec<f32>>> {
if vectors.is_empty() || k == 0 {
return Ok(Vec::new());
}
let dim = vectors[0].len();
let mut centroids = Vec::with_capacity(k);
for i in 0..k {
let idx = (i * vectors.len()) / k; centroids.push(vectors[idx].clone());
}
for _iteration in 0..self.config.kmeans_iterations {
let mut new_centroids = vec![vec![0.0; dim]; k];
let mut counts = vec![0; k];
for vector in vectors {
let closest = self.find_closest_centroid(vector, ¢roids) as usize;
for (j, &val) in vector.iter().enumerate() {
new_centroids[closest][j] += val;
}
counts[closest] += 1;
}
for (i, centroid) in new_centroids.iter_mut().enumerate() {
if counts[i] > 0 {
for val in centroid.iter_mut() {
*val /= counts[i] as f32;
}
}
}
centroids = new_centroids;
}
Ok(centroids)
}
fn find_closest_centroid(&self, vector: &[f32], centroids: &[Vec<f32>]) -> u8 {
let mut best_distance_squared = f32::INFINITY;
let mut best_idx = 0;
for (i, centroid) in centroids.iter().enumerate() {
let distance_squared = self.euclidean_distance_squared(vector, centroid);
if distance_squared < best_distance_squared {
best_distance_squared = distance_squared;
best_idx = i;
}
}
best_idx as u8
}
fn euclidean_distance_squared(&self, a: &[f32], b: &[f32]) -> f32 {
debug_assert_eq!(a.len(), b.len(), "Vector dimensions must match");
let mut sum = 0.0f32;
for (&x, &y) in a.iter().zip(b.iter()) {
let diff = x - y;
sum += diff * diff;
}
sum
}
fn update_stats(
&mut self,
original: &[Vec<f32>],
compressed: &[CompressedVector],
elapsed: std::time::Duration,
) {
let original_bytes: usize = original
.iter()
.map(|v| v.len() * std::mem::size_of::<f32>())
.sum();
let compressed_bytes: usize = compressed
.iter()
.map(|c| c.data.len() + std::mem::size_of::<CompressionMetadata>())
.sum();
self.stats.original_size_bytes += original_bytes;
self.stats.compressed_size_bytes += compressed_bytes;
self.stats.vectors_processed += original.len();
if self.stats.original_size_bytes > 0 {
self.stats.compression_ratio =
self.stats.original_size_bytes as f64 / self.stats.compressed_size_bytes as f64;
}
self.stats.avg_compression_time_us = elapsed.as_micros() as f64 / original.len() as f64;
}
pub fn stats(&self) -> &CompressionStats {
&self.stats
}
}
#[cfg(test)]
mod tests {
use super::*;
fn create_test_vectors(count: usize, dimensions: usize) -> Vec<Vec<f32>> {
use rand::prelude::*;
let mut rng = thread_rng();
(0..count)
.map(|_| (0..dimensions).map(|_| rng.gen_range(-1.0..1.0)).collect())
.collect()
}
#[test]
fn test_scalar_quantization() {
let config = CompressionConfig::default();
let mut compressor = VectorCompressor::new(config);
let vectors = create_test_vectors(10, 384);
let compressed = compressor.compress_batch(&vectors).unwrap();
let decompressed = compressor.decompress_batch(&compressed).unwrap();
assert_eq!(vectors.len(), decompressed.len());
assert_eq!(vectors[0].len(), decompressed[0].len());
let stats = compressor.stats();
assert!(stats.compression_ratio > 1.0);
}
#[test]
fn test_no_compression() {
let config = CompressionConfig {
algorithm: CompressionAlgorithm::None,
..Default::default()
};
let mut compressor = VectorCompressor::new(config);
let vectors = create_test_vectors(5, 100);
let compressed = compressor.compress_batch(&vectors).unwrap();
let decompressed = compressor.decompress_batch(&compressed).unwrap();
for (orig, decomp) in vectors.iter().zip(decompressed.iter()) {
assert_eq!(orig.len(), decomp.len());
for (a, b) in orig.iter().zip(decomp.iter()) {
assert!((a - b).abs() < 1e-6);
}
}
}
#[test]
fn test_compression_stats() {
let config = CompressionConfig::default();
let compressor = VectorCompressor::new(config);
let stats = compressor.stats();
assert_eq!(stats.vectors_processed, 0);
assert_eq!(stats.compression_ratio, 0.0);
}
}