Skip to main content

molgfx_math/bounds/
bvh.rs

1//! Deterministic linear bounding-volume hierarchy.
2//!
3//! Construction radix-sorts `n` primitive bounds by a 63-bit Morton key in
4//! `O(n)` and emits `O(n)` compact nodes. Traversal is `O(log n + hits)` for
5//! coherent spatial data and accepts caller-owned scratch vectors, so repeated
6//! queries allocate nothing. The 32-byte node is directly uploadable to storage
7//! buffers.
8//!
9//! The key is 21 bits per axis rather than ten. Ten bits addresses a 1024³
10//! lattice, so past roughly ten million primitives a dense region collapses
11//! into one cell, every key there is equal, and the split degenerates to a
12//! median that ignores geometry — the hierarchy stops separating what it is
13//! built to separate. Twenty-one bits addresses a 2,097,152³ lattice, which
14//! keeps neighbouring primitives distinguishable at the scales this engine
15//! claims.
16
17use super::BvhBuildError;
18use super::source::BvhSource;
19use crate::{Aabb, Vec3};
20use std::ops::Range;
21
22use super::build;
23use build::{checked_index, checked_u32};
24
25#[cfg(test)]
26#[path = "bvh_tests.rs"]
27mod tests;
28
29#[path = "bvh_query.rs"]
30mod query;
31#[cfg(test)]
32use query::point_box_distance_squared;
33
34pub(super) const LEAF_SIZE: usize = 4;
35
36/// Cells per axis in the Morton lattice: 21 bits, minus one so the top index is
37/// reachable from a unit coordinate of exactly one.
38pub(super) const MORTON_LEVELS: u32 = (1 << 21) - 1;
39
40/// Primitive count above which the key sort partitions across cores. Below it
41/// the pool would cost more than the sort.
42const PARALLEL_SORT_THRESHOLD: usize = 1 << 16;
43pub(super) const COUNT_SHIFT: u32 = 29;
44pub(super) const INDEX_MASK: u32 = (1 << COUNT_SHIFT) - 1;
45pub(super) const MAX_BRANCH_LEFT: u32 = INDEX_MASK - 1;
46
47/// One compact BVH node. Internal children are adjacent; leaves address the
48/// hierarchy's primitive-index array.
49#[repr(C, align(16))]
50#[derive(Clone, Copy, PartialEq, Debug, Default, bytemuck::Pod, bytemuck::Zeroable)]
51pub struct BvhNode {
52    /// Minimum corner followed by left-child or first-primitive index.
53    pub min_left: [f32; 4],
54    /// Maximum corner followed by the largest primitive half-extent.
55    pub max_radius: [f32; 4],
56}
57
58impl BvhNode {
59    // `build_range` validates both offsets before reaching this encoder.
60    #[inline]
61    pub(super) fn leaf(bounds: Aabb, first: u32, count: u32, max_radius: f32) -> Self {
62        let metadata = first | (count << COUNT_SHIFT);
63        Self {
64            min_left: [
65                bounds.min.x,
66                bounds.min.y,
67                bounds.min.z,
68                f32::from_bits(metadata),
69            ],
70            max_radius: [bounds.max.x, bounds.max.y, bounds.max.z, max_radius],
71        }
72    }
73
74    // `build_range` validates the child offset before reaching this encoder.
75    #[inline]
76    pub(super) fn branch(bounds: Aabb, left: u32, max_radius: f32) -> Self {
77        Self {
78            min_left: [
79                bounds.min.x,
80                bounds.min.y,
81                bounds.min.z,
82                f32::from_bits(left),
83            ],
84            max_radius: [bounds.max.x, bounds.max.y, bounds.max.z, max_radius],
85        }
86    }
87
88    /// Node bounds.
89    #[must_use]
90    #[inline]
91    pub fn bounds(self) -> Aabb {
92        Aabb::new(
93            Vec3::from_array([self.min_left[0], self.min_left[1], self.min_left[2]]),
94            Vec3::from_array([self.max_radius[0], self.max_radius[1], self.max_radius[2]]),
95        )
96    }
97
98    /// Largest primitive half-extent below this node.
99    #[must_use]
100    #[inline]
101    pub fn maximum_radius(self) -> f32 {
102        self.max_radius[3]
103    }
104
105    /// True for a leaf node.
106    #[must_use]
107    #[inline]
108    pub fn is_leaf(self) -> bool {
109        self.min_left[3].to_bits() >> COUNT_SHIFT != 0
110    }
111
112    /// Adjacent child indices for an internal node.
113    #[must_use]
114    #[inline]
115    pub fn children(self) -> Option<(u32, u32)> {
116        if self.is_leaf() {
117            return None;
118        }
119        let left = self.min_left[3].to_bits() & INDEX_MASK;
120        Some((left, left + 1))
121    }
122
123    /// Range into `Bvh::primitive_indices` for a leaf.
124    #[must_use]
125    #[inline]
126    pub fn primitive_range(self) -> Option<Range<u32>> {
127        if !self.is_leaf() {
128            return None;
129        }
130        let metadata = self.min_left[3].to_bits();
131        let first = metadata & INDEX_MASK;
132        let count = metadata >> COUNT_SHIFT;
133        Some(first..first + count)
134    }
135}
136
137/// Flat hierarchy and stable source-primitive permutation.
138#[derive(Clone, PartialEq, Debug, Default)]
139pub struct Bvh {
140    /// Root-first compact nodes.
141    pub nodes: Vec<BvhNode>,
142    /// Source primitive indices addressed by leaves.
143    pub primitive_indices: Vec<u32>,
144    /// Per-node stackless-traversal successor, parallel to `nodes`.
145    pub escape: Vec<u32>,
146}
147
148/// Reusable construction workspace. Keeping it beside a dynamic structure
149/// makes trajectory-driven hierarchy rebuilds allocation-free after warmup.
150#[derive(Clone, Debug, Default)]
151pub struct BvhBuildScratch {
152    entries: Vec<MortonEntry>,
153    radix: Vec<MortonEntry>,
154    stack: Vec<u32>,
155}
156
157impl Bvh {
158    /// Escape value that ends a stackless traversal.
159    pub const ESCAPE_END: u32 = u32::MAX;
160
161    /// Builds a hierarchy over finite, non-empty bounds. Invalid primitives
162    /// are omitted instead of poisoning the structure-wide bounds.
163    ///
164    /// # Errors
165    ///
166    /// Returns [`BvhBuildError`] when a source row, node or primitive offset
167    /// cannot fit the compact GPU representation.
168    pub fn build<S: BvhSource + ?Sized>(source: &S) -> Result<Self, BvhBuildError> {
169        let mut hierarchy = Self::default();
170        hierarchy.rebuild(source, &mut BvhBuildScratch::default())?;
171        Ok(hierarchy)
172    }
173
174    /// Rebuilds in existing storage, retaining node, permutation and sort
175    /// capacities for moving coordinates.
176    ///
177    /// # Errors
178    ///
179    /// Returns [`BvhBuildError`] when a source row, node or primitive offset
180    /// cannot fit the compact GPU representation.
181    pub fn rebuild<S: BvhSource + ?Sized>(
182        &mut self,
183        source: &S,
184        scratch: &mut BvhBuildScratch,
185    ) -> Result<(), BvhBuildError> {
186        self.nodes.clear();
187        self.primitive_indices.clear();
188        self.escape.clear();
189        scratch.entries.clear();
190        let count = checked_u32("BVH source primitive", source.len(), u32::MAX)?;
191        // Union is componentwise minimum and maximum: exact, associative and
192        // commutative, so splitting it across cores cannot move the result.
193        let scene_bounds = build::scene_bounds(source, count);
194        if scene_bounds.is_empty() {
195            return Ok(());
196        }
197        build::extend_keys(&mut scratch.entries, source, count, scene_bounds);
198        checked_index("BVH primitive table", scratch.entries.len(), INDEX_MASK)?;
199        // Entries were pushed in ascending source order and the sort is stable,
200        // so ordering by code alone reproduces the `(code, source)` order the
201        // hierarchy is specified against without widening the key.
202        super::radix::sort_by_code(
203            &mut scratch.entries,
204            &mut scratch.radix,
205            PARALLEL_SORT_THRESHOLD,
206        );
207        let node_capacity = scratch.entries.len().checked_mul(2).ok_or(BvhBuildError {
208            resource: "BVH node capacity",
209            index: u64::MAX,
210            maximum: INDEX_MASK,
211        })?;
212        checked_index("BVH node table", node_capacity, INDEX_MASK)?;
213        self.nodes.reserve(node_capacity);
214        self.primitive_indices.reserve(scratch.entries.len());
215        self.nodes.push(BvhNode::default());
216        build::build_range(&scratch.entries, source, 0, scratch.entries.len(), 0, self)?;
217        self.thread_escapes(&mut scratch.stack);
218        Ok(())
219    }
220
221    /// Refits leaf and branch bounds without changing topology or primitive
222    /// order. The work is `O(resident primitives + nodes)` and reuses every
223    /// allocation retained by the hierarchy.
224    ///
225    /// # Errors
226    ///
227    /// Returns [`BvhBuildError`] when a primitive used by the existing
228    /// topology is absent or no longer has finite, ordered bounds.
229    pub fn refit<S: BvhSource + ?Sized>(&mut self, source: &S) -> Result<(), BvhBuildError> {
230        for node_index in 0..self.nodes.len() {
231            let Some(node) = self.nodes.get(node_index).copied() else {
232                continue;
233            };
234            let Some(range) = node.primitive_range() else {
235                continue;
236            };
237            let first = range.start;
238            let count = range.end - range.start;
239            let mut aggregate = Aabb::EMPTY;
240            let mut maximum_radius = 0.0f32;
241            for offset in range {
242                let Some(&source_row) = self.primitive_indices.get(offset as usize) else {
243                    return Err(refit_error("BVH primitive table", u64::from(offset)));
244                };
245                if source_row as usize >= source.len() {
246                    return Err(refit_error("BVH source primitive", u64::from(source_row)));
247                }
248                let bound = source.bound(source_row);
249                if !build::valid_bound(&bound) {
250                    return Err(refit_error("BVH source bounds", u64::from(source_row)));
251                }
252                aggregate = aggregate.union(&bound);
253                maximum_radius = maximum_radius.max(bound.half_extents().max_element());
254            }
255            self.nodes[node_index] = BvhNode::leaf(aggregate, first, count, maximum_radius);
256        }
257        for node_index in (0..self.nodes.len()).rev() {
258            let Some(node) = self.nodes.get(node_index).copied() else {
259                continue;
260            };
261            let Some((left, right)) = node.children() else {
262                continue;
263            };
264            let Some(left_node) = self.nodes.get(left as usize).copied() else {
265                return Err(refit_error("BVH left child", u64::from(left)));
266            };
267            let Some(right_node) = self.nodes.get(right as usize).copied() else {
268                return Err(refit_error("BVH right child", u64::from(right)));
269            };
270            self.nodes[node_index] = BvhNode::branch(
271                left_node.bounds().union(&right_node.bounds()),
272                left,
273                left_node.maximum_radius().max(right_node.maximum_radius()),
274            );
275        }
276        Ok(())
277    }
278
279    /// Threads the finished hierarchy for stackless traversal in `O(nodes)`.
280    ///
281    /// `escape[n]` names the node a walk moves to once it has skipped or
282    /// finished `n`'s subtree. A walk that hits a node descends into its first
283    /// child and otherwise follows the link, so no per-ray stack is needed —
284    /// which is what makes the hierarchy usable from a fragment shader, where
285    /// a stack would cost registers on every lane.
286    ///
287    /// The left child escapes into its sibling and the sibling inherits the
288    /// parent's escape, so the chain leaves any subtree exactly once. The root
289    /// escapes to `ESCAPE_END`, which terminates the walk.
290    fn thread_escapes(&mut self, stack: &mut Vec<u32>) {
291        self.escape.resize(self.nodes.len(), Self::ESCAPE_END);
292        stack.clear();
293        if self.nodes.is_empty() {
294            return;
295        }
296        stack.push(0);
297        while let Some(index) = stack.pop() {
298            let Some(node) = self.nodes.get(index as usize).copied() else {
299                continue;
300            };
301            let Some((left, right)) = node.children() else {
302                continue;
303            };
304            let Some(&after) = self.escape.get(index as usize) else {
305                continue;
306            };
307            if let Some(slot) = self.escape.get_mut(left as usize) {
308                *slot = right;
309            }
310            if let Some(slot) = self.escape.get_mut(right as usize) {
311                *slot = after;
312            }
313            stack.push(right);
314            stack.push(left);
315        }
316    }
317}
318
319#[inline]
320fn refit_error(resource: &'static str, index: u64) -> BvhBuildError {
321    BvhBuildError {
322        resource,
323        index,
324        maximum: u32::MAX,
325    }
326}
327
328/// One primitive's spatial key and the row it came from.
329#[derive(Clone, Copy, Debug)]
330pub(super) struct MortonEntry {
331    pub(super) code: u64,
332    pub(super) source: u32,
333}
334
335impl MortonEntry {
336    /// Fill value for scratch storage the sort is about to overwrite. It is
337    /// never observed: every slot is written before it is read back.
338    pub(super) const PLACEHOLDER: Self = Self { code: 0, source: 0 };
339}