use xz_embed::*;
struct Lcg {
state: u64,
}
impl Lcg {
const fn new(seed: u64) -> Self {
Self { state: seed }
}
fn next_f32(&mut self) -> f32 {
self.state = self.state.wrapping_mul(6364136223846793005).wrapping_add(1442695040888963407);
let bits = (self.state >> 32) as u32;
(bits as f64 / (u32::MAX as f64 + 1.0) * 2.0 - 1.0) as f32
}
}
fn generate_random_vectors(num_vectors: usize, dim: usize, seed: u64) -> Vec<Vec<f32>> {
let mut rng = Lcg::new(seed);
(0..num_vectors).map(|_| (0..dim).map(|_| rng.next_f32()).collect()).collect()
}
#[test]
fn test_scalar_round_trip_max_error() {
let dim = 64;
let samples = generate_random_vectors(100, dim, 42);
let mut ranges: Vec<(f32, f32)> = vec![(f32::MAX, f32::MIN); dim];
for sample in &samples {
for d in 0..dim {
if sample[d] < ranges[d].0 {
ranges[d].0 = sample[d];
}
if sample[d] > ranges[d].1 {
ranges[d].1 = sample[d];
}
}
}
let sq = ScalarQuantizer::from_samples(&samples, 8);
let compressed = sq.compress(&samples);
assert_eq!(compressed.len(), samples.len());
assert_eq!(compressed[0].len(), dim);
let reconstructed = sq.decompress(&compressed);
assert_eq!(reconstructed.len(), samples.len());
let mut max_relative_error = 0.0f32;
for (orig, recon) in samples.iter().zip(reconstructed.iter()) {
for d in 0..dim {
let range = ranges[d].1 - ranges[d].0;
if range < f32::EPSILON {
continue; }
let abs_err = (orig[d] - recon[d]).abs();
let rel_err = abs_err / range;
if rel_err > max_relative_error {
max_relative_error = rel_err;
}
}
}
assert!(
max_relative_error < 0.01,
"max relative error {max_relative_error:.6} exceeds 1% threshold"
);
}
#[test]
fn test_scalar_round_trip_deterministic() {
let vectors = vec![vec![0.0, 0.5, 1.0], vec![-1.0, 0.0, 1.0], vec![0.25, -0.75, 0.1]];
let sq = ScalarQuantizer::from_samples(&vectors, 8);
let compressed1 = sq.compress(&vectors);
let compressed2 = sq.compress(&vectors);
assert_eq!(compressed1, compressed2);
let decompressed1 = sq.decompress(&compressed1);
let decompressed2 = sq.decompress(&compressed1);
assert_eq!(decompressed1, decompressed2);
}
#[test]
fn test_scalar_round_trip_known_values() {
let sq = ScalarQuantizer::new(vec![(0.0, 1.0)], 8);
let compressed = sq.compress(&[vec![0.5]]);
let reconstructed = sq.decompress(&compressed);
assert!((reconstructed[0][0] - 0.5).abs() < 0.01);
}
#[test]
fn test_scalar_quantizer_empty_samples() {
let sq = ScalarQuantizer::from_samples(&[], 8);
assert!(sq.compress(&[]).is_empty());
assert!(sq.decompress(&[]).is_empty());
assert!(sq.compress(&[vec![1.0, 2.0]]).is_empty());
}
#[test]
fn test_scalar_quantizer_single_sample() {
let sq = ScalarQuantizer::from_samples(&[vec![0.5, -0.3, 0.0]], 8);
let compressed = sq.compress(&[vec![0.5, -0.3, 0.0]]);
assert_eq!(compressed.len(), 1);
assert_eq!(compressed[0].len(), 3);
let decompressed = sq.decompress(&compressed);
assert!((decompressed[0][0] - 0.5).abs() < f32::EPSILON);
}
#[test]
fn test_scalar_quantizer_exact_boundaries() {
let sq = ScalarQuantizer::new(vec![(-1.0, 1.0)], 8);
let compressed = sq.compress(&[vec![-1.0], vec![1.0]]);
assert_eq!(compressed[0][0], 0);
assert_eq!(compressed[1][0], 255);
let decompressed = sq.decompress(&compressed);
assert!((decompressed[0][0] + 1.0).abs() < f32::EPSILON);
assert!((decompressed[1][0] - 1.0).abs() < f32::EPSILON);
}
#[test]
fn test_scalar_compress_preserves_dimension_count() {
let samples = generate_random_vectors(10, 16, 123);
let sq = ScalarQuantizer::from_samples(&samples, 8);
let compressed = sq.compress(&samples);
assert_eq!(compressed.len(), 10);
for code in &compressed {
assert_eq!(code.len(), 16);
}
}
#[test]
fn test_scalar_decompress_reverses_compress() {
let samples = generate_random_vectors(20, 8, 99);
let sq = ScalarQuantizer::from_samples(&samples, 8);
let compressed = sq.compress(&samples);
let decompressed = sq.decompress(&compressed);
let recompressed = sq.compress(&decompressed);
assert_eq!(compressed, recompressed);
}
#[test]
fn test_scalar_quantizer_zero_range_dimensions() {
let samples = vec![vec![0.5, 1.0, 2.0], vec![0.5, 2.0, 3.0], vec![0.5, 3.0, 4.0]];
let sq = ScalarQuantizer::from_samples(&samples, 8);
let compressed = sq.compress(&samples);
let decompressed = sq.decompress(&compressed);
for v in &decompressed {
assert!((v[0] - 0.5).abs() < f32::EPSILON);
}
}
#[test]
fn test_pq_train_compress_decompress() {
let dim = 12;
let num_sub = 3; let bits = 8;
let samples = generate_random_vectors(200, dim, 42);
let pq = ProductQuantizer::train(&samples, num_sub, bits).unwrap();
let compressed = pq.compress(&samples);
assert_eq!(compressed.len(), samples.len());
for code in &compressed {
assert_eq!(code.len(), num_sub);
}
let decompressed = pq.decompress(&compressed);
assert_eq!(decompressed.len(), samples.len());
for v in &decompressed {
assert_eq!(v.len(), dim);
}
}
#[test]
fn test_pq_reconstruction_idempotent() {
let dim = 8;
let num_sub = 4; let bits = 4;
let samples = generate_random_vectors(100, dim, 77);
let pq = ProductQuantizer::train(&samples, num_sub, bits).unwrap();
let compressed = pq.compress(&samples);
let decompressed = pq.decompress(&compressed);
let recompressed = pq.compress(&decompressed);
assert_eq!(compressed, recompressed);
}
#[test]
fn test_pq_empty_samples() {
let result = ProductQuantizer::train(&[], 4, 8);
assert!(result.is_err());
let err = result.unwrap_err();
assert!(matches!(err, EmbedError::Config(_)));
}
#[test]
fn test_pq_zero_sub_vectors() {
let samples = vec![vec![1.0, 2.0, 3.0, 4.0]];
let result = ProductQuantizer::train(&samples, 0, 8);
assert!(result.is_err());
let err = result.unwrap_err();
assert!(matches!(err, EmbedError::Config(_)));
}
#[test]
fn test_pq_dimension_not_divisible() {
let samples = vec![
vec![1.0, 2.0, 3.0, 4.0, 5.0], ];
let result = ProductQuantizer::train(&samples, 2, 8);
assert!(result.is_err());
let err = result.unwrap_err();
assert!(matches!(err, EmbedError::Config(_)));
}
#[test]
fn test_pq_bits_exceeds_max() {
let samples = vec![vec![1.0, 2.0, 3.0, 4.0]];
let result = ProductQuantizer::train(&samples, 2, 9); assert!(result.is_err());
let err = result.unwrap_err();
match err {
EmbedError::Config(msg) => {
assert!(msg.contains("bits"), "expected message about bits limit, got: {msg}");
}
_ => panic!("expected Config error"),
}
}
#[test]
fn test_pq_single_sample() {
let dim = 8;
let num_sub = 2; let bits = 2; let samples = vec![vec![0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8]];
let pq = ProductQuantizer::train(&samples, num_sub, bits).unwrap();
let compressed = pq.compress(&samples);
assert_eq!(compressed.len(), 1);
let decompressed = pq.decompress(&compressed);
assert_eq!(decompressed.len(), 1);
assert_eq!(decompressed[0].len(), dim);
}
#[test]
fn test_pq_k_greater_than_samples() {
let dim = 4;
let num_sub = 2; let bits = 8; let samples = generate_random_vectors(5, dim, 123);
let pq = ProductQuantizer::train(&samples, num_sub, bits).unwrap();
let compressed = pq.compress(&samples);
assert_eq!(compressed.len(), 5);
let decompressed = pq.decompress(&compressed);
assert_eq!(decompressed.len(), 5);
}
#[test]
fn test_pq_compress_empty_vectors() {
let samples = generate_random_vectors(50, 6, 99);
let pq = ProductQuantizer::train(&samples, 3, 8).unwrap();
let compressed = pq.compress(&[]);
assert!(compressed.is_empty());
}
#[test]
fn test_pq_decompress_empty() {
let samples = generate_random_vectors(50, 6, 99);
let pq = ProductQuantizer::train(&samples, 3, 8).unwrap();
let decompressed = pq.decompress(&[]);
assert!(decompressed.is_empty());
}
#[test]
fn test_pq_codebook_coverage() {
let dim = 6;
let num_sub = 3;
let bits = 4; let samples = generate_random_vectors(100, dim, 55);
let pq = ProductQuantizer::train(&samples, num_sub, bits).unwrap();
let compressed = pq.compress(&samples);
for code in &compressed {
for &c in code {
assert!(c < 16, "code index {c} exceeds maximum for {bits} bits");
}
}
}
#[test]
fn test_pq_noop_quantizer_round_trip() {
let vectors = vec![vec![1.0f32, 2.0, 3.0], vec![4.0, 5.0, 6.0], vec![7.0, 8.0, 9.0]];
let nq = xz_embed::quantize::NoopQuantizer;
let compressed = nq.compress(&vectors);
let decompressed = nq.decompress(&compressed);
assert_eq!(vectors, decompressed);
assert_eq!(compressed[0].len(), 12);
}
#[test]
fn test_scalar_quantizer_matches_trait() {
fn _assert_trait<T: VectorQuantizer>(_: &T) {}
let sq = ScalarQuantizer::from_samples(&[vec![1.0, 2.0]], 8);
_assert_trait(&sq);
}
#[test]
fn test_product_quantizer_matches_trait() {
fn _assert_trait<T: VectorQuantizer>(_: &T) {}
let pq = ProductQuantizer::train(&[vec![1.0, 2.0, 3.0, 4.0]], 2, 8).unwrap();
_assert_trait(&pq);
}
#[test]
fn test_vector_quantizer_object_safety() {
fn use_dyn(q: &dyn VectorQuantizer, v: &[Vec<f32>]) {
let _ = q.compress(v);
}
let sq = ScalarQuantizer::from_samples(&[vec![1.0, 2.0]], 8);
use_dyn(&sq, &[vec![0.5, -0.5]]);
let pq = ProductQuantizer::train(&[vec![1.0, 2.0, 3.0, 4.0]], 2, 8).unwrap();
use_dyn(&pq, &[vec![1.0, 2.0, 3.0, 4.0]]);
}