const D: usize = 6;
const LEAF_SIZE: usize = 24;
#[cfg(feature = "parallel")]
const PARALLEL_BUILD_CUTOFF: usize = 32 * 1024;
#[repr(C)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) struct Point {
pub v: [i16; D],
pub id: u32,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) struct Nearest {
pub id: u32,
pub distance_squared: u64,
}
pub(crate) struct KdTree {
points: Vec<Point>,
}
impl KdTree {
pub(crate) fn build(points: Vec<Point>) -> Self {
Self::build_with::<true>(points)
}
#[allow(dead_code)] pub(crate) fn build_serial(points: Vec<Point>) -> Self {
Self::build_with::<false>(points)
}
fn build_with<const PARALLEL: bool>(mut points: Vec<Point>) -> Self {
build_recursive::<PARALLEL>(&mut points, 0);
KdTree { points }
}
pub(crate) fn nearest(&self, query: &[i16; D]) -> Option<Nearest> {
if self.points.is_empty() {
return None;
}
let mut best = Nearest {
id: u32::MAX,
distance_squared: u64::MAX,
};
search_recursive(&self.points, query, 0, &mut best);
Some(best)
}
#[allow(dead_code)] pub(crate) fn len(&self) -> usize {
self.points.len()
}
#[allow(dead_code)] pub(crate) fn is_empty(&self) -> bool {
self.points.is_empty()
}
}
fn build_recursive<const PARALLEL: bool>(points: &mut [Point], depth: usize) {
if points.len() <= LEAF_SIZE {
return;
}
let axis = depth % D;
let middle = points.len() / 2;
let (left, _pivot, right) =
points.select_nth_unstable_by(middle, |a, b| a.v[axis].cmp(&b.v[axis]));
build_children::<PARALLEL>(left, right, depth);
}
#[cfg(feature = "parallel")]
fn build_children<const PARALLEL: bool>(left: &mut [Point], right: &mut [Point], depth: usize) {
if !PARALLEL || left.len() + right.len() < PARALLEL_BUILD_CUTOFF {
build_recursive::<PARALLEL>(left, depth + 1);
build_recursive::<PARALLEL>(right, depth + 1);
} else {
rayon::join(
|| build_recursive::<PARALLEL>(left, depth + 1),
|| build_recursive::<PARALLEL>(right, depth + 1),
);
}
}
#[cfg(not(feature = "parallel"))]
fn build_children<const PARALLEL: bool>(left: &mut [Point], right: &mut [Point], depth: usize) {
build_recursive::<PARALLEL>(left, depth + 1);
build_recursive::<PARALLEL>(right, depth + 1);
}
#[inline(always)]
fn consider(best: &mut Nearest, candidate: &Point, distance_squared: u64) {
if distance_squared < best.distance_squared
|| (distance_squared == best.distance_squared && candidate.id < best.id)
{
best.distance_squared = distance_squared;
best.id = candidate.id;
}
}
fn search_recursive(points: &[Point], query: &[i16; D], depth: usize, best: &mut Nearest) {
if points.is_empty() {
return;
}
if points.len() <= LEAF_SIZE {
for p in points {
consider(best, p, distance_squared(query, &p.v));
}
return;
}
let middle = points.len() / 2;
let pivot = &points[middle];
consider(best, pivot, distance_squared(query, &pivot.v));
let axis = depth % D;
let delta = query[axis] as i64 - pivot.v[axis] as i64;
let (near, far) = if delta <= 0 {
(&points[..middle], &points[middle + 1..])
} else {
(&points[middle + 1..], &points[..middle])
};
search_recursive(near, query, depth + 1, best);
if (delta * delta) as u64 <= best.distance_squared {
search_recursive(far, query, depth + 1, best);
}
}
#[inline(always)]
fn distance_squared(a: &[i16; D], b: &[i16; D]) -> u64 {
let mut sum = 0i64;
for k in 0..D {
let d = a[k] as i64 - b[k] as i64;
sum += d * d;
}
sum as u64
}
#[cfg(test)]
mod tests {
use super::*;
use crate::rng::Rng;
fn brute_force(points: &[Point], query: &[i16; D]) -> Option<Nearest> {
let mut best = Nearest {
id: u32::MAX,
distance_squared: u64::MAX,
};
for p in points {
consider(&mut best, p, distance_squared(query, &p.v));
}
(!points.is_empty()).then_some(best)
}
fn random_points(rng: &mut Rng, count: usize, spread: i16) -> Vec<Point> {
(0..count)
.map(|i| {
let mut v = [0i16; D];
for slot in v.iter_mut() {
*slot = (rng.next_u64() % (spread as u64 * 2 + 1)) as i16 - spread;
}
Point { v, id: i as u32 }
})
.collect()
}
fn random_key(rng: &mut Rng, spread: i16) -> [i16; D] {
let mut v = [0i16; D];
for slot in v.iter_mut() {
*slot = (rng.next_u64() % (spread as u64 * 2 + 1)) as i16 - spread;
}
v
}
#[test]
fn nearest_agrees_with_brute_force() {
let mut rng = Rng::seed_from_u64(0xA11CE);
for &count in &[1usize, 2, 23, 24, 25, 49, 500, 2000] {
let points = random_points(&mut rng, count, 1000);
let tree = KdTree::build(points.clone());
assert_eq!(tree.len(), count);
for _ in 0..200 {
let query = random_key(&mut rng, 1200);
assert_eq!(
tree.nearest(&query),
brute_force(&points, &query),
"disagreed on a {count}-point tree at query {query:?}"
);
}
}
}
#[test]
fn nearest_agrees_with_brute_force_when_ties_are_everywhere() {
let mut rng = Rng::seed_from_u64(0xB0B);
for &count in &[30usize, 200, 1500] {
let points = random_points(&mut rng, count, 2);
let tree = KdTree::build(points.clone());
for _ in 0..500 {
let query = random_key(&mut rng, 3);
assert_eq!(
tree.nearest(&query),
brute_force(&points, &query),
"disagreed on a {count}-point tree with heavy ties"
);
}
}
}
#[test]
fn identical_points_resolve_to_the_lowest_id() {
let points: Vec<Point> = (0..100).map(|i| Point { v: [7; D], id: i }).collect();
let tree = KdTree::build(points);
let found = tree.nearest(&[9; D]).expect("tree is not empty");
assert_eq!(found.id, 0);
assert_eq!(found.distance_squared, 4 * D as u64);
}
#[test]
fn the_answer_does_not_depend_on_build_order() {
let mut rng = Rng::seed_from_u64(0xC0FFEE);
let points = random_points(&mut rng, 800, 6);
let queries: Vec<[i16; D]> = (0..300).map(|_| random_key(&mut rng, 8)).collect();
let reference = KdTree::build(points.clone());
let expected: Vec<_> = queries.iter().map(|q| reference.nearest(q)).collect();
for _ in 0..8 {
let mut shuffled = points.clone();
rng.shuffle(&mut shuffled);
let tree = KdTree::build(shuffled);
let got: Vec<_> = queries.iter().map(|q| tree.nearest(q)).collect();
assert_eq!(
got, expected,
"a reordered build produced different answers"
);
}
}
#[test]
fn the_parallel_build_matches_the_serial_one() {
let mut rng = Rng::seed_from_u64(0xD00D);
let points = random_points(&mut rng, 5000, 400);
let parallel = KdTree::build(points.clone());
let serial = KdTree::build_serial(points);
assert_eq!(
parallel.points, serial.points,
"the two builds produced different trees"
);
}
#[test]
fn degenerate_trees_behave() {
let empty = KdTree::build(Vec::new());
assert!(empty.is_empty());
assert_eq!(empty.len(), 0);
assert_eq!(empty.nearest(&[0; D]), None);
let single = KdTree::build(vec![Point {
v: [1, 2, 3, 4, 5, 6],
id: 42,
}]);
assert!(!single.is_empty());
let found = single
.nearest(&[1, 2, 3, 4, 5, 6])
.expect("tree is not empty");
assert_eq!(
found,
Nearest {
id: 42,
distance_squared: 0
}
);
}
#[test]
fn extreme_keys_do_not_overflow() {
let points = vec![
Point {
v: [i16::MIN; D],
id: 0,
},
Point {
v: [i16::MAX; D],
id: 1,
},
];
let tree = KdTree::build(points.clone());
let query = [i16::MAX; D];
let found = tree.nearest(&query).expect("tree is not empty");
assert_eq!(
found,
Nearest {
id: 1,
distance_squared: 0
}
);
let far = distance_squared(&[i16::MIN; D], &[i16::MAX; D]);
assert_eq!(far, 6 * 65_535u64 * 65_535);
assert_eq!(brute_force(&points, &query), Some(found));
}
#[test]
fn a_point_is_sixteen_bytes() {
assert_eq!(std::mem::size_of::<Point>(), 16);
}
}