use crate::bbox::Interval;
use crate::build::{cpp_max, cpp_min, middle_split, SubtreeBuilder};
use crate::data_source::DataSource;
use crate::node::Node;
use crate::scalar::{IndexType, Scalar};
const PARALLEL_CUTOFF: usize = 4096;
fn ceil_log2(x: usize) -> u32 {
if x <= 1 {
0
} else {
(usize::BITS) - (x - 1).leading_zeros()
}
}
pub(crate) fn depth_budget_for(n: usize) -> usize {
let ratio = (n / PARALLEL_CUTOFF).max(2);
let budget = 2 * ceil_log2(ratio) as usize + 8;
budget.min(64)
}
pub(crate) fn build_subtree_parallel<T, DS, Idx>(
ds: &DS,
dim: usize,
leaf_max_size: usize,
base: u32,
vind: &mut [Idx],
bbox: &mut [Interval<T>],
depth_budget: usize,
) -> (u32, Vec<Node<T>>)
where
T: Scalar,
DS: DataSource<T> + ?Sized + Sync,
Idx: IndexType,
{
let n = vind.len();
if n <= PARALLEL_CUTOFF || depth_budget == 0 || n <= leaf_max_size {
let mut arena = Vec::new();
let root = {
let mut builder = SubtreeBuilder {
ds,
dim,
leaf_max_size,
base,
vind,
arena: &mut arena,
};
builder.build(bbox)
};
return (root, arena);
}
let (split_index, cutfeat, cutval) = middle_split(ds, dim, vind, bbox);
let cutfeat = cutfeat as usize;
let mut left_bbox = bbox.to_vec();
left_bbox[cutfeat].high = cutval;
let mut right_bbox = bbox.to_vec();
right_bbox[cutfeat].low = cutval;
let (lv, rv) = vind.split_at_mut(split_index);
let right_base = base + split_index as u32;
let ((left_root, mut left_arena), (right_root, mut right_arena)) = rayon::join(
|| {
build_subtree_parallel(
ds,
dim,
leaf_max_size,
base,
lv,
&mut left_bbox,
depth_budget - 1,
)
},
|| {
build_subtree_parallel(
ds,
dim,
leaf_max_size,
right_base,
rv,
&mut right_bbox,
depth_budget - 1,
)
},
);
let l_len = left_arena.len() as u32;
let mut arena = Vec::with_capacity(1 + left_arena.len() + right_arena.len());
arena.push(Node::split(cutfeat as u32, T::default(), T::default()));
for node in left_arena.iter_mut() {
if !node.is_leaf() {
node.offset_children(1);
}
}
arena.append(&mut left_arena);
for node in right_arena.iter_mut() {
if !node.is_leaf() {
node.offset_children(1 + l_len);
}
}
arena.append(&mut right_arena);
arena[0].set_children(1 + left_root, 1 + l_len + right_root);
arena[0].set_div_bounds(left_bbox[cutfeat].high, right_bbox[cutfeat].low);
for d in 0..dim {
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),
};
}
(0, arena)
}
pub(crate) fn build_tree_parallel<T, DS, Idx>(
ds: &DS,
dim: usize,
leaf_max_size: usize,
vind: &mut [Idx],
bbox: &mut [Interval<T>],
) -> Vec<Node<T>>
where
T: Scalar,
DS: DataSource<T> + ?Sized + Sync,
Idx: IndexType,
{
let depth_budget = depth_budget_for(vind.len());
let (root, arena) = build_subtree_parallel(ds, dim, leaf_max_size, 0, vind, bbox, depth_budget);
debug_assert_eq!(root, 0, "root is expected to always be arena index 0");
arena
}
#[cfg(all(test, feature = "parallel"))]
mod tests {
use super::*;
use crate::build::init_vind;
use crate::dim::ConstDim;
use crate::metric::L2;
use crate::result_set::{KnnResultSet, ResultSet};
use crate::search::{find_neighbors, SearchCtx};
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)
}
}
fn seeded_points<const N: usize>(seed: u64, n: usize, scale: f64) -> Vec<[f64; N]> {
let mut rng = Lcg(seed);
(0..n)
.map(|_| {
let mut p = [0.0f64; N];
for v in p.iter_mut() {
*v = rng.next_f64() * scale;
}
p
})
.collect()
}
fn build_sequential<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
];
crate::bbox::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 build_parallel<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
];
crate::bbox::compute_bounding_box(&points, dim, &mut bbox);
let arena = build_tree_parallel(&points, dim, leaf_max_size, &mut vind, &mut bbox);
(arena, vind, 0, bbox)
}
fn build_parallel_with_pool<const N: usize>(
points: &[[f64; N]],
leaf_max_size: usize,
n_threads: usize,
) -> (Vec<Node<f64>>, Vec<u32>, u32, Vec<Interval<f64>>) {
let pool = rayon::ThreadPoolBuilder::new()
.num_threads(n_threads)
.build()
.expect("rayon pool");
pool.install(|| build_parallel(points, leaf_max_size))
}
fn assert_builds_bit_identical<const N: usize>(points: &[[f64; N]], leaf_max_size: usize) {
let (seq_arena, seq_vind, seq_root, seq_bbox) = build_sequential(points, leaf_max_size);
let (auto_arena, auto_vind, auto_root, auto_bbox) = build_parallel(points, leaf_max_size);
assert_eq!(seq_vind, auto_vind, "vind mismatch (Auto)");
assert_eq!(seq_root, auto_root, "root mismatch (Auto)");
assert_eq!(seq_arena, auto_arena, "arena mismatch (Auto)");
assert_eq!(seq_bbox, auto_bbox, "root_bbox mismatch (Auto)");
let (t2_arena, t2_vind, t2_root, t2_bbox) =
build_parallel_with_pool(points, leaf_max_size, 2);
assert_eq!(seq_vind, t2_vind, "vind mismatch (Threads(2))");
assert_eq!(seq_root, t2_root, "root mismatch (Threads(2))");
assert_eq!(seq_arena, t2_arena, "arena mismatch (Threads(2))");
assert_eq!(seq_bbox, t2_bbox, "root_bbox mismatch (Threads(2))");
}
#[test]
fn bit_identical_to_sequential_5000_uniform_dim3() {
let points = seeded_points::<3>(0xC0FFEE_u64, 5000, 1000.0);
assert_builds_bit_identical(&points, 10);
}
#[test]
fn bit_identical_to_sequential_20000_uniform_dim3() {
let points = seeded_points::<3>(0xC0FFEE2_u64, 20_000, 1000.0);
assert!(
points.len() > 4 * PARALLEL_CUTOFF,
"test setup: n must exceed 4x PARALLEL_CUTOFF to force a composed (multi-level) merge"
);
assert_builds_bit_identical(&points, 10);
}
#[test]
fn query_equivalence_200_knn_queries_k10() {
let points = seeded_points::<3>(0xC0FFEE_u64, 5000, 1000.0);
let (seq_arena, seq_vind, _root, seq_bbox) = build_sequential(&points, 10);
let (auto_arena, auto_vind, _root2, auto_bbox) = build_parallel(&points, 10);
let seq_ctx = SearchCtx {
ds: &points.as_slice(),
metric: &L2,
dim: ConstDim::<3>,
nodes: &seq_arena,
vind: &seq_vind,
root_bbox: &seq_bbox,
};
let auto_ctx = SearchCtx {
ds: &points.as_slice(),
metric: &L2,
dim: ConstDim::<3>,
nodes: &auto_arena,
vind: &auto_vind,
root_bbox: &auto_bbox,
};
let mut rng = Lcg(0xA5A5_u64);
let params = crate::params::SearchParams::default();
for _ in 0..200 {
let query = [
rng.next_f64() * 1000.0,
rng.next_f64() * 1000.0,
rng.next_f64() * 1000.0,
];
let k = 10;
let mut seq_idx = vec![0u32; k];
let mut seq_dist = vec![0.0f64; k];
let mut seq_rs = KnnResultSet::<f64, u32>::new(&mut seq_idx, &mut seq_dist);
let mut scratch = vec![0.0f64; 3];
find_neighbors(
&seq_ctx,
&mut seq_rs,
&query,
¶ms,
&crate::filter::AcceptAll,
&mut scratch,
);
let seq_found = seq_rs.size();
let mut auto_idx = vec![0u32; k];
let mut auto_dist = vec![0.0f64; k];
let mut auto_rs = KnnResultSet::<f64, u32>::new(&mut auto_idx, &mut auto_dist);
let mut scratch2 = vec![0.0f64; 3];
find_neighbors(
&auto_ctx,
&mut auto_rs,
&query,
¶ms,
&crate::filter::AcceptAll,
&mut scratch2,
);
let auto_found = auto_rs.size();
assert_eq!(seq_found, auto_found);
assert_eq!(seq_idx, auto_idx);
assert_eq!(seq_dist, auto_dist);
}
}
#[test]
fn bit_identical_duplicate_heavy_5000_points_50_percent() {
let mut rng = Lcg(0xDEADBEEF_u64);
let unique_n = 2500usize;
let mut points: Vec<[f64; 3]> = (0..unique_n)
.map(|_| {
[
rng.next_f64() * 500.0,
rng.next_f64() * 500.0,
rng.next_f64() * 500.0,
]
})
.collect();
for i in 0..(5000 - unique_n) {
points.push(points[i % unique_n]);
}
assert_eq!(points.len(), 5000);
assert_builds_bit_identical(&points, 10);
}
#[test]
fn small_n_below_cutoff_matches_sequential_trivially() {
let points = seeded_points::<3>(0x1234_u64, 200, 50.0);
assert!(points.len() <= PARALLEL_CUTOFF);
assert_builds_bit_identical(&points, 10);
}
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 check_tree(
points: &[[f64; 1]],
arena: &[Node<f64>],
vind: &[u32],
root: u32,
leaf_max_size: usize,
) {
let dim = 1usize;
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 mismatch"
);
assert_eq!(
right_bbox[cutfeat].low,
node.div_high(),
"node {idx}: div_high mismatch"
);
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);
}
#[test]
#[ignore]
fn heavy_exponential_parallel_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_parallel(&points, 10);
let depth = max_depth(&arena, root);
assert!(
depth > 2_000,
"expected a deeply degenerate tree, got depth {depth}"
);
check_tree(&points, &arena, &vind, root, 10);
}
}