use std::collections::HashMap;
#[derive(Debug, Clone)]
pub struct QuantizationParams {
pub min: Vec<f32>,
pub max: Vec<f32>,
pub dim: usize,
}
impl QuantizationParams {
pub fn fit(vectors: &[Vec<f32>]) -> crate::errors::Result<Self> {
if vectors.is_empty() {
return Err(crate::errors::RagError::InvalidConfig(
"Cannot fit quantization on empty vector set".to_string(),
));
}
let dim = vectors[0].len();
let mut min = vec![f32::INFINITY; dim];
let mut max = vec![f32::NEG_INFINITY; dim];
for v in vectors {
if v.len() != dim {
return Err(crate::errors::RagError::InvalidConfig(format!(
"Vector dimension mismatch: expected {dim}, got {}",
v.len()
)));
}
for (i, x) in v.iter().enumerate() {
if *x < min[i] {
min[i] = *x;
}
if *x > max[i] {
max[i] = *x;
}
}
}
for i in 0..dim {
if max[i] <= min[i] {
let eps = 1e-6_f32;
let mid = min[i];
min[i] = mid - eps;
max[i] = mid + eps;
}
}
Ok(Self { min, max, dim })
}
pub fn quantize(&self, vector: &[f32]) -> Vec<i8> {
let mut out = vec![0i8; self.dim];
for (i, slot) in out.iter_mut().enumerate() {
let v = vector.get(i).copied().unwrap_or(self.min[i]);
let range = self.max[i] - self.min[i];
let normalized = ((v - self.min[i]) / range).clamp(0.0, 1.0);
*slot = (normalized * 255.0 - 128.0).round() as i8;
}
out
}
pub fn dequantize(&self, codes: &[i8]) -> Vec<f32> {
let mut out = vec![0.0_f32; self.dim];
for (i, slot) in out.iter_mut().enumerate() {
let c = codes.get(i).copied().unwrap_or(0) as f32;
let range = self.max[i] - self.min[i];
*slot = self.min[i] + ((c + 128.0) / 255.0) * range;
}
out
}
}
#[derive(Debug, Clone)]
pub struct QuantizedIndex {
pub params: QuantizationParams,
codes: HashMap<String, Vec<i8>>,
}
impl QuantizedIndex {
pub fn new(params: QuantizationParams) -> Self {
Self {
params,
codes: HashMap::new(),
}
}
pub fn add(&mut self, id: impl Into<String>, vector: &[f32]) {
let q = self.params.quantize(vector);
self.codes.insert(id.into(), q);
}
pub fn remove(&mut self, id: &str) -> bool {
self.codes.remove(id).is_some()
}
pub fn len(&self) -> usize {
self.codes.len()
}
pub fn is_empty(&self) -> bool {
self.codes.is_empty()
}
pub fn search(&self, query: &[f32], top_k: usize) -> Vec<(String, f32)> {
let q = self.params.quantize(query);
let qd = self.params.dequantize(&q);
let mut scored: Vec<(String, f32)> = self
.codes
.iter()
.map(|(id, codes)| {
let dq = self.params.dequantize(codes);
let dot: f32 = dq.iter().zip(qd.iter()).map(|(a, b)| a * b).sum();
(id.clone(), dot)
})
.collect();
scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
scored.truncate(top_k);
scored
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_fit_and_quantize_roundtrip() {
let vectors = vec![
vec![-1.0, 0.0, 1.0],
vec![0.5, -0.5, 0.5],
vec![2.0, 1.0, -1.0],
];
let params = QuantizationParams::fit(&vectors).unwrap();
let q = params.quantize(&[0.0, 0.0, 0.0]);
assert_eq!(q.len(), 3);
let dq = params.dequantize(&q);
for (orig, back) in [0.0_f32, 0.0, 0.0].iter().zip(dq.iter()) {
assert!((orig - back).abs() < 0.05, "orig={orig} back={back}");
}
}
#[test]
fn test_constant_dimension() {
let vectors = vec![vec![5.0, 1.0], vec![5.0, 2.0]];
let params = QuantizationParams::fit(&vectors).unwrap();
let q = params.quantize(&[5.0, 1.5]);
assert_eq!(q.len(), 2);
}
#[test]
fn test_fit_empty_errors() {
assert!(QuantizationParams::fit(&[]).is_err());
}
#[test]
fn test_fit_dim_mismatch_errors() {
let vectors = vec![vec![1.0, 2.0], vec![1.0]];
assert!(QuantizationParams::fit(&vectors).is_err());
}
#[test]
fn test_quantized_index_search_order() {
let vectors = vec![vec![1.0, 0.0], vec![0.0, 1.0], vec![1.0, 1.0]];
let params = QuantizationParams::fit(&vectors).unwrap();
let mut index = QuantizedIndex::new(params);
index.add("a", &[1.0, 0.0]);
index.add("b", &[0.0, 1.0]);
index.add("c", &[1.0, 1.0]);
let res = index.search(&[1.0, 1.0], 3);
assert_eq!(res.len(), 3);
assert_eq!(res[0].0, "c");
}
#[test]
fn test_quantized_index_remove() {
let vectors = vec![vec![1.0, 0.0], vec![0.0, 1.0]];
let params = QuantizationParams::fit(&vectors).unwrap();
let mut index = QuantizedIndex::new(params);
index.add("a", &[1.0, 0.0]);
assert_eq!(index.len(), 1);
assert!(index.remove("a"));
assert!(index.is_empty());
assert!(!index.remove("a"));
}
#[test]
fn test_quantized_index_empty_search() {
let params = QuantizationParams::fit(&[vec![1.0, 0.0]]).unwrap();
let index = QuantizedIndex::new(params);
let res = index.search(&[1.0, 0.0], 5);
assert!(res.is_empty());
}
}