use super::metric::l2_sq_kernel;
use super::VecPoint;
use multiversion::multiversion;
#[multiversion(targets = "simd")]
fn scan_sq(points: &[VecPoint], query: &[f32], out: &mut Vec<(usize, f32)>) {
out.clear();
out.extend(
points
.iter()
.enumerate()
.map(|(i, p)| (i, l2_sq_kernel(&p.data, query))),
);
}
pub(super) fn topk(points: &[VecPoint], query: &[f32], k: usize) -> (Vec<usize>, Vec<f32>) {
let mut scored: Vec<(usize, f32)> = Vec::with_capacity(points.len());
scan_sq(points, query, &mut scored);
let kk = k.min(scored.len());
if kk == 0 {
return (Vec::new(), Vec::new());
}
scored.select_nth_unstable_by(kk - 1, |a, b| a.1.total_cmp(&b.1));
scored.truncate(kk);
scored.sort_unstable_by(|a, b| a.1.total_cmp(&b.1));
let indices = scored.iter().map(|&(i, _)| i).collect();
let distances = scored.iter().map(|&(_, d2)| d2.sqrt()).collect();
(indices, distances)
}