use super::knn_distinct_pbsamples_in_batch;
use legume_numeric::matrix::knn_match::ColumnDict;
use nalgebra::DMatrix;
use std::collections::HashSet;
#[test]
fn adaptive_recovers_knn_distinct_pbsamples() {
let knn = 10usize;
let mut feats: Vec<f32> = Vec::new();
let mut cell_to_pbsamp: Vec<usize> = Vec::new();
for pb in 0..3 {
for _ in 0..20 {
feats.push(pb as f32);
cell_to_pbsamp.push(pb);
}
}
for pb in 3..15 {
for _ in 0..3 {
feats.push(pb as f32);
cell_to_pbsamp.push(pb);
}
}
let n = feats.len();
let names: Vec<usize> = (0..n).collect(); let mat = DMatrix::<f32>::from_row_slice(1, n, &feats);
let bknn = ColumnDict::<usize>::from_dmatrix(mat, names);
let query = vec![0.0f32];
let hits = knn_distinct_pbsamples_in_batch(
&bknn,
&query,
knn,
&cell_to_pbsamp,
usize::MAX - 1,
None,
None,
)
.unwrap();
let distinct: HashSet<usize> = hits.iter().map(|&(p, _)| p).collect();
assert_eq!(
distinct.len(),
knn,
"should recover exactly knn distinct pb-samples, got {distinct:?}"
);
assert!(
distinct.iter().any(|&p| p >= 3),
"must reach pb-samples beyond the dense near cluster"
);
assert!(hits.windows(2).all(|w| w[0].1 <= w[1].1));
}
#[test]
fn adaptive_returns_all_when_fewer_than_knn() {
let knn = 10usize;
let mut feats: Vec<f32> = Vec::new();
let mut cell_to_pbsamp: Vec<usize> = Vec::new();
for pb in 0..3 {
for _ in 0..5 {
feats.push(pb as f32);
cell_to_pbsamp.push(pb);
}
}
let n = feats.len();
let names: Vec<usize> = (0..n).collect();
let mat = DMatrix::<f32>::from_row_slice(1, n, &feats);
let bknn = ColumnDict::<usize>::from_dmatrix(mat, names);
let query = vec![0.0f32];
let hits = knn_distinct_pbsamples_in_batch(
&bknn,
&query,
knn,
&cell_to_pbsamp,
usize::MAX - 1,
None,
None,
)
.unwrap();
let distinct: HashSet<usize> = hits.iter().map(|&(p, _)| p).collect();
assert_eq!(distinct.len(), 3);
}