use super::Compressor;
use super::CV;
use crate::metric::StdMetric;
use crate::vec::VecData;
use crate::CoreNN;
use crate::Mode;
use itertools::Itertools;
use linfa::traits::FitWith;
use linfa::traits::Predict;
use linfa::DatasetBase;
use linfa::Float;
use linfa_clustering::IncrKMeansError;
use linfa_clustering::KMeans;
use linfa_clustering::KMeansInit;
use linfa_nn::distance::L2Dist;
use ndarray::s;
use ndarray::Array2;
use ndarray::ArrayView1;
use ndarray::ArrayView2;
use rand::seq::IteratorRandom;
use rand::thread_rng;
use rayon::iter::IntoParallelIterator;
use rayon::iter::ParallelIterator;
use serde::Deserialize;
use serde::Serialize;
use std::cmp::min;
use std::sync::Arc;
#[derive(Debug, Deserialize, Serialize)]
pub struct ProductQuantizer<T: Float> {
dims: usize,
subspace_codebooks: Vec<KMeans<T, L2Dist>>,
}
impl<T: Float> ProductQuantizer<T> {
pub fn train(mat: &ArrayView2<T>, subspaces: usize) -> Self {
let batch_size = 128;
let dims = mat.shape()[1];
assert_eq!(dims % subspaces, 0);
let subdims = dims / subspaces;
assert!(!mat.iter().any(|x| x.is_nan()));
let subspace_codebooks = (0..subspaces)
.into_par_iter()
.map(|i| {
let submat = mat.slice(s![.., i * subdims..(i + 1) * subdims]);
let obs = DatasetBase::from(submat).shuffle(&mut thread_rng());
let clf = KMeans::params(256).init_method(KMeansInit::KMeansPara);
let mut cur: Option<KMeans<T, L2Dist>> = None;
for batch in obs.sample_chunks(batch_size).cycle() {
match clf.fit_with(cur, &batch) {
Ok(model) => {
cur = Some(model);
break;
}
Err(IncrKMeansError::NotConverged(model)) => {
cur = Some(model);
}
Err(err) => {
panic!("unexpected K-means error: {}", err);
}
};
}
cur.unwrap()
})
.collect::<Vec<_>>();
ProductQuantizer {
dims,
subspace_codebooks,
}
}
pub fn train_from_corenn(corenn: &CoreNN) -> ProductQuantizer<f32> {
let samp_sz = min(corenn.cfg.pq_sample_size, corenn.count.get());
let ids = (0..corenn.count.get()).choose_multiple(&mut thread_rng(), samp_sz);
let Mode::Uncompressed(nodes) = &*corenn.mode.read() else {
unreachable!();
};
let samp_nodes = nodes
.multi_get(&ids)
.into_iter()
.filter_map(|n| n)
.collect_vec();
let actual_samp_sz = samp_nodes.len();
let mut mat = Array2::zeros((actual_samp_sz, corenn.cfg.dim));
for (i, node) in samp_nodes.into_iter().enumerate() {
mat.row_mut(i).assign(&node.vector.to_f32());
}
let pq = ProductQuantizer::train(&mat.view(), corenn.cfg.pq_subspaces);
tracing::info!(
sample_inputs = actual_samp_sz,
subspaces = corenn.cfg.pq_subspaces,
"trained PQ"
);
pq
}
pub fn encode(&self, vec: &ArrayView1<T>) -> Vec<u8> {
assert_eq!(vec.shape()[0], self.dims);
let subspaces = self.subspace_codebooks.len();
let subdims = self.dims / subspaces;
let mut code = vec![0; subspaces];
for (i, codebook) in self.subspace_codebooks.iter().enumerate() {
let subvec = vec.slice(s![i * subdims..(i + 1) * subdims]);
let obs = DatasetBase::new(subvec, ());
let label = codebook.predict(&obs);
code[i] = u8::try_from(label).unwrap();
}
code
}
}
impl Compressor for ProductQuantizer<f32> {
fn into_compressed(&self, v: VecData) -> CV {
let v = v.into_f32();
let view = ArrayView1::from(&v);
Arc::new(self.encode(&view))
}
fn dist(&self, metric: StdMetric, a: &CV, b: &CV) -> f64 {
let a_codes = a.downcast_ref::<Vec<u8>>().unwrap();
let b_codes = b.downcast_ref::<Vec<u8>>().unwrap();
assert_eq!(a_codes.len(), b_codes.len());
let num_subspaces = a_codes.len();
match metric {
StdMetric::L2 => {
let mut total_dist_sq_f64 = 0.0_f64;
for i in 0..num_subspaces {
let codebook = &self.subspace_codebooks[i]; let centroids = codebook.centroids(); let centroid_a_sub = centroids.row(a_codes[i].into()); let centroid_b_sub = centroids.row(b_codes[i].into());
let mut sub_dist_sq_f64 = 0.0_f64;
for k in 0..centroid_a_sub.len() {
let diff_f32 = centroid_a_sub[k] - centroid_b_sub[k];
let diff_f64: f64 = diff_f32.into(); sub_dist_sq_f64 += diff_f64 * diff_f64;
}
total_dist_sq_f64 += sub_dist_sq_f64;
}
total_dist_sq_f64.sqrt()
}
StdMetric::Cosine => {
let mut total_dot_product_f64 = 0.0_f64;
let mut total_norm_sq_a_f64 = 0.0_f64;
let mut total_norm_sq_b_f64 = 0.0_f64;
for i in 0..num_subspaces {
let codebook = &self.subspace_codebooks[i];
let centroids = codebook.centroids();
let centroid_a_sub = centroids.row(a_codes[i] as usize);
let centroid_b_sub = centroids.row(b_codes[i] as usize);
total_dot_product_f64 += (centroid_a_sub.dot(¢roid_b_sub)) as f64;
total_norm_sq_a_f64 += (centroid_a_sub.dot(¢roid_a_sub)) as f64; total_norm_sq_b_f64 += (centroid_b_sub.dot(¢roid_b_sub)) as f64; }
const EPSILON_SQ_NORM: f64 = 1e-12;
if total_norm_sq_a_f64 < EPSILON_SQ_NORM && total_norm_sq_b_f64 < EPSILON_SQ_NORM {
return 0.0; }
if total_norm_sq_a_f64 < EPSILON_SQ_NORM || total_norm_sq_b_f64 < EPSILON_SQ_NORM {
return 1.0;
}
let norm_a_f64 = total_norm_sq_a_f64.sqrt();
let norm_b_f64 = total_norm_sq_b_f64.sqrt();
const EPSILON_NORM: f64 = 1e-6; if norm_a_f64 < EPSILON_NORM || norm_b_f64 < EPSILON_NORM {
return 1.0; }
let cosine_similarity = total_dot_product_f64 / (norm_a_f64 * norm_b_f64);
let clamped_similarity = cosine_similarity.max(-1.0).min(1.0);
1.0 - clamped_similarity
}
}
}
}