use ordered_float::{FloatCore, OrderedFloat};
use std::collections::{BinaryHeap, HashMap};
use crate::error::{Failed, FailedError};
use crate::linalg::basic::arrays::{Array2, ArrayView1};
use crate::metrics::distance::PairwiseDistance;
use crate::numbers::floatnum::FloatNumber;
use crate::numbers::realnum::RealNumber;
#[derive(Debug, Clone)]
pub struct CosinePairParameters {
pub top_k: Option<usize>,
pub approximate: bool,
}
#[expect(clippy::derivable_impls)]
impl Default for CosinePairParameters {
fn default() -> Self {
Self {
top_k: None,
approximate: false,
}
}
}
#[derive(Debug, Clone)]
pub struct CosinePair<'a, T: RealNumber + FloatNumber, M: Array2<T>> {
pub samples: &'a M,
pub distances: HashMap<usize, PairwiseDistance<T>>,
pub neighbours: Vec<usize>,
row_norms: Vec<f64>,
pub parameters: CosinePairParameters,
}
impl<'a, T: RealNumber + FloatNumber + FloatCore, M: Array2<T>> CosinePair<'a, T, M> {
pub fn new(m: &'a M) -> Result<Self, Failed> {
Self::with_parameters(m, CosinePairParameters::default())
}
pub fn with_top_k(m: &'a M, top_k: usize) -> Result<Self, Failed> {
Self::with_parameters(
m,
CosinePairParameters {
top_k: Some(top_k),
approximate: false,
},
)
}
pub fn with_parameters(m: &'a M, parameters: CosinePairParameters) -> Result<Self, Failed> {
if m.shape().0 < 2 {
return Err(Failed::because(
FailedError::FindFailed,
"min number of rows should be 2",
));
}
let row_norms = (0..m.shape().0).map(|i| m.get_row(i).norm2()).collect();
let mut init = Self {
samples: m,
distances: HashMap::with_capacity(m.shape().0),
neighbours: Vec::with_capacity(m.shape().0),
row_norms,
parameters,
};
init.init();
Ok(init)
}
fn ordered_float(value: T) -> OrderedFloat<T> {
OrderedFloat(value)
}
fn extract_float(ordered: OrderedFloat<T>) -> T {
ordered.into_inner()
}
fn cosine_distance_with_norms(
row_i: &dyn ArrayView1<T>,
norm_i: f64,
row_j: &dyn ArrayView1<T>,
norm_j: f64,
) -> T {
let similarity = if norm_i == 0.0 || norm_j == 0.0 {
f64::MIN
} else {
row_i.dot(row_j).to_f64().unwrap() / (norm_i * norm_j)
};
T::from(1.0 - similarity).unwrap()
}
fn row_distance(&self, i: usize, j: usize) -> T {
let row_i = self.samples.get_row(i);
let row_j = self.samples.get_row(j);
Self::cosine_distance_with_norms(
row_i.as_ref(),
self.row_norms[i],
row_j.as_ref(),
self.row_norms[j],
)
}
fn init(&mut self) {
let len = self.samples.shape().0;
let mut distances = HashMap::with_capacity(len);
let mut neighbours = Vec::with_capacity(len);
neighbours.extend(0..len);
let mut best: Vec<Option<(OrderedFloat<T>, usize)>> = vec![None; len];
for i in 0..len {
for j in (i + 1)..len {
let distance = Self::ordered_float(self.row_distance(i, j));
if best[i].is_none_or(|(d, _)| distance < d) {
best[i] = Some((distance, j));
}
if best[j].is_none_or(|(d, _)| distance < d) {
best[j] = Some((distance, i));
}
}
}
for (i, best_of_i) in best.iter().enumerate() {
let (distance, neighbour) = best_of_i.expect("every row has at least one neighbour");
distances.insert(
i,
PairwiseDistance {
node: i,
neighbour: Some(neighbour),
distance: Some(Self::extract_float(distance)),
},
);
}
self.distances = distances;
self.neighbours = neighbours;
}
pub fn query_row_top_k(
&self,
query_row_index: usize,
k: usize,
) -> Result<Vec<(T, usize)>, Failed> {
if query_row_index >= self.samples.shape().0 {
return Err(Failed::because(
FailedError::FindFailed,
"Query row index out of bounds",
));
}
if k == 0 {
return Ok(Vec::new());
}
let n = self.samples.shape().0;
let max_candidates = self.parameters.top_k.unwrap_or(n);
let actual_k: usize = k.min(max_candidates);
let mut heap = BinaryHeap::with_capacity(actual_k + 1);
let query_row = self.samples.get_row(query_row_index);
let query_norm = self.row_norms[query_row_index];
let score_candidate = |heap: &mut BinaryHeap<(OrderedFloat<T>, usize)>, index: usize| {
let row = self.samples.get_row(index);
let distance = Self::cosine_distance_with_norms(
query_row.as_ref(),
query_norm,
row.as_ref(),
self.row_norms[index],
);
heap.push((Self::ordered_float(distance), index));
if heap.len() > actual_k {
heap.pop();
}
};
match (self.parameters.approximate, self.parameters.top_k) {
(true, Some(top_k)) => {
let step = (n / top_k).max(1);
for candidate in (0..n)
.step_by(step)
.filter(|&i| i != query_row_index)
.take(top_k)
{
score_candidate(&mut heap, candidate);
}
}
_ => {
for candidate in (0..n).filter(|&i| i != query_row_index) {
score_candidate(&mut heap, candidate);
}
}
}
let mut neighbors: Vec<(T, usize)> = heap
.into_iter()
.map(|(distance, index)| (Self::extract_float(distance), index))
.collect();
neighbors.sort_by(|a, b| {
Self::ordered_float(a.0)
.cmp(&Self::ordered_float(b.0))
.then(a.1.cmp(&b.1))
});
Ok(neighbors)
}
pub fn query_row(&self, query_row_index: usize, k: usize) -> Result<Vec<(T, usize)>, Failed> {
if query_row_index >= self.samples.shape().0 {
return Err(Failed::because(
FailedError::FindFailed,
"Query row index out of bounds",
));
}
if k == 0 {
return Ok(Vec::new());
}
let mut distances = self.distances_from(query_row_index);
distances.sort_by(|a, b| {
a.distance
.unwrap()
.partial_cmp(&b.distance.unwrap())
.unwrap_or(std::cmp::Ordering::Equal)
});
let neighbors: Vec<(T, usize)> = distances
.into_iter()
.take(k)
.map(|pd| (pd.distance.unwrap(), pd.neighbour.unwrap()))
.collect();
Ok(neighbors)
}
pub fn query(&self, query_vector: &Vec<T>, k: usize) -> Result<Vec<(T, usize)>, Failed> {
if query_vector.len() != self.samples.shape().1 {
return Err(Failed::because(
FailedError::FindFailed,
"Query vector dimension mismatch",
));
}
if k == 0 {
return Ok(Vec::new());
}
let query_norm = query_vector.norm2();
let mut distances = Vec::<PairwiseDistance<T>>::with_capacity(self.samples.shape().0);
for i in 0..self.samples.shape().0 {
let dataset_point = self.samples.get_row(i);
distances.push(PairwiseDistance {
node: i, neighbour: Some(i),
distance: Some(Self::cosine_distance_with_norms(
query_vector,
query_norm,
dataset_point.as_ref(),
self.row_norms[i],
)),
});
}
distances.sort_by(|a, b| {
a.distance
.unwrap()
.partial_cmp(&b.distance.unwrap())
.unwrap_or(std::cmp::Ordering::Equal)
});
let neighbors: Vec<(T, usize)> = distances
.into_iter()
.take(k)
.map(|pd| (pd.distance.unwrap(), pd.node))
.collect();
Ok(neighbors)
}
pub fn query_optimized(
&self,
query_row_index: usize,
k: usize,
) -> Result<Vec<(T, usize)>, Failed> {
self.query_row(query_row_index, k)
}
#[allow(dead_code)]
pub fn closest_pair(&self) -> PairwiseDistance<T> {
let mut a = self.neighbours[0]; let mut d = self.distances[&a].distance;
for p in self.neighbours.iter() {
if self.distances[p].distance < d {
a = *p; d = self.distances[p].distance;
}
}
let b = self.distances[&a].neighbour;
PairwiseDistance {
node: a,
neighbour: b,
distance: d,
}
}
#[allow(dead_code)]
pub fn ordered_pairs(&self) -> std::vec::IntoIter<&PairwiseDistance<T>> {
let mut distances = self
.distances
.values()
.collect::<Vec<&PairwiseDistance<T>>>();
distances.sort_by(|a, b| a.partial_cmp(b).unwrap());
distances.into_iter()
}
#[allow(dead_code)]
fn distances_from(&self, index_row: usize) -> Vec<PairwiseDistance<T>> {
let mut distances = Vec::<PairwiseDistance<T>>::with_capacity(self.samples.shape().0);
let query_row = self.samples.get_row(index_row);
let query_norm = self.row_norms[index_row];
for other in self.neighbours.iter() {
if index_row != *other {
let row = self.samples.get_row(*other);
distances.push(PairwiseDistance {
node: index_row,
neighbour: Some(*other),
distance: Some(Self::cosine_distance_with_norms(
query_row.as_ref(),
query_norm,
row.as_ref(),
self.row_norms[*other],
)),
})
}
}
distances
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::linalg::basic::arrays::Array1;
use crate::linalg::basic::{arrays::Array, matrix::DenseMatrix};
use crate::metrics::distance::Distance;
use crate::metrics::distance::cosine::Cosine;
use approx::{assert_relative_eq, relative_eq};
#[cfg_attr(
all(target_arch = "wasm32", not(target_os = "wasi")),
wasm_bindgen_test::wasm_bindgen_test
)]
#[test]
fn cosine_pair_initialization() {
let x = DenseMatrix::<f64>::from_2d_array(&[
&[5.1, 3.5, 1.4, 0.2],
&[4.9, 3.0, 1.4, 0.2],
&[4.7, 3.2, 1.3, 0.2],
&[4.6, 3.1, 1.5, 0.2],
&[5.0, 3.6, 1.4, 0.2],
&[5.4, 3.9, 1.7, 0.4],
])
.unwrap();
let cosine_pair = CosinePair::new(&x);
assert!(cosine_pair.is_ok());
let cp = cosine_pair.unwrap();
assert_eq!(cp.samples.shape().0, 6);
assert_eq!(cp.distances.len(), 6);
assert_eq!(cp.neighbours.len(), 6);
assert!(!cp.distances.is_empty());
assert!(!cp.neighbours.is_empty());
}
#[cfg_attr(
all(target_arch = "wasm32", not(target_os = "wasi")),
wasm_bindgen_test::wasm_bindgen_test
)]
#[test]
fn cosine_pair_minimum_rows_error() {
let x = DenseMatrix::<f64>::from_2d_array(&[&[5.1, 3.5, 1.4, 0.2]]).unwrap();
let result = CosinePair::new(&x);
assert!(result.is_err());
if let Err(e) = result {
let expected_error =
Failed::because(FailedError::FindFailed, "min number of rows should be 2");
assert_eq!(e, expected_error);
}
}
#[cfg_attr(
all(target_arch = "wasm32", not(target_os = "wasi")),
wasm_bindgen_test::wasm_bindgen_test
)]
#[test]
fn cosine_pair_closest_pair() {
let x = DenseMatrix::<f64>::from_2d_array(&[
&[1.0, 0.0],
&[0.0, 1.0],
&[1.0, 1.0],
&[2.0, 2.0], ])
.unwrap();
let cosine_pair = CosinePair::new(&x).unwrap();
let closest_pair = cosine_pair.closest_pair();
assert!(closest_pair.distance.is_some());
assert!(closest_pair.neighbour.is_some());
let distance = closest_pair.distance.unwrap();
assert!(distance >= 0.0 && distance <= 2.0); }
#[cfg_attr(
all(target_arch = "wasm32", not(target_os = "wasi")),
wasm_bindgen_test::wasm_bindgen_test
)]
#[test]
fn cosine_pair_identical_vectors() {
let x = DenseMatrix::<f64>::from_2d_array(&[
&[1.0, 2.0, 3.0],
&[1.0, 2.0, 3.0], &[4.0, 5.0, 6.0],
])
.unwrap();
let cosine_pair = CosinePair::new(&x).unwrap();
let closest_pair = cosine_pair.closest_pair();
let distance = closest_pair.distance.unwrap();
assert!((distance - 0.0).abs() < 1e-8);
}
#[cfg_attr(
all(target_arch = "wasm32", not(target_os = "wasi")),
wasm_bindgen_test::wasm_bindgen_test
)]
#[test]
fn cosine_pair_orthogonal_vectors() {
let x = DenseMatrix::<f64>::from_2d_array(&[
&[1.0, 0.0],
&[0.0, 1.0], &[2.0, 3.0],
])
.unwrap();
let cosine_pair = CosinePair::new(&x).unwrap();
let distances_from_first = cosine_pair.distances_from(0);
let orthogonal_distance = distances_from_first
.iter()
.find(|pd| pd.neighbour == Some(1))
.unwrap()
.distance
.unwrap();
assert!((orthogonal_distance - 1.0).abs() < 1e-8);
}
#[cfg_attr(
all(target_arch = "wasm32", not(target_os = "wasi")),
wasm_bindgen_test::wasm_bindgen_test
)]
#[test]
fn cosine_pair_ordered_pairs() {
let x = DenseMatrix::<f64>::from_2d_array(&[
&[1.0, 2.0],
&[2.0, 1.0],
&[3.0, 4.0],
&[4.0, 3.0],
])
.unwrap();
let cosine_pair = CosinePair::new(&x).unwrap();
let ordered_pairs: Vec<_> = cosine_pair.ordered_pairs().collect();
assert_eq!(ordered_pairs.len(), 4);
for i in 1..ordered_pairs.len() {
let prev_distance = ordered_pairs[i - 1].distance.unwrap();
let curr_distance = ordered_pairs[i].distance.unwrap();
assert!(prev_distance <= curr_distance);
}
}
#[cfg_attr(
all(target_arch = "wasm32", not(target_os = "wasi")),
wasm_bindgen_test::wasm_bindgen_test
)]
#[test]
fn cosine_pair_query_row() {
let x = DenseMatrix::<f64>::from_2d_array(&[
&[1.0, 0.0, 0.0],
&[0.0, 1.0, 0.0],
&[0.0, 0.0, 1.0],
&[1.0, 1.0, 0.0],
&[0.0, 1.0, 1.0],
])
.unwrap();
let cosine_pair = CosinePair::new(&x).unwrap();
let neighbors = cosine_pair.query_row(0, 2).unwrap();
assert_eq!(neighbors.len(), 2);
assert!(neighbors[0].0 <= neighbors[1].0);
for (distance, _) in &neighbors {
assert!(*distance >= 0.0 && *distance <= 2.0);
}
}
#[test]
fn cosine_pair_query_row_bounds_error() {
let x = DenseMatrix::<f64>::from_2d_array(&[&[1.0, 2.0], &[3.0, 4.0]]).unwrap();
let cosine_pair = CosinePair::new(&x).unwrap();
let result = cosine_pair.query_row(5, 1);
assert!(result.is_err());
if let Err(e) = result {
let expected_error =
Failed::because(FailedError::FindFailed, "Query row index out of bounds");
assert_eq!(e, expected_error);
}
}
#[test]
fn cosine_pair_query_row_k_zero() {
let x =
DenseMatrix::<f64>::from_2d_array(&[&[1.0, 2.0], &[3.0, 4.0], &[5.0, 6.0]]).unwrap();
let cosine_pair = CosinePair::new(&x).unwrap();
let neighbors = cosine_pair.query_row(0, 0).unwrap();
assert_eq!(neighbors.len(), 0);
}
#[cfg_attr(
all(target_arch = "wasm32", not(target_os = "wasi")),
wasm_bindgen_test::wasm_bindgen_test
)]
#[test]
fn cosine_pair_query_external_vector() {
let x = DenseMatrix::<f64>::from_2d_array(&[
&[1.0, 0.0, 0.0],
&[0.0, 1.0, 0.0],
&[0.0, 0.0, 1.0],
&[1.0, 1.0, 0.0],
])
.unwrap();
let cosine_pair = CosinePair::new(&x).unwrap();
let query_vector = vec![1.0, 0.5, 0.0];
let neighbors = cosine_pair.query(&query_vector, 2).unwrap();
assert_eq!(neighbors.len(), 2);
assert!(neighbors[0].0 <= neighbors[1].0);
for (distance, index) in &neighbors {
assert!(*distance >= 0.0 && *distance <= 2.0);
assert!(*index < x.shape().0);
}
}
#[test]
fn cosine_pair_query_dimension_mismatch() {
let x = DenseMatrix::<f64>::from_2d_array(&[&[1.0, 2.0, 3.0], &[4.0, 5.0, 6.0]]).unwrap();
let cosine_pair = CosinePair::new(&x).unwrap();
let query_vector = vec![1.0, 2.0]; let result = cosine_pair.query(&query_vector, 1);
assert!(result.is_err());
if let Err(e) = result {
let expected_error =
Failed::because(FailedError::FindFailed, "Query vector dimension mismatch");
assert_eq!(e, expected_error);
}
}
#[cfg_attr(
all(target_arch = "wasm32", not(target_os = "wasi")),
wasm_bindgen_test::wasm_bindgen_test
)]
#[test]
fn cosine_pair_query_k_zero_external() {
let x = DenseMatrix::<f64>::from_2d_array(&[&[1.0, 2.0], &[3.0, 4.0]]).unwrap();
let cosine_pair = CosinePair::new(&x).unwrap();
let query_vector = vec![1.0, 1.0];
let neighbors = cosine_pair.query(&query_vector, 0).unwrap();
assert_eq!(neighbors.len(), 0);
}
#[test]
fn cosine_pair_large_dataset() {
let x = DenseMatrix::<f64>::from_2d_array(&[
&[5.1, 3.5, 1.4, 0.2],
&[4.9, 3.0, 1.4, 0.2],
&[4.7, 3.2, 1.3, 0.2],
&[4.6, 3.1, 1.5, 0.2],
&[5.0, 3.6, 1.4, 0.2],
&[5.4, 3.9, 1.7, 0.4],
&[4.6, 3.4, 1.4, 0.3],
&[5.0, 3.4, 1.5, 0.2],
&[4.4, 2.9, 1.4, 0.2],
&[4.9, 3.1, 1.5, 0.1],
&[7.0, 3.2, 4.7, 1.4],
&[6.4, 3.2, 4.5, 1.5],
&[6.9, 3.1, 4.9, 1.5],
&[5.5, 2.3, 4.0, 1.3],
&[6.5, 2.8, 4.6, 1.5],
])
.unwrap();
let cosine_pair = CosinePair::new(&x).unwrap();
assert_eq!(cosine_pair.samples.shape().0, 15);
assert_eq!(cosine_pair.distances.len(), 15);
assert_eq!(cosine_pair.neighbours.len(), 15);
let closest_pair = cosine_pair.closest_pair();
assert!(closest_pair.distance.is_some());
assert!(closest_pair.neighbour.is_some());
let distance = closest_pair.distance.unwrap();
assert!(distance >= 0.0 && distance <= 2.0);
}
#[test]
fn query_row_top_k_top_k_limiting() {
let x = DenseMatrix::<f64>::from_2d_array(&[
&[1.0, 0.0, 0.0], &[0.0, 1.0, 0.0], &[0.0, 0.0, 1.0], &[1.0, 1.0, 0.0], &[0.5, 0.0, 0.0], &[2.0, 0.0, 0.0], &[0.0, 1.0, 1.0], &[3.0, 3.0, 3.0], ])
.unwrap();
let cosine_pair = CosinePair::with_top_k(&x, 4).unwrap();
let neighbors = cosine_pair.query_row_top_k(0, 3).unwrap();
assert_eq!(neighbors.len(), 3);
for i in 1..neighbors.len() {
assert!(
neighbors[i - 1].0 <= neighbors[i].0,
"Distances should be in ascending order: {} <= {}",
neighbors[i - 1].0,
neighbors[i].0
);
}
for (distance, index) in &neighbors {
assert!(
*distance >= 0.0 && *distance <= 2.0,
"Cosine distance {} should be between 0 and 2",
distance
);
assert!(
*index < x.shape().0,
"Neighbor index {} should be less than dataset size {}",
index,
x.shape().0
);
assert!(
*index != 0,
"Neighbor index should not include query point itself"
);
}
let closest_distance = neighbors[0].0;
assert!(
closest_distance < 0.01,
"Closest parallel vector should have distance close to 0, got {}",
closest_distance
);
let cosine_pair_full = CosinePair::new(&x).unwrap();
let neighbors_full = cosine_pair_full.query_row(0, 3).unwrap();
assert_eq!(neighbors.len(), neighbors_full.len());
let closest_idx_fast = neighbors[0].1;
let closest_idx_full = neighbors_full[0].1;
let closest_dist_fast = neighbors[0].0;
let closest_dist_full = neighbors_full[0].0;
if closest_idx_fast == closest_idx_full {
assert!(relative_eq!(
closest_dist_fast,
closest_dist_full,
epsilon = 1e-10
));
} else {
assert!(relative_eq!(
closest_dist_fast,
closest_dist_full,
epsilon = 1e-6
));
}
}
#[test]
fn query_row_top_k_performance_vs_accuracy() {
let large_dataset = DenseMatrix::<f32>::from_2d_array(&[
&[1.0f32, 2.0, 3.0, 4.0], &[1.1f32, 2.1, 3.1, 4.1], &[1.05f32, 2.05, 3.05, 4.05], &[2.0f32, 4.0, 6.0, 8.0], &[0.5f32, 1.0, 1.5, 2.0], &[-1.0f32, -2.0, -3.0, -4.0], &[4.0f32, 3.0, 2.0, 1.0], &[0.0f32, 0.0, 0.0, 0.1], &[10.0f32, 20.0, 30.0, 40.0], &[1.0f32, 0.0, 0.0, 0.0], &[0.0f32, 2.0, 0.0, 0.0], &[0.0f32, 0.0, 3.0, 0.0], ])
.unwrap();
let cosine_pair_limited = CosinePair::with_top_k(&large_dataset, 5).unwrap();
let neighbors_limited = cosine_pair_limited.query_row_top_k(0, 4).unwrap();
assert_eq!(neighbors_limited.len(), 4);
let result_oob = cosine_pair_limited.query_row_top_k(15, 2);
assert!(result_oob.is_err());
if let Err(e) = result_oob {
assert_eq!(
e,
Failed::because(FailedError::FindFailed, "Query row index out of bounds")
);
}
let neighbors_zero = cosine_pair_limited.query_row_top_k(0, 0).unwrap();
assert_eq!(neighbors_zero.len(), 0);
let neighbors_large_k = cosine_pair_limited.query_row_top_k(0, 20).unwrap();
assert!(neighbors_large_k.len() <= 11);
for i in 1..neighbors_limited.len() {
assert!(
neighbors_limited[i - 1].0 <= neighbors_limited[i].0,
"Distance ordering violation at position {}: {} > {}",
i,
neighbors_limited[i - 1].0,
neighbors_limited[i].0
);
}
let closest_distance = neighbors_limited[0].0;
assert!(
closest_distance < 0.1,
"Closest neighbor should be nearly parallel, distance: {}",
closest_distance
);
let cosine_pair_full = CosinePair::new(&large_dataset).unwrap();
let neighbors_full = cosine_pair_full.query_row(0, 4).unwrap();
let dist_diff = (neighbors_limited[0].0 - neighbors_full[0].0).abs();
assert!(
dist_diff < 0.01,
"Fast and full algorithms should give similar closest distances. Diff: {}",
dist_diff
);
let mut indices: Vec<usize> = neighbors_limited.iter().map(|(_, idx)| *idx).collect();
indices.sort();
indices.dedup();
assert_eq!(
indices.len(),
neighbors_limited.len(),
"All neighbor indices should be unique"
);
for &idx in &indices {
assert!(
idx < large_dataset.shape().0,
"Neighbor index {} should be valid",
idx
);
assert!(idx != 0, "Neighbor should not include query point itself");
}
for (distance, _) in &neighbors_limited {
assert!(!distance.is_nan(), "Distance should not be NaN");
assert!(distance.is_finite(), "Distance should be finite");
assert!(*distance >= 0.0, "Distance should be non-negative");
}
}
#[test]
fn cosine_pair_float_precision() {
let x = DenseMatrix::<f32>::from_2d_array(&[
&[1.0f32, 2.0, 3.0],
&[4.0f32, 5.0, 6.0],
&[7.0f32, 8.0, 9.0],
])
.unwrap();
let cosine_pair = CosinePair::new(&x).unwrap();
let closest_pair = cosine_pair.closest_pair();
assert!(closest_pair.distance.is_some());
let distance = closest_pair.distance.unwrap();
assert!(distance >= 0.0 && distance <= 2.0);
let neighbors = cosine_pair.query_row(0, 2).unwrap();
assert_eq!(neighbors.len(), 2);
assert_eq!(neighbors[0].1, 1);
assert_relative_eq!(neighbors[0].0, 0.025368154);
assert_eq!(neighbors[1].1, 2);
assert_relative_eq!(neighbors[1].0, 0.040588055);
}
#[cfg_attr(
all(target_arch = "wasm32", not(target_os = "wasi")),
wasm_bindgen_test::wasm_bindgen_test
)]
#[test]
fn cosine_pair_distances_from() {
let x = DenseMatrix::<f64>::from_2d_array(&[
&[1.0, 0.0],
&[0.0, 1.0],
&[1.0, 1.0],
&[2.0, 0.0],
])
.unwrap();
let cosine_pair = CosinePair::new(&x).unwrap();
let distances = cosine_pair.distances_from(0);
assert_eq!(distances.len(), 3);
for pd in &distances {
assert_eq!(pd.node, 0);
assert!(pd.neighbour.is_some());
assert!(pd.distance.is_some());
assert!(pd.neighbour.unwrap() != 0); }
}
#[cfg_attr(
all(target_arch = "wasm32", not(target_os = "wasi")),
wasm_bindgen_test::wasm_bindgen_test
)]
#[test]
fn cosine_pair_consistency_check() {
let x = DenseMatrix::<f64>::from_2d_array(&[
&[1.0, 2.0, 3.0],
&[4.0, 5.0, 6.0],
&[7.0, 8.0, 9.0],
&[2.0, 3.0, 4.0],
])
.unwrap();
let cosine_pair = CosinePair::new(&x).unwrap();
let neighbors_internal = cosine_pair.query_row(0, 2).unwrap();
let neighbors_optimized = cosine_pair.query_optimized(0, 2).unwrap();
assert_eq!(neighbors_internal.len(), neighbors_optimized.len());
for i in 0..neighbors_internal.len() {
let (dist1, idx1) = neighbors_internal[i];
let (dist2, idx2) = neighbors_optimized[i];
assert!((dist1 - dist2).abs() < 1e-10);
assert_eq!(idx1, idx2);
}
}
fn closest_pair_brute_force(
cosine_pair: &CosinePair<'_, f64, DenseMatrix<f64>>,
) -> PairwiseDistance<f64> {
use itertools::Itertools;
let m = cosine_pair.samples.shape().0;
let mut closest_pair = PairwiseDistance {
node: 0,
neighbour: None,
distance: Some(f64::MAX),
};
for pair in (0..m).combinations(2) {
let d = Cosine::new().distance(
&Vec::from_iterator(
cosine_pair.samples.get_row(pair[0]).iterator(0).copied(),
cosine_pair.samples.shape().1,
),
&Vec::from_iterator(
cosine_pair.samples.get_row(pair[1]).iterator(0).copied(),
cosine_pair.samples.shape().1,
),
);
if d < closest_pair.distance.unwrap() {
closest_pair.node = pair[0];
closest_pair.neighbour = Some(pair[1]);
closest_pair.distance = Some(d);
}
}
closest_pair
}
#[test]
fn cosine_pair_vs_brute_force() {
let x = DenseMatrix::<f64>::from_2d_array(&[
&[1.0, 2.0, 3.0],
&[4.0, 5.0, 6.0],
&[7.0, 8.0, 9.0],
&[1.1, 2.1, 3.1], ])
.unwrap();
let cosine_pair = CosinePair::new(&x).unwrap();
let cp_result = cosine_pair.closest_pair();
let brute_result = closest_pair_brute_force(&cosine_pair);
assert!((cp_result.distance.unwrap() - brute_result.distance.unwrap()).abs() < 1e-10);
}
fn mixed_direction_rows() -> DenseMatrix<f64> {
DenseMatrix::<f64>::from_2d_array(&[
&[1.0, 0.0, 0.0], &[0.9, 0.1, 0.0], &[0.0, 1.0, 0.0], &[0.95, 0.05, 0.0], &[0.0, 0.0, 1.0], &[0.99, 0.01, 0.0], &[0.0, 1.0, 1.0], &[0.8, 0.2, 0.0], ])
.unwrap()
}
fn cosine_distance_between_rows(m: &DenseMatrix<f64>, i: usize, j: usize) -> f64 {
Cosine::new().distance(
&Vec::from_iterator(m.get_row(i).iterator(0).copied(), m.shape().1),
&Vec::from_iterator(m.get_row(j).iterator(0).copied(), m.shape().1),
)
}
fn brute_force_top_k(m: &DenseMatrix<f64>, row: usize, k: usize) -> Vec<(f64, usize)> {
let mut scored: Vec<(f64, usize)> = (0..m.shape().0)
.filter(|&j| j != row)
.map(|j| (cosine_distance_between_rows(m, row, j), j))
.collect();
scored.sort_by(|a, b| a.0.partial_cmp(&b.0).unwrap().then(a.1.cmp(&b.1)));
scored.truncate(k);
scored
}
#[test]
fn query_row_top_k_is_exact_when_approximate_is_false_with_top_k() {
let x = mixed_direction_rows();
let cosine_pair = CosinePair::with_top_k(&x, 4).unwrap();
let neighbors = cosine_pair.query_row_top_k(0, 3).unwrap();
let expected = brute_force_top_k(&x, 0, 3);
assert_eq!(neighbors.len(), expected.len());
for (got, want) in neighbors.iter().zip(expected.iter()) {
assert_eq!(got.1, want.1);
assert_relative_eq!(got.0, want.0, epsilon = 1e-12);
}
let indices: Vec<usize> = neighbors.iter().map(|&(_, i)| i).collect();
assert_eq!(indices, vec![5, 3, 1]);
}
#[test]
fn query_row_top_k_is_exact_when_approximate_is_false_with_parameters() {
let x = mixed_direction_rows();
let cosine_pair = CosinePair::with_parameters(
&x,
CosinePairParameters {
top_k: Some(3),
approximate: false,
},
)
.unwrap();
let neighbors = cosine_pair.query_row_top_k(0, 3).unwrap();
let expected = brute_force_top_k(&x, 0, 3);
assert_eq!(neighbors.len(), expected.len());
for (got, want) in neighbors.iter().zip(expected.iter()) {
assert_eq!(got.1, want.1);
assert_relative_eq!(got.0, want.0, epsilon = 1e-12);
}
}
#[test]
fn query_row_top_k_exact_matches_query_row_from_new() {
let x = mixed_direction_rows();
let limited = CosinePair::with_top_k(&x, 4).unwrap();
let full = CosinePair::new(&x).unwrap();
for row in [0usize, 1, 3, 5, 7] {
let fast = limited.query_row_top_k(row, 3).unwrap();
let exact = full.query_row(row, 3).unwrap();
assert_eq!(fast.len(), exact.len(), "row {}", row);
for (got, want) in fast.iter().zip(exact.iter()) {
assert_eq!(got.1, want.1, "row {}", row);
assert_relative_eq!(got.0, want.0, epsilon = 1e-12);
}
}
}
#[test]
fn with_top_k_build_keeps_true_closest_neighbour_per_row() {
let x = mixed_direction_rows();
let cosine_pair = CosinePair::with_top_k(&x, 4).unwrap();
for i in 0..x.shape().0 {
let expected = brute_force_top_k(&x, i, 1)[0];
let stored = cosine_pair.distances[&i];
assert_eq!(stored.neighbour, Some(expected.1), "row {}", i);
assert_relative_eq!(stored.distance.unwrap(), expected.0, epsilon = 1e-12);
}
}
#[test]
fn with_top_k_full_parity_with_new() {
let x = mixed_direction_rows();
let n = x.shape().0;
let limited = CosinePair::with_top_k(&x, n - 1).unwrap();
let full = CosinePair::new(&x).unwrap();
for i in 0..n {
let a = limited.distances[&i];
let b = full.distances[&i];
assert_eq!(a.neighbour, b.neighbour, "row {}", i);
assert_relative_eq!(a.distance.unwrap(), b.distance.unwrap(), epsilon = 1e-12);
}
for row in [0usize, 1, 3, 5, 7] {
let fast = limited.query_row_top_k(row, 3).unwrap();
let exact = full.query_row(row, 3).unwrap();
assert_eq!(fast.len(), exact.len(), "row {}", row);
for (got, want) in fast.iter().zip(exact.iter()) {
assert_eq!(got.1, want.1, "row {}", row);
assert_relative_eq!(got.0, want.0, epsilon = 1e-12);
}
}
}
#[test]
fn build_closest_neighbour_matches_brute_force_per_row() {
let x = mixed_direction_rows();
let cosine_pair = CosinePair::new(&x).unwrap();
for i in 0..x.shape().0 {
let expected = brute_force_top_k(&x, i, 1)[0];
let stored = cosine_pair.distances[&i];
assert_eq!(stored.neighbour, Some(expected.1), "row {}", i);
assert_relative_eq!(stored.distance.unwrap(), expected.0, epsilon = 1e-12);
}
}
#[test]
fn query_row_top_k_samples_strided_candidates_when_approximate_is_true() {
let x = mixed_direction_rows();
let cosine_pair = CosinePair::with_parameters(
&x,
CosinePairParameters {
top_k: Some(4),
approximate: true,
},
)
.unwrap();
let neighbors = cosine_pair.query_row_top_k(0, 3).unwrap();
assert_eq!(neighbors.len(), 3);
let indices: Vec<usize> = neighbors.iter().map(|&(_, i)| i).collect();
let mut sorted_indices = indices.clone();
sorted_indices.sort_unstable();
assert_eq!(sorted_indices, vec![2, 4, 6]);
for (distance, index) in &neighbors {
assert_relative_eq!(*distance, 1.0, epsilon = 1e-12);
assert_ne!(*index, 0);
}
for i in 1..neighbors.len() {
assert!(neighbors[i - 1].0 <= neighbors[i].0);
}
}
}