use crate::bbox::Interval;
use crate::data_source::DataSource;
use crate::node::Node;
use crate::scalar::{IndexType, Scalar};
pub(crate) fn init_vind<Idx: IndexType>(n: usize) -> Vec<Idx> {
assert!(
n <= u32::MAX as usize,
"init_vind: point count {n} exceeds u32::MAX (leaf offsets are u32)"
);
(0..n).map(Idx::from_usize).collect()
}
#[inline]
pub(crate) fn cpp_min<T: Scalar>(a: T, b: T) -> T {
if b < a {
b
} else {
a
}
}
#[inline]
pub(crate) fn cpp_max<T: Scalar>(a: T, b: T) -> T {
if a < b {
b
} else {
a
}
}
fn compute_min_max<T, DS, Idx>(ds: &DS, ind: &[Idx], dim: usize) -> (T, T)
where
T: Scalar,
DS: DataSource<T> + ?Sized,
Idx: IndexType,
{
let count = ind.len();
let mut local_min = ds.point_component(ind[0].to_usize(), dim);
let mut local_max = local_min;
const UNROLL: usize = 4;
let mut k: usize = 1;
while k + UNROLL <= count {
let v0 = ds.point_component(ind[k].to_usize(), dim);
let v1 = ds.point_component(ind[k + 1].to_usize(), dim);
let v2 = ds.point_component(ind[k + 2].to_usize(), dim);
let v3 = ds.point_component(ind[k + 3].to_usize(), dim);
local_min = cpp_min(cpp_min(cpp_min(cpp_min(local_min, v0), v1), v2), v3);
local_max = cpp_max(cpp_max(cpp_max(cpp_max(local_max, v0), v1), v2), v3);
k += UNROLL;
}
while k < count {
let val = ds.point_component(ind[k].to_usize(), dim);
local_min = cpp_min(local_min, val);
local_max = cpp_max(local_max, val);
k += 1;
}
(local_min, local_max)
}
pub(crate) fn plane_split<T, DS, Idx>(
ds: &DS,
ind: &mut [Idx],
cutfeat: usize,
cutval: T,
) -> (usize, usize)
where
T: Scalar,
DS: DataSource<T> + ?Sized,
Idx: IndexType,
{
let count = ind.len();
let mut left: usize = 0;
let mut mid: usize = 0;
let mut right: isize = count as isize - 1;
while (mid as isize) <= right {
let val = ds.point_component(ind[mid].to_usize(), cutfeat);
if val < cutval {
ind.swap(left, mid);
left += 1;
mid += 1;
} else if val > cutval {
ind.swap(mid, right as usize);
right -= 1;
} else {
mid += 1;
}
}
(left, mid)
}
#[allow(clippy::needless_range_loop)]
pub(crate) fn middle_split<T, DS, Idx>(
ds: &DS,
dim: usize,
ind: &mut [Idx],
bbox: &[Interval<T>],
) -> (usize, u32, T)
where
T: Scalar,
DS: DataSource<T> + ?Sized,
Idx: IndexType,
{
let count = ind.len();
let eps = T::from_f64(0.00001);
let one = T::from_f64(1.0);
let two = T::from_f64(2.0);
let mut max_span = bbox[0].high - bbox[0].low;
for d in 1..dim {
let span = bbox[d].high - bbox[d].low;
if span > max_span {
max_span = span;
}
}
let mut cutfeat: usize = 0;
let mut max_spread = T::from_f64(-1.0);
let mut min_elem = T::default();
let mut max_elem = T::default();
let threshold = (one - eps) * max_span;
for d in 0..dim {
if bbox[d].high - bbox[d].low < threshold {
continue;
}
let (local_min, local_max) = compute_min_max(ds, ind, d);
let spread = local_max - local_min;
if spread > max_spread {
cutfeat = d;
max_spread = spread;
min_elem = local_min;
max_elem = local_max;
}
}
let mut split_val = (bbox[cutfeat].low + bbox[cutfeat].high) / two;
if split_val < min_elem {
split_val = min_elem;
}
if split_val > max_elem {
split_val = max_elem;
}
let cutval = split_val;
let (lim1, lim2) = plane_split(ds, ind, cutfeat, cutval);
let half = count / 2;
let index = if lim1 > half {
lim1
} else if lim2 < half {
lim2
} else {
half
};
(index, cutfeat as u32, cutval)
}
pub(crate) struct SubtreeBuilder<'a, T: Scalar, DS: DataSource<T> + ?Sized, Idx: Copy> {
pub ds: &'a DS,
pub dim: usize,
pub leaf_max_size: usize,
pub base: u32,
pub vind: &'a mut [Idx],
pub arena: &'a mut Vec<Node<T>>,
}
enum Work<T> {
Build {
left: usize,
right: usize,
bbox: Vec<Interval<T>>,
},
Finalize { node: u32, cutfeat: usize },
}
impl<'a, T, DS, Idx> SubtreeBuilder<'a, T, DS, Idx>
where
T: Scalar,
DS: DataSource<T> + ?Sized,
Idx: IndexType,
{
#[allow(clippy::needless_range_loop)]
pub(crate) fn build(&mut self, bbox: &mut [Interval<T>]) -> u32 {
assert!(
!self.vind.is_empty(),
"SubtreeBuilder::build called with empty vind"
);
debug_assert_eq!(bbox.len(), self.dim);
let mut bbox_pool: Vec<Vec<Interval<T>>> = Vec::new();
let mut work: Vec<Work<T>> = vec![Work::Build {
left: 0,
right: self.vind.len(),
bbox: bbox.to_vec(),
}];
let mut results: Vec<(u32, Vec<Interval<T>>)> = Vec::new();
while let Some(item) = work.pop() {
match item {
Work::Build {
left,
right,
bbox: mut sub_bbox,
} => {
let count = right - left;
if count <= self.leaf_max_size {
debug_assert!(left <= u32::MAX as usize && right <= u32::MAX as usize);
let node_idx = self.arena.len() as u32;
self.arena.push(Node::leaf(
self.base + left as u32,
self.base + right as u32,
));
for d in 0..self.dim {
let v = self.ds.point_component(self.vind[left].to_usize(), d);
sub_bbox[d] = Interval { low: v, high: v };
}
for k in (left + 1)..right {
for d in 0..self.dim {
let v = self.ds.point_component(self.vind[k].to_usize(), d);
if sub_bbox[d].low > v {
sub_bbox[d].low = v;
}
if sub_bbox[d].high < v {
sub_bbox[d].high = v;
}
}
}
results.push((node_idx, sub_bbox));
} else {
let (split_index, cutfeat, cutval) =
middle_split(self.ds, self.dim, &mut self.vind[left..right], &sub_bbox);
let cutfeat = cutfeat as usize;
let node_idx = self.arena.len() as u32;
self.arena
.push(Node::split(cutfeat as u32, T::default(), T::default()));
let mut left_bbox = match bbox_pool.pop() {
Some(mut buf) => {
buf.clear();
buf.extend_from_slice(&sub_bbox);
buf
}
None => sub_bbox.clone(),
};
left_bbox[cutfeat].high = cutval;
let mut right_bbox = sub_bbox;
right_bbox[cutfeat].low = cutval;
work.push(Work::Finalize {
node: node_idx,
cutfeat,
});
work.push(Work::Build {
left: left + split_index,
right,
bbox: right_bbox,
});
work.push(Work::Build {
left,
right: left + split_index,
bbox: left_bbox,
});
}
}
Work::Finalize { node, cutfeat } => {
let (right_node, right_bbox) = results.pop().expect("right result missing");
let (left_node, mut left_bbox) = results.pop().expect("left result missing");
self.arena[node as usize].set_children(left_node, right_node);
self.arena[node as usize]
.set_div_bounds(left_bbox[cutfeat].high, right_bbox[cutfeat].low);
for d in 0..self.dim {
left_bbox[d] = Interval {
low: cpp_min(left_bbox[d].low, right_bbox[d].low),
high: cpp_max(left_bbox[d].high, right_bbox[d].high),
};
}
bbox_pool.push(right_bbox);
results.push((node, left_bbox));
}
}
}
let (root_node, root_bbox) = results.pop().expect("build produced no result");
debug_assert!(results.is_empty());
bbox.copy_from_slice(&root_bbox);
root_node
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::bbox::compute_bounding_box;
fn build_full_tree<const N: usize>(
points: &[[f64; N]],
leaf_max_size: usize,
) -> (Vec<Node<f64>>, Vec<u32>, u32, Vec<Interval<f64>>) {
let dim = N;
let n = points.len();
let mut vind: Vec<u32> = init_vind(n);
let mut bbox = vec![
Interval {
low: 0.0,
high: 0.0
};
dim
];
compute_bounding_box(&points, dim, &mut bbox);
let mut arena = Vec::new();
let root = {
let mut builder = SubtreeBuilder {
ds: &points,
dim,
leaf_max_size,
base: 0,
vind: &mut vind,
arena: &mut arena,
};
builder.build(&mut bbox)
};
(arena, vind, root, bbox)
}
fn check_tree<const N: usize>(
points: &[[f64; N]],
arena: &[Node<f64>],
vind: &[u32],
root: u32,
leaf_max_size: usize,
) {
let dim = N;
let n = points.len();
assert!(
arena.len() <= 2 * n,
"node count {} exceeds 2n ({})",
arena.len(),
2 * n
);
enum St {
Visit(u32),
Process(u32),
}
let mut seen = vec![false; n];
let mut stack = vec![St::Visit(root)];
let mut results: Vec<Vec<Interval<f64>>> = Vec::new();
while let Some(item) = stack.pop() {
match item {
St::Visit(idx) => {
let node = &arena[idx as usize];
if node.is_leaf() {
let (l, r) = node.leaf_range();
assert!(
r - l <= leaf_max_size,
"leaf [{l},{r}) has {} points > leaf_max_size {leaf_max_size}",
r - l
);
let mut leaf_bbox = vec![
Interval {
low: 0.0,
high: 0.0
};
dim
];
for (i, k) in (l..r).enumerate() {
let pt = vind[k] as usize;
assert!(!seen[pt], "point {pt} appears in more than one leaf");
seen[pt] = true;
for d in 0..dim {
let v = points[pt][d];
if i == 0 {
leaf_bbox[d] = Interval { low: v, high: v };
} else {
if v < leaf_bbox[d].low {
leaf_bbox[d].low = v;
}
if v > leaf_bbox[d].high {
leaf_bbox[d].high = v;
}
}
}
}
results.push(leaf_bbox);
} else {
let (c1, c2) = node.children();
stack.push(St::Process(idx));
stack.push(St::Visit(c2));
stack.push(St::Visit(c1));
}
}
St::Process(idx) => {
let node = &arena[idx as usize];
let cutfeat = node.split_dim();
let right_bbox = results.pop().expect("missing right subtree bbox");
let left_bbox = results.pop().expect("missing left subtree bbox");
assert_eq!(
left_bbox[cutfeat].high,
node.div_low(),
"node {idx}: div_low != max of left subtree's coord[{cutfeat}]"
);
assert_eq!(
right_bbox[cutfeat].low,
node.div_high(),
"node {idx}: div_high != min of right subtree's coord[{cutfeat}]"
);
let mut combined = vec![
Interval {
low: 0.0,
high: 0.0
};
dim
];
for d in 0..dim {
combined[d] = Interval {
low: cpp_min(left_bbox[d].low, right_bbox[d].low),
high: cpp_max(left_bbox[d].high, right_bbox[d].high),
};
}
results.push(combined);
}
}
}
assert!(
seen.iter().all(|&b| b),
"not every point index appears in a leaf"
);
assert_eq!(results.len(), 1);
}
fn max_depth(arena: &[Node<f64>], root: u32) -> usize {
let mut stack = vec![(root, 1usize)];
let mut max_d = 0usize;
while let Some((idx, d)) = stack.pop() {
if d > max_d {
max_d = d;
}
let node = &arena[idx as usize];
if !node.is_leaf() {
let (c1, c2) = node.children();
stack.push((c1, d + 1));
stack.push((c2, d + 1));
}
}
max_d
}
fn assert_plane_split_postcondition(values: &[f64], cutval: f64) {
let points: Vec<[f64; 1]> = values.iter().map(|&v| [v]).collect();
let mut ind: Vec<u32> = (0..values.len() as u32).collect();
let original: Vec<u32> = ind.clone();
let (lim1, lim2) = plane_split(&points.as_slice(), &mut ind, 0, cutval);
for &i in &ind[..lim1] {
assert!(points[i as usize][0] < cutval, "ind[..lim1] violated");
}
for &i in &ind[lim1..lim2] {
assert_eq!(points[i as usize][0], cutval, "ind[lim1..lim2] violated");
}
for &i in &ind[lim2..] {
assert!(points[i as usize][0] > cutval, "ind[lim2..] violated");
}
let mut sorted_ind = ind.clone();
sorted_ind.sort_unstable();
let mut sorted_orig = original.clone();
sorted_orig.sort_unstable();
assert_eq!(sorted_ind, sorted_orig, "ind is not a permutation of input");
}
#[test]
fn plane_split_postcondition_with_duplicates() {
assert_plane_split_postcondition(&[5.0, 3.0, 5.0, 1.0, 5.0, 8.0, 1.0, 5.0], 5.0);
}
#[test]
fn plane_split_postcondition_all_less() {
assert_plane_split_postcondition(&[1.0, 2.0, 3.0, 4.0], 10.0);
}
#[test]
fn plane_split_postcondition_all_greater() {
assert_plane_split_postcondition(&[11.0, 12.0, 13.0, 14.0], 10.0);
}
#[test]
fn plane_split_postcondition_all_equal() {
assert_plane_split_postcondition(&[7.0, 7.0, 7.0, 7.0, 7.0], 7.0);
}
#[test]
fn plane_split_postcondition_single_element() {
assert_plane_split_postcondition(&[42.0], 42.0);
}
#[test]
fn plane_split_exact_permutation_hand_simulated() {
let values = [5.0f64, 3.0, 5.0, 1.0, 5.0, 8.0, 1.0, 5.0];
let points: Vec<[f64; 1]> = values.iter().map(|&v| [v]).collect();
let mut ind: Vec<u32> = (0..8).collect();
let (lim1, lim2) = plane_split(&points.as_slice(), &mut ind, 0, 5.0);
assert_eq!(ind, vec![1, 3, 6, 0, 4, 7, 2, 5]);
assert_eq!((lim1, lim2), (3, 7));
}
#[test]
fn middle_split_picks_second_dim_when_spans_tie_but_spread_larger() {
let points: &[[f64; 2]] = &[[2.0, 1.0], [8.0, 1.0], [2.0, 9.0], [8.0, 9.0]];
let mut ind: Vec<u32> = (0..4).collect();
let bbox = [
Interval {
low: 0.0,
high: 10.0,
},
Interval {
low: 0.0,
high: 10.0,
},
];
let (index, cutfeat, cutval) = middle_split(&points, 2, &mut ind, &bbox);
assert_eq!(cutfeat, 1, "should pick dim 1 (larger actual spread)");
assert_eq!(cutval, 5.0);
assert_eq!(index, 2);
}
#[test]
fn middle_split_clamps_split_val_to_min_elem() {
let points: &[[f64; 1]] = &[[80.0], [81.0], [82.0]];
let mut ind: Vec<u32> = (0..3).collect();
let bbox = [Interval {
low: 0.0,
high: 100.0,
}];
let (index, cutfeat, cutval) = middle_split(&points, 1, &mut ind, &bbox);
assert_eq!(cutfeat, 0);
assert_eq!(cutval, 80.0);
assert_eq!(index, 1);
}
#[test]
fn middle_split_lim1_greater_than_half_selects_lim1() {
let points: &[[f64; 1]] = &[[1.0], [1.0], [1.0], [1.0], [5.0], [9.0]];
let mut ind: Vec<u32> = (0..6).collect();
let bbox = [Interval {
low: 1.0,
high: 9.0,
}];
let (index, cutfeat, cutval) = middle_split(&points, 1, &mut ind, &bbox);
assert_eq!(cutfeat, 0);
assert_eq!(cutval, 5.0);
assert_eq!(index, 4);
}
#[test]
fn middle_split_lim2_less_than_half_selects_lim2() {
let points: &[[f64; 1]] = &[[1.0], [9.0], [9.0], [9.0], [9.0], [9.0]];
let mut ind: Vec<u32> = (0..6).collect();
let bbox = [Interval {
low: 1.0,
high: 9.0,
}];
let (index, cutfeat, cutval) = middle_split(&points, 1, &mut ind, &bbox);
assert_eq!(cutfeat, 0);
assert_eq!(cutval, 5.0);
assert_eq!(index, 1);
}
#[test]
fn leaf_equal_boundary_is_a_single_leaf() {
let points: Vec<[f64; 2]> = (0..10).map(|i| [i as f64, (i * 2) as f64]).collect();
let (arena, vind, root, _bbox) = build_full_tree(&points, 10);
assert_eq!(arena.len(), 1);
assert!(arena[root as usize].is_leaf());
assert_eq!(arena[root as usize].leaf_range(), (0, 10));
assert_eq!(vind.len(), 10);
}
struct Lcg(u64);
impl Lcg {
fn next_f64(&mut self) -> f64 {
self.0 = self
.0
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
((self.0 >> 11) as f64) / ((1u64 << 53) as f64)
}
}
#[test]
fn invariant_holds_on_100_uniform_random_points_dim3() {
let mut rng = Lcg(0xC0FFEE_u64);
let points: Vec<[f64; 3]> = (0..100)
.map(|_| {
[
rng.next_f64() * 100.0,
rng.next_f64() * 100.0,
rng.next_f64() * 100.0,
]
})
.collect();
let (arena, vind, root, _bbox) = build_full_tree(&points, 10);
check_tree(&points, &arena, &vind, root, 10);
}
#[test]
fn all_identical_points_build_balanced_tree() {
let points: Vec<[f64; 3]> = vec![[3.5, -2.0, 7.25]; 1000];
let (arena, vind, root, _bbox) = build_full_tree(&points, 10);
check_tree(&points, &arena, &vind, root, 10);
for node in &arena {
if node.is_leaf() {
let (l, r) = node.leaf_range();
assert!(r - l <= 10);
} else {
assert_eq!(node.div_low(), node.div_high());
}
}
let depth = max_depth(&arena, root);
assert!(depth <= 12, "tree too deep for balanced splits: {depth}");
assert_eq!(vind.len(), 1000);
}
#[test]
fn leaf_max_size_one_gives_singleton_leaves() {
let points: Vec<[f64; 2]> = (0..64).map(|i| [i as f64, (i * i) as f64 % 37.0]).collect();
let (arena, vind, root, _bbox) = build_full_tree(&points, 1);
check_tree(&points, &arena, &vind, root, 1);
for node in &arena {
if node.is_leaf() {
let (l, r) = node.leaf_range();
assert_eq!(r - l, 1);
}
}
}
#[test]
fn leaf_max_size_at_least_n_gives_single_node() {
let points: Vec<[f64; 2]> = (0..64).map(|i| [i as f64, (i * i) as f64 % 37.0]).collect();
let (arena, _vind, root, _bbox) = build_full_tree(&points, 1000);
assert_eq!(arena.len(), 1);
assert!(arena[root as usize].is_leaf());
}
#[test]
fn single_point_tree() {
let points: Vec<[f64; 3]> = vec![[1.5, -2.5, 3.5]];
let (arena, vind, root, bbox) = build_full_tree(&points, 10);
assert_eq!(arena.len(), 1);
assert!(arena[root as usize].is_leaf());
assert_eq!(vind, vec![0]);
assert_eq!(
bbox[0],
Interval {
low: 1.5,
high: 1.5
}
);
assert_eq!(
bbox[1],
Interval {
low: -2.5,
high: -2.5
}
);
assert_eq!(
bbox[2],
Interval {
low: 3.5,
high: 3.5
}
);
}
#[test]
fn build_tightens_a_loose_input_bbox() {
let points: Vec<[f64; 2]> = vec![[1.0, 5.0], [2.0, 6.0], [3.0, 4.0], [1.5, 5.5]];
let mut vind: Vec<u32> = init_vind(points.len());
let mut bbox = vec![
Interval {
low: -100.0,
high: 100.0,
},
Interval {
low: -100.0,
high: 100.0,
},
];
let mut arena = Vec::new();
let root = {
let mut builder = SubtreeBuilder {
ds: &points.as_slice(),
dim: 2,
leaf_max_size: 10,
base: 0,
vind: &mut vind,
arena: &mut arena,
};
builder.build(&mut bbox)
};
assert!(arena[root as usize].is_leaf());
assert_eq!(
bbox[0],
Interval {
low: 1.0,
high: 3.0
}
);
assert_eq!(
bbox[1],
Interval {
low: 4.0,
high: 6.0
}
);
}
#[test]
#[ignore]
fn heavy_exponential_build_1m() {
let n = 1_000_000usize;
let mut spine: Vec<f64> = Vec::new();
let mut v = 2f64.powi(1023);
while v > 0.0 && spine.len() < n - 1 {
spine.push(v);
v /= 2.0;
}
let mut values = spine;
values.resize(n, 0.0);
let points: Vec<[f64; 1]> = values.into_iter().map(|v| [v]).collect();
let (arena, vind, root, _bbox) = build_full_tree(&points, 10);
let depth = max_depth(&arena, root);
assert!(
depth > 2_000,
"expected a deeply degenerate tree to exercise stack safety, got depth {depth}"
);
check_tree(&points, &arena, &vind, root, 10);
}
#[test]
#[ignore]
fn heavy_exponential_build_1m_dim8() {
const DIM: usize = 8;
let n = 1_000_000usize;
let mut ladder: Vec<f64> = Vec::new();
let mut v = 2f64.powi(1023);
while v > 0.0 {
ladder.push(v);
v /= 2.0;
}
let spine_len = (ladder.len() * DIM).min(n - 1);
let mut points: Vec<[f64; DIM]> = Vec::with_capacity(n);
for k in 0..spine_len {
let d = k % DIM;
let m = k / DIM;
let mut coords = [0.0f64; DIM];
coords[d] = ladder[m];
points.push(coords);
}
points.resize(n, [0.0f64; DIM]);
let (arena, vind, root, _bbox) = build_full_tree(&points, 10);
let depth = max_depth(&arena, root);
assert!(
depth > 10_000,
"expected round-robin degenerate tree to exceed the single-axis ceiling, got depth {depth}"
);
check_tree(&points, &arena, &vind, root, 10);
}
}