#[cfg(test)] mod tests { use std::sync::atomic::AtomicBool; use common::counter::hardware_counter::HardwareCounterCell; use quantization::encoded_storage::TestEncodedStorageBuilder; use quantization::encoded_vectors::{DistanceType, EncodedVectors, VectorParameters}; use quantization::encoded_vectors_binary::{ self, BitsStoreType, EncodedVectorsBin, Encoding, QueryEncoding, }; use rand::{RngExt, SeedableRng}; use strum::IntoEnumIterator; use crate::metrics::dot_similarity; fn generate_number(rng: &mut rand::rngs::StdRng) -> f32 { rng.random_range(-1.0..1.0) } fn generate_vector(dim: usize, rng: &mut rand::rngs::StdRng) -> Vec { (0..dim) .map(|_| generate_number(rng) / (dim as f32).sqrt()) .collect() } fn get_top(scores: &[f32], count: usize, invert: bool) -> Vec { let mut indices: Vec = (0..scores.len()).collect(); indices.sort_by(|&a, &b| scores[b].partial_cmp(&scores[a]).unwrap()); if invert { indices.reverse(); } indices.into_iter().take(count).collect() } fn match_count(ids1: &[usize], ids2: &[usize]) -> usize { ids1.iter().filter(|&&id| ids2.contains(&id)).count() } #[test] fn test_binary_dot() { test_binary_dot_impl::(0, false); test_binary_dot_impl::(600, false); test_binary_dot_impl::(601, false); test_binary_dot_impl::(600, false); } #[test] fn test_binary_dot_inverted() { test_binary_dot_impl::(700, true); test_binary_dot_impl::(700, true); } fn test_binary_dot_impl(vector_dim: usize, invert: bool) { let vectors_count = 1000; let mut rng = rand::rngs::StdRng::seed_from_u64(42); let mut vector_data: Vec> = Vec::new(); for _ in 0..vectors_count { vector_data.push(generate_vector(vector_dim, &mut rng)); } let encodings = [ Encoding::OneBit, Encoding::OneAndHalfBits, Encoding::TwoBits, ]; let encoded: Vec<_> = encodings .iter() .map(|&encoding| { let quantized_vector_size = encoded_vectors_binary::get_quantized_vector_size_from_params::( vector_dim, encoding, ); EncodedVectorsBin::::encode( vector_data.iter(), TestEncodedStorageBuilder::new(None, quantized_vector_size), &VectorParameters { dim: vector_dim, deprecated_count: None, distance_type: DistanceType::Dot, invert, }, encoding, QueryEncoding::SameAsStorage, None, &AtomicBool::new(false), ) .unwrap() }) .collect(); let top = 10; let query: Vec = generate_vector(vector_dim, &mut rng); let orig_scores: Vec = vector_data .iter() .map(|vector| dot_similarity(&query, vector)) .collect(); let original_top = get_top(&orig_scores, top, invert); let tops = encoded .iter() .map(|encoded| { let query_encoded = encoded.encode_query(&query); let scores: Vec = (0..vector_data.len()) .map(|index| { encoded.score_point( &query_encoded, index as u32, &HardwareCounterCell::new(), ) }) .collect(); let tops = get_top(&scores, top, false); match_count(&original_top, &tops) }) .collect::>(); // Check if encoding has more accuracy than previous one for i in 1..tops.len() { assert!( tops[i] >= tops[i - 1], "Encoding {} has less accuracy than encoding {}", i, i - 1 ); } } #[test] fn test_binary_dot_asymetric() { test_binary_dot_asymentric_impl::(0, Encoding::OneBit, false); test_binary_dot_asymentric_impl::(0, Encoding::OneBit, false); test_binary_dot_asymentric_impl::(1024, Encoding::OneBit, false); test_binary_dot_asymentric_impl::(601, Encoding::OneBit, false); test_binary_dot_asymentric_impl::(600, Encoding::OneBit, false); test_binary_dot_asymentric_impl::(1024, Encoding::OneAndHalfBits, false); test_binary_dot_asymentric_impl::(601, Encoding::OneAndHalfBits, false); test_binary_dot_asymentric_impl::(600, Encoding::OneAndHalfBits, false); test_binary_dot_asymentric_impl::(1024, Encoding::TwoBits, false); test_binary_dot_asymentric_impl::(701, Encoding::TwoBits, false); test_binary_dot_asymentric_impl::(700, Encoding::TwoBits, false); } #[test] fn test_binary_dot_inverted_asymetric() { test_binary_dot_asymentric_impl::(0, Encoding::OneBit, true); test_binary_dot_asymentric_impl::(0, Encoding::OneBit, true); test_binary_dot_asymentric_impl::(1024, Encoding::OneBit, true); test_binary_dot_asymentric_impl::(601, Encoding::OneBit, true); test_binary_dot_asymentric_impl::(600, Encoding::OneBit, true); } fn test_binary_dot_asymentric_impl( vector_dim: usize, encoding: Encoding, invert: bool, ) { let vectors_count = 1000; let mut rng = rand::rngs::StdRng::seed_from_u64(43); let mut vector_data: Vec> = Vec::new(); for _ in 0..vectors_count { vector_data.push(generate_vector(vector_dim, &mut rng)); } let encoded: Vec<_> = QueryEncoding::iter() .map(|query_encoding| { let quantized_vector_size = encoded_vectors_binary::get_quantized_vector_size_from_params::( vector_dim, encoding, ); EncodedVectorsBin::::encode( vector_data.iter(), TestEncodedStorageBuilder::new(None, quantized_vector_size), &VectorParameters { dim: vector_dim, deprecated_count: None, distance_type: DistanceType::Dot, invert, }, encoding, query_encoding, None, &AtomicBool::new(false), ) .unwrap() }) .collect(); let top = 10; let query: Vec = generate_vector(vector_dim, &mut rng); let orig_scores: Vec = vector_data .iter() .map(|vector| dot_similarity(&query, vector)) .collect(); let original_top = get_top(&orig_scores, top, invert); let tops = encoded .iter() .map(|encoded| { let query_encoded = encoded.encode_query(&query); let scores: Vec = (0..vector_data.len()) .map(|index| { encoded.score_point( &query_encoded, index as u32, &HardwareCounterCell::new(), ) }) .collect(); let tops = get_top(&scores, top, false); match_count(&original_top, &tops) }) .collect::>(); // Check if encoding has more accuracy than previous one for i in 1..tops.len() { assert!( tops[i] >= tops[0], "Encoding {i} has less accuracy than original encoding", ); } } }