use wide::i32x8;
use crate::vector::core::distance::DistanceMetric;
use crate::vector::core::quantization::{PqParams, QuantizedVectorMeta, ScalarQuantParams};
pub const SIMD_BLOCK: usize = 32;
#[inline]
pub const fn padded_dim(dim: usize) -> usize {
(dim + SIMD_BLOCK - 1) & !(SIMD_BLOCK - 1)
}
#[derive(Debug, Clone)]
pub struct QuantizedQuery {
pub q_data: Vec<u8>,
pub sum_q: u32,
pub norm_query: f32,
pub scale: f32,
pub offset: f32,
pub dim: usize,
pub pad_dim: usize,
}
impl QuantizedQuery {
pub fn prepare(query: &[f32], params: &ScalarQuantParams) -> Self {
let dim = query.len();
let mut q_data = params.quantize_slice(query);
let sum_q: u32 = q_data.iter().map(|&x| x as u32).sum();
let norm_query = q_data
.iter()
.map(|&q| {
let dq = params.dequantize_value(q);
dq * dq
})
.sum::<f32>()
.sqrt();
let pad_dim = padded_dim(dim);
q_data.resize(pad_dim, 0);
Self {
q_data,
sum_q,
norm_query,
scale: params.scale,
offset: params.offset,
dim,
pad_dim,
}
}
}
pub fn distance_quantized(
metric: DistanceMetric,
query: &QuantizedQuery,
cand: &[u8],
cand_meta: QuantizedVectorMeta,
) -> f32 {
debug_assert_eq!(cand.len(), query.pad_dim);
match metric {
DistanceMetric::Cosine => {
let approx_dot = approx_dot_product(query, cand, cand_meta.sum_q);
let denom = cand_meta.norm_q * query.norm_query;
if denom == 0.0 {
1.0
} else {
let cosine = (approx_dot / denom).clamp(-1.0, 1.0);
1.0 - cosine
}
}
DistanceMetric::Angular => {
let approx_dot = approx_dot_product(query, cand, cand_meta.sum_q);
let denom = cand_meta.norm_q * query.norm_query;
if denom == 0.0 {
std::f32::consts::PI
} else {
let cosine = (approx_dot / denom).clamp(-1.0, 1.0);
cosine.acos()
}
}
DistanceMetric::Euclidean => {
let sq_diff = sq_diff_u8_to_i32(&query.q_data, cand);
(query.scale * query.scale * sq_diff as f32).sqrt()
}
DistanceMetric::Manhattan => {
let abs_diff = abs_diff_u8_to_i32(&query.q_data, cand);
query.scale * abs_diff as f32
}
DistanceMetric::DotProduct => -approx_dot_product(query, cand, cand_meta.sum_q),
}
}
#[inline]
fn approx_dot_product(query: &QuantizedQuery, cand: &[u8], cand_sum_q: u32) -> f32 {
let dot_q = dot_u8_to_i32(&query.q_data, cand);
let n = query.dim as f32;
let off = query.offset;
let scale = query.scale;
n * off * off
+ scale * off * (query.sum_q as f32 + cand_sum_q as f32)
+ scale * scale * dot_q as f32
}
#[inline]
pub fn dot_u8_to_i32(a: &[u8], b: &[u8]) -> i32 {
debug_assert_eq!(a.len(), b.len());
#[cfg(target_arch = "x86_64")]
if crate::vector::core::sq_int8_avx2::is_avx2_supported() {
return unsafe { crate::vector::core::sq_int8_avx2::dot_u8_to_i32_avx2(a, b) };
}
#[cfg(target_arch = "aarch64")]
if crate::vector::core::sq_int8_neon::is_neon_supported() {
return unsafe { crate::vector::core::sq_int8_neon::dot_u8_to_i32_neon(a, b) };
}
dot_u8_to_i32_scalar(a, b)
}
#[doc(hidden)]
pub fn dot_u8_to_i32_scalar(a: &[u8], b: &[u8]) -> i32 {
debug_assert_eq!(a.len(), b.len());
let mut acc = i32x8::ZERO;
let chunks_a = a.chunks_exact(8);
let chunks_b = b.chunks_exact(8);
let rem_a = chunks_a.remainder();
let rem_b = chunks_b.remainder();
for (ca, cb) in chunks_a.zip(chunks_b) {
let va = i32x8::from([
ca[0] as i32,
ca[1] as i32,
ca[2] as i32,
ca[3] as i32,
ca[4] as i32,
ca[5] as i32,
ca[6] as i32,
ca[7] as i32,
]);
let vb = i32x8::from([
cb[0] as i32,
cb[1] as i32,
cb[2] as i32,
cb[3] as i32,
cb[4] as i32,
cb[5] as i32,
cb[6] as i32,
cb[7] as i32,
]);
acc += va * vb;
}
let mut total: i32 = acc.reduce_add();
for (x, y) in rem_a.iter().zip(rem_b.iter()) {
total += (*x as i32) * (*y as i32);
}
total
}
#[inline]
pub fn sq_diff_u8_to_i32(a: &[u8], b: &[u8]) -> i32 {
debug_assert_eq!(a.len(), b.len());
#[cfg(target_arch = "x86_64")]
if crate::vector::core::sq_int8_avx2::is_avx2_supported() {
return unsafe { crate::vector::core::sq_int8_avx2::sq_diff_u8_to_i32_avx2(a, b) };
}
#[cfg(target_arch = "aarch64")]
if crate::vector::core::sq_int8_neon::is_neon_supported() {
return unsafe { crate::vector::core::sq_int8_neon::sq_diff_u8_to_i32_neon(a, b) };
}
sq_diff_u8_to_i32_scalar(a, b)
}
#[doc(hidden)]
pub fn sq_diff_u8_to_i32_scalar(a: &[u8], b: &[u8]) -> i32 {
debug_assert_eq!(a.len(), b.len());
let mut acc = i32x8::ZERO;
let chunks_a = a.chunks_exact(8);
let chunks_b = b.chunks_exact(8);
let rem_a = chunks_a.remainder();
let rem_b = chunks_b.remainder();
for (ca, cb) in chunks_a.zip(chunks_b) {
let va = i32x8::from([
ca[0] as i32,
ca[1] as i32,
ca[2] as i32,
ca[3] as i32,
ca[4] as i32,
ca[5] as i32,
ca[6] as i32,
ca[7] as i32,
]);
let vb = i32x8::from([
cb[0] as i32,
cb[1] as i32,
cb[2] as i32,
cb[3] as i32,
cb[4] as i32,
cb[5] as i32,
cb[6] as i32,
cb[7] as i32,
]);
let diff = va - vb;
acc += diff * diff;
}
let mut total: i32 = acc.reduce_add();
for (x, y) in rem_a.iter().zip(rem_b.iter()) {
let d = (*x as i32) - (*y as i32);
total += d * d;
}
total
}
#[inline]
pub fn abs_diff_u8_to_i32(a: &[u8], b: &[u8]) -> i32 {
debug_assert_eq!(a.len(), b.len());
#[cfg(target_arch = "x86_64")]
if crate::vector::core::sq_int8_avx2::is_avx2_supported() {
return unsafe { crate::vector::core::sq_int8_avx2::abs_diff_u8_to_i32_avx2(a, b) };
}
#[cfg(target_arch = "aarch64")]
if crate::vector::core::sq_int8_neon::is_neon_supported() {
return unsafe { crate::vector::core::sq_int8_neon::abs_diff_u8_to_i32_neon(a, b) };
}
abs_diff_u8_to_i32_scalar(a, b)
}
#[doc(hidden)]
pub fn abs_diff_u8_to_i32_scalar(a: &[u8], b: &[u8]) -> i32 {
debug_assert_eq!(a.len(), b.len());
let mut acc = i32x8::ZERO;
let chunks_a = a.chunks_exact(8);
let chunks_b = b.chunks_exact(8);
let rem_a = chunks_a.remainder();
let rem_b = chunks_b.remainder();
for (ca, cb) in chunks_a.zip(chunks_b) {
let va = i32x8::from([
ca[0] as i32,
ca[1] as i32,
ca[2] as i32,
ca[3] as i32,
ca[4] as i32,
ca[5] as i32,
ca[6] as i32,
ca[7] as i32,
]);
let vb = i32x8::from([
cb[0] as i32,
cb[1] as i32,
cb[2] as i32,
cb[3] as i32,
cb[4] as i32,
cb[5] as i32,
cb[6] as i32,
cb[7] as i32,
]);
let diff = va - vb;
acc += diff.abs();
}
let mut total: i32 = acc.reduce_add();
for (x, y) in rem_a.iter().zip(rem_b.iter()) {
total += ((*x as i32) - (*y as i32)).abs();
}
total
}
#[derive(Debug, Clone)]
pub struct PqQuery {
pub lut: Vec<f32>,
pub query_norm_sq: f32,
pub params: PqParams,
}
impl PqQuery {
pub fn prepare(query: &[f32], params: PqParams, codebook: &[f32]) -> Self {
debug_assert_eq!(query.len(), params.original_dim());
debug_assert_eq!(codebook.len(), params.codebook_len());
let m = params.m as usize;
let k = params.k as usize;
let sub_dim = params.sub_dim as usize;
let mut lut = vec![0.0_f32; m * k];
for sub in 0..m {
let q_sub = &query[sub * sub_dim..(sub + 1) * sub_dim];
let cb_base = sub * k * sub_dim;
for ki in 0..k {
let c = &codebook[cb_base + ki * sub_dim..cb_base + (ki + 1) * sub_dim];
let mut acc = 0.0_f32;
for d in 0..sub_dim {
let diff = q_sub[d] - c[d];
acc += diff * diff;
}
lut[sub * k + ki] = acc;
}
}
let query_norm_sq: f32 = query.iter().map(|x| x * x).sum();
Self {
lut,
query_norm_sq,
params,
}
}
}
pub fn distance_pq_adc(metric: DistanceMetric, query: &PqQuery, codes: &[u8]) -> f32 {
debug_assert_eq!(codes.len(), query.params.m as usize);
let m = query.params.m as usize;
let k = query.params.k as usize;
let mut l2_sq = 0.0_f32;
for (sub, &code) in codes.iter().enumerate().take(m) {
l2_sq += query.lut[sub * k + code as usize];
}
match metric {
DistanceMetric::Euclidean => l2_sq.max(0.0).sqrt(),
DistanceMetric::Cosine => {
(l2_sq * 0.5).clamp(0.0, 2.0)
}
_ => f32::INFINITY,
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::vector::core::vector::Vector;
fn scalar_dot(a: &[u8], b: &[u8]) -> i32 {
a.iter()
.zip(b.iter())
.map(|(&x, &y)| (x as i32) * (y as i32))
.sum()
}
fn scalar_sq_diff(a: &[u8], b: &[u8]) -> i32 {
a.iter()
.zip(b.iter())
.map(|(&x, &y)| {
let d = (x as i32) - (y as i32);
d * d
})
.sum()
}
fn scalar_abs_diff(a: &[u8], b: &[u8]) -> i32 {
a.iter()
.zip(b.iter())
.map(|(&x, &y)| ((x as i32) - (y as i32)).abs())
.sum()
}
fn pseudo_random_u8(seed: u32, len: usize) -> Vec<u8> {
let mut state = seed.wrapping_mul(0x9E37_79B9).wrapping_add(0xDEAD_BEEF);
(0..len)
.map(|_| {
state = state.wrapping_mul(1103515245).wrapping_add(12345);
(state >> 16) as u8
})
.collect()
}
fn pseudo_random_f32(seed: u32, len: usize, lo: f32, hi: f32) -> Vec<f32> {
let bytes = pseudo_random_u8(seed, len);
let range = hi - lo;
bytes
.into_iter()
.map(|b| lo + (b as f32 / 255.0) * range)
.collect()
}
#[test]
fn dot_simd_matches_scalar_for_various_lengths() {
for &dim in &[
1, 7, 8, 9, 16, 17, 31, 32, 33, 64, 65, 96, 100, 128, 384, 768,
] {
let a = pseudo_random_u8(1, dim);
let b = pseudo_random_u8(2, dim);
assert_eq!(
dot_u8_to_i32(&a, &b),
scalar_dot(&a, &b),
"dim = {dim}: SIMD dot disagrees with scalar"
);
}
}
#[test]
fn sq_diff_simd_matches_scalar_for_various_lengths() {
for &dim in &[
1, 7, 8, 9, 16, 17, 31, 32, 33, 64, 65, 96, 100, 128, 384, 768,
] {
let a = pseudo_random_u8(3, dim);
let b = pseudo_random_u8(4, dim);
assert_eq!(
sq_diff_u8_to_i32(&a, &b),
scalar_sq_diff(&a, &b),
"dim = {dim}: SIMD sq_diff disagrees with scalar"
);
}
}
#[test]
fn abs_diff_simd_matches_scalar_for_various_lengths() {
for &dim in &[
1, 7, 8, 9, 16, 17, 31, 32, 33, 64, 65, 96, 100, 128, 384, 768,
] {
let a = pseudo_random_u8(5, dim);
let b = pseudo_random_u8(6, dim);
assert_eq!(
abs_diff_u8_to_i32(&a, &b),
scalar_abs_diff(&a, &b),
"dim = {dim}: SIMD abs_diff disagrees with scalar"
);
}
}
#[test]
fn scalar_fallbacks_match_reference_for_various_lengths() {
for &dim in &[
1, 7, 8, 9, 16, 17, 31, 32, 33, 64, 65, 96, 100, 128, 384, 768,
] {
let a = pseudo_random_u8(7, dim);
let b = pseudo_random_u8(8, dim);
assert_eq!(
dot_u8_to_i32_scalar(&a, &b),
scalar_dot(&a, &b),
"dot dim={dim}"
);
assert_eq!(
sq_diff_u8_to_i32_scalar(&a, &b),
scalar_sq_diff(&a, &b),
"sq_diff dim={dim}"
);
assert_eq!(
abs_diff_u8_to_i32_scalar(&a, &b),
scalar_abs_diff(&a, &b),
"abs_diff dim={dim}"
);
}
}
#[cfg(target_arch = "x86_64")]
#[test]
fn avx2_kernels_match_scalar_when_supported() {
use crate::vector::core::sq_int8_avx2::{
abs_diff_u8_to_i32_avx2, dot_u8_to_i32_avx2, is_avx2_supported, sq_diff_u8_to_i32_avx2,
};
if !is_avx2_supported() {
return;
}
for &dim in &[
1, 7, 8, 9, 16, 17, 31, 32, 33, 64, 65, 96, 100, 128, 384, 768, 4096,
] {
let a = pseudo_random_u8(9, dim);
let b = pseudo_random_u8(10, dim);
unsafe {
assert_eq!(
dot_u8_to_i32_avx2(&a, &b),
scalar_dot(&a, &b),
"dot dim={dim}"
);
assert_eq!(
sq_diff_u8_to_i32_avx2(&a, &b),
scalar_sq_diff(&a, &b),
"sq_diff dim={dim}"
);
assert_eq!(
abs_diff_u8_to_i32_avx2(&a, &b),
scalar_abs_diff(&a, &b),
"abs_diff dim={dim}"
);
}
}
let a = vec![255u8; 4096];
let b = vec![0u8; 4096];
unsafe {
assert_eq!(dot_u8_to_i32_avx2(&a, &a), 255 * 255 * 4096);
assert_eq!(sq_diff_u8_to_i32_avx2(&a, &b), 255 * 255 * 4096);
assert_eq!(abs_diff_u8_to_i32_avx2(&a, &b), 255 * 4096);
}
}
#[test]
fn dot_extremes_do_not_overflow_at_realistic_dim() {
let dim = 4096;
let a = vec![255u8; dim];
let b = vec![255u8; dim];
let expected = 255i32 * 255 * dim as i32;
assert_eq!(dot_u8_to_i32(&a, &b), expected);
}
#[test]
fn cosine_quantized_matches_f32_within_tolerance() {
let dim = 128;
let a_f32 = pseudo_random_f32(11, dim, -1.0, 1.0);
let b_f32 = pseudo_random_f32(12, dim, -1.0, 1.0);
let training = vec![Vector::new(a_f32.clone()), Vector::new(b_f32.clone())];
let params = ScalarQuantParams::train(&training).unwrap();
let q_a = params.quantize_slice(&a_f32);
let meta_a = QuantizedVectorMeta::from_quantized(&q_a, ¶ms);
let prepared = QuantizedQuery::prepare(&b_f32, ¶ms);
let approx_cosine_dist =
distance_quantized(DistanceMetric::Cosine, &prepared, &q_a, meta_a);
let exact_cosine_dist = DistanceMetric::Cosine.distance(&a_f32, &b_f32).unwrap();
let delta = (approx_cosine_dist - exact_cosine_dist).abs();
assert!(
delta < 0.02,
"cosine: quantized = {approx_cosine_dist}, f32 = {exact_cosine_dist}, |Δ| = {delta}"
);
}
#[test]
fn cosine_matches_f32_for_non_block_dim_with_padding() {
let dim = 100; let a_f32 = pseudo_random_f32(31, dim, -1.0, 1.0);
let b_f32 = pseudo_random_f32(32, dim, -1.0, 1.0);
let params =
ScalarQuantParams::train(&[Vector::new(a_f32.clone()), Vector::new(b_f32.clone())])
.unwrap();
let mut q_a = params.quantize_slice(&a_f32);
let meta_a = QuantizedVectorMeta::from_quantized(&q_a, ¶ms);
q_a.resize(padded_dim(dim), 0);
let prepared = QuantizedQuery::prepare(&b_f32, ¶ms);
assert_eq!(prepared.pad_dim, 128);
assert_eq!(q_a.len(), prepared.pad_dim);
let approx = distance_quantized(DistanceMetric::Cosine, &prepared, &q_a, meta_a);
let exact = DistanceMetric::Cosine.distance(&a_f32, &b_f32).unwrap();
assert!(
(approx - exact).abs() < 0.02,
"padded cosine: quantized = {approx}, f32 = {exact}"
);
}
#[test]
fn euclidean_quantized_matches_f32_within_tolerance() {
let dim = 128;
let a_f32 = pseudo_random_f32(21, dim, -1.0, 1.0);
let b_f32 = pseudo_random_f32(22, dim, -1.0, 1.0);
let params =
ScalarQuantParams::train(&[Vector::new(a_f32.clone()), Vector::new(b_f32.clone())])
.unwrap();
let q_a = params.quantize_slice(&a_f32);
let meta_a = QuantizedVectorMeta::from_quantized(&q_a, ¶ms);
let prepared = QuantizedQuery::prepare(&b_f32, ¶ms);
let approx_dist = distance_quantized(DistanceMetric::Euclidean, &prepared, &q_a, meta_a);
let exact_dist = DistanceMetric::Euclidean.distance(&a_f32, &b_f32).unwrap();
let rel_err = (approx_dist - exact_dist).abs() / exact_dist.max(1e-6);
assert!(
rel_err < 0.05,
"euclidean: quantized = {approx_dist}, f32 = {exact_dist}, rel_err = {rel_err}"
);
}
#[test]
fn manhattan_quantized_matches_f32_within_tolerance() {
let dim = 128;
let a_f32 = pseudo_random_f32(31, dim, -1.0, 1.0);
let b_f32 = pseudo_random_f32(32, dim, -1.0, 1.0);
let params =
ScalarQuantParams::train(&[Vector::new(a_f32.clone()), Vector::new(b_f32.clone())])
.unwrap();
let q_a = params.quantize_slice(&a_f32);
let meta_a = QuantizedVectorMeta::from_quantized(&q_a, ¶ms);
let prepared = QuantizedQuery::prepare(&b_f32, ¶ms);
let approx_dist = distance_quantized(DistanceMetric::Manhattan, &prepared, &q_a, meta_a);
let exact_dist = DistanceMetric::Manhattan.distance(&a_f32, &b_f32).unwrap();
let rel_err = (approx_dist - exact_dist).abs() / exact_dist.max(1e-6);
assert!(
rel_err < 0.02,
"manhattan: quantized = {approx_dist}, f32 = {exact_dist}, rel_err = {rel_err}"
);
}
#[test]
fn dot_product_quantized_matches_f32_within_tolerance() {
let dim = 128;
let a_f32 = pseudo_random_f32(41, dim, -1.0, 1.0);
let b_f32 = pseudo_random_f32(42, dim, -1.0, 1.0);
let params =
ScalarQuantParams::train(&[Vector::new(a_f32.clone()), Vector::new(b_f32.clone())])
.unwrap();
let q_a = params.quantize_slice(&a_f32);
let meta_a = QuantizedVectorMeta::from_quantized(&q_a, ¶ms);
let prepared = QuantizedQuery::prepare(&b_f32, ¶ms);
let approx_dist = distance_quantized(DistanceMetric::DotProduct, &prepared, &q_a, meta_a);
let exact_dist = DistanceMetric::DotProduct.distance(&a_f32, &b_f32).unwrap();
let abs_err = (approx_dist - exact_dist).abs();
assert!(
abs_err < 0.5,
"dot_product: quantized = {approx_dist}, f32 = {exact_dist}, |Δ| = {abs_err}"
);
}
#[test]
fn angular_quantized_matches_f32_within_tolerance() {
let dim = 128;
let a_f32 = pseudo_random_f32(51, dim, -1.0, 1.0);
let b_f32 = pseudo_random_f32(52, dim, -1.0, 1.0);
let params =
ScalarQuantParams::train(&[Vector::new(a_f32.clone()), Vector::new(b_f32.clone())])
.unwrap();
let q_a = params.quantize_slice(&a_f32);
let meta_a = QuantizedVectorMeta::from_quantized(&q_a, ¶ms);
let prepared = QuantizedQuery::prepare(&b_f32, ¶ms);
let approx_dist = distance_quantized(DistanceMetric::Angular, &prepared, &q_a, meta_a);
let exact_dist = DistanceMetric::Angular.distance(&a_f32, &b_f32).unwrap();
let abs_err = (approx_dist - exact_dist).abs();
assert!(
abs_err < 0.05,
"angular: quantized = {approx_dist}, f32 = {exact_dist}, |Δ| = {abs_err}"
);
}
#[test]
fn cosine_handles_zero_norm_query() {
let dim = 16;
let params = ScalarQuantParams {
offset: 0.0,
scale: 1.0,
};
let mut a: Vec<u8> = (0..dim as u8).collect();
let meta_a = QuantizedVectorMeta::from_quantized(&a, ¶ms);
a.resize(padded_dim(dim), 0);
let zero_query = vec![0.0_f32; dim];
let prepared = QuantizedQuery::prepare(&zero_query, ¶ms);
assert_eq!(prepared.norm_query, 0.0);
let dist = distance_quantized(DistanceMetric::Cosine, &prepared, &a, meta_a);
assert_eq!(
dist, 1.0,
"zero-norm query should yield max cosine distance, got {dist}"
);
}
#[test]
fn prepared_query_caches_segment_params() {
let dim = 8;
let query = vec![0.5_f32; dim];
let params = ScalarQuantParams {
offset: -1.0,
scale: 2.0 / 255.0,
};
let prepared = QuantizedQuery::prepare(&query, ¶ms);
assert_eq!(prepared.dim, dim);
assert_eq!(prepared.offset, params.offset);
assert_eq!(prepared.scale, params.scale);
assert_eq!(prepared.pad_dim, padded_dim(dim));
assert_eq!(prepared.q_data.len(), prepared.pad_dim);
let expected_sum: u32 = prepared.q_data.iter().map(|&x| x as u32).sum();
assert_eq!(prepared.sum_q, expected_sum);
}
use crate::vector::core::quantization::{pq_decode, pq_encode, pq_train_codebook};
fn small_pq_setup() -> (PqParams, Vec<f32>, Vec<Vector>) {
let dim = 4;
let m = 2;
let params = PqParams::from_dim_and_m(dim, m).unwrap();
let training: Vec<Vector> = vec![
Vector::new(vec![5.0, 5.0, 10.0, 10.0]),
Vector::new(vec![-5.0, -5.0, -10.0, -10.0]),
Vector::new(vec![5.1, 5.1, 10.1, 10.1]),
Vector::new(vec![-4.9, -4.9, -9.9, -9.9]),
];
let codebook = pq_train_codebook(dim, params, &training).unwrap();
(params, codebook, training)
}
#[test]
fn pq_lut_has_expected_dimensions() {
let (params, codebook, _) = small_pq_setup();
let query = vec![5.0_f32, 5.0, 10.0, 10.0];
let pq_query = PqQuery::prepare(&query, params, &codebook);
let m = params.m as usize;
let k = params.k as usize;
assert_eq!(pq_query.lut.len(), m * k);
for sub in 0..m {
let row = &pq_query.lut[sub * k..(sub + 1) * k];
let min = row.iter().cloned().fold(f32::INFINITY, f32::min);
assert!(
min < 1.0,
"sub-vector {sub}: nearest centroid should have L2² near 0, got {min}"
);
}
}
#[test]
fn pq_adc_euclidean_matches_decoded_l2() {
let (params, codebook, training) = small_pq_setup();
let encoded: Vec<Vec<u8>> = training
.iter()
.map(|v| pq_encode(&v.data, params, &codebook))
.collect();
let query = vec![5.0_f32, 5.0, 9.5, 9.5];
let pq_query = PqQuery::prepare(&query, params, &codebook);
for codes in &encoded {
let approx = distance_pq_adc(DistanceMetric::Euclidean, &pq_query, codes);
let decoded = pq_decode(codes, params, &codebook);
let mut ref_l2_sq = 0.0_f32;
for (q, d) in query.iter().zip(decoded.iter()) {
let diff = q - d;
ref_l2_sq += diff * diff;
}
let ref_l2 = ref_l2_sq.sqrt();
assert!(
(approx - ref_l2).abs() < 1e-3,
"ADC Euclidean = {approx}, reference = {ref_l2}, codes = {codes:?}"
);
}
}
#[test]
fn pq_adc_cosine_matches_l2_over_two_for_unit_norm() {
let dim = 4;
let m = 2;
let params = PqParams::from_dim_and_m(dim, m).unwrap();
fn unit_norm(v: &mut [f32]) {
let n: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
if n > 0.0 {
for x in v.iter_mut() {
*x /= n;
}
}
}
let mut training: Vec<Vector> = vec![
Vector::new({
let mut v = vec![1.0_f32, 0.5, 0.2, 0.1];
unit_norm(&mut v);
v
}),
Vector::new({
let mut v = vec![-1.0_f32, -0.5, -0.2, -0.1];
unit_norm(&mut v);
v
}),
Vector::new({
let mut v = vec![0.9_f32, 0.45, 0.2, 0.1];
unit_norm(&mut v);
v
}),
Vector::new({
let mut v = vec![-0.9_f32, -0.45, -0.2, -0.1];
unit_norm(&mut v);
v
}),
];
let codebook = pq_train_codebook(dim, params, &training).unwrap();
let mut query = vec![0.95_f32, 0.48, 0.2, 0.1];
unit_norm(&mut query);
let pq_query = PqQuery::prepare(&query, params, &codebook);
for v in training.iter_mut() {
let codes = pq_encode(&v.data, params, &codebook);
let cos = distance_pq_adc(DistanceMetric::Cosine, &pq_query, &codes);
let euc = distance_pq_adc(DistanceMetric::Euclidean, &pq_query, &codes);
let expected = (euc * euc) * 0.5;
assert!(
(cos - expected).abs() < 1e-3,
"cosine {cos} vs euc²/2 {expected} (euc = {euc})"
);
assert!((0.0..=2.0).contains(&cos), "cosine {cos} outside [0, 2]");
}
}
#[test]
fn pq_adc_unsupported_metric_returns_infinity() {
let (params, codebook, training) = small_pq_setup();
let codes = pq_encode(&training[0].data, params, &codebook);
let query = vec![1.0_f32, 1.0, 1.0, 1.0];
let pq_query = PqQuery::prepare(&query, params, &codebook);
for metric in [
DistanceMetric::Manhattan,
DistanceMetric::DotProduct,
DistanceMetric::Angular,
] {
let d = distance_pq_adc(metric, &pq_query, &codes);
assert!(
d.is_infinite(),
"metric {metric:?} should yield +inf, got {d}"
);
}
}
}