Skip to main content

henad_core/
spatial_hash.rs

1//! The counting-sort [`SpatialHash`] that agents use to find their neighbours, and the [`HashGrid`]
2//! cell geometry both backends share.
3
4use crate::authoring::model::field::Extent;
5use crate::authoring::primitives::space::{Boundary, axis_delta};
6
7/// Cell geometry on its own, for a caller that needs the grid without the buckets. A GPU model
8/// mirrors it into its step uniform so its query walks the same grid as the CPU sort.
9#[derive(Clone, Copy, Debug, PartialEq)]
10pub struct HashGrid {
11    /// Cells along x.
12    pub grid_w: u32,
13    /// Cells along y.
14    pub grid_h: u32,
15    /// World width one cell spans.
16    pub cell_w: f32,
17    /// World height one cell spans.
18    pub cell_h: f32,
19}
20
21/// Maximum number of cells one index can hold, on either backend.
22///
23/// The grid is `(world / cell_size)^2`, and the world size and the cell size are both model parameters, so a small
24/// cell on a large world requests a grid that no machine can hold. A coarser cell only makes a query scan candidates
25/// that it then rejects, whereas an unbounded grid requests gigabytes.
26pub const MAX_INDEX_CELLS: u64 = 1 << 22;
27
28impl HashGrid {
29    /// Returns the geometry that fits whole cells to the world.
30    ///
31    /// A query walks in cell index space, so cells all have to span the same distance or the wrap seam gets
32    /// under-covered.
33    ///
34    /// A grid that would hold more than [`MAX_INDEX_CELLS`] cells is coarsened by one factor on both
35    /// axes, and a `cell_size` that is not positive counts as 1. This is the one place the geometry is
36    /// decided. [`SpatialHash`] and its GPU counterpart both build from here, so neither backend can walk a
37    /// grid that the other backend did not sort.
38    pub fn new(extent: Extent, cell_size: f32) -> Self {
39        let cell_size = if cell_size > 0.0 { cell_size } else { 1.0 };
40        let mut grid_w = (extent.w / cell_size).floor().max(1.0) as u32;
41        let mut grid_h = (extent.h / cell_size).floor().max(1.0) as u32;
42
43        if u64::from(grid_w) * u64::from(grid_h) > MAX_INDEX_CELLS {
44            // Both axes by the same factor, so cells stay as square as the world lets them. Two
45            // floors of the same divisor cannot leave the product above the cap.
46            let scale = (f64::from(grid_w) * f64::from(grid_h) / MAX_INDEX_CELLS as f64).sqrt();
47            grid_w = ((f64::from(grid_w) / scale).floor() as u32).max(1);
48            grid_h = ((f64::from(grid_h) / scale).floor() as u32).max(1);
49        }
50
51        Self {
52            grid_w,
53            grid_h,
54            cell_w: extent.w / grid_w as f32,
55            cell_h: extent.h / grid_h as f32,
56        }
57    }
58
59    /// Number of cells, `grid_w * grid_h`.
60    pub fn num_cells(&self) -> u32 {
61        self.grid_w * self.grid_h
62    }
63}
64
65/// Flat counting-sort grid over agent positions, rebuilt every tick.
66pub struct SpatialHash {
67    /// Requested cell size, kept only to detect changes.
68    cell_size: f32,
69    // Actual cell extents, which tile the world exactly.
70    cell_w: f32,
71    cell_h: f32,
72    cell_w_inv: f32,
73    cell_h_inv: f32,
74    grid_w: u32,
75    grid_h: u32,
76    world_w: f32,
77    world_h: f32,
78    /// Flat cell index of each agent.
79    agent_cells: Vec<u32>,
80    /// Agents sorted by cell index.
81    sorted_agents: Vec<u32>,
82    /// Start index of each cell in `sorted_agents`.
83    cell_start: Vec<u32>,
84}
85
86/// Prints the cell geometry and the agent count, not the sorted agents.
87impl std::fmt::Debug for SpatialHash {
88    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
89        f.debug_struct("SpatialHash")
90            .field("cell_size", &self.cell_size)
91            .field("grid_w", &self.grid_w)
92            .field("grid_h", &self.grid_h)
93            .field("world_w", &self.world_w)
94            .field("world_h", &self.world_h)
95            .field("agent_count", &self.agent_cells.len())
96            .finish_non_exhaustive()
97    }
98}
99
100impl SpatialHash {
101    /// Creates an empty hash over a `world_w` by `world_h` world, with cells about `cell_size` wide.
102    pub fn new(cell_size: f32, world_w: f32, world_h: f32) -> Self {
103        // `HashGrid` fits whole cells to the world and caps how many there are. Sharing it keeps this
104        // sort and the GPU sort walking the same grid.
105        let grid = HashGrid::new(Extent { w: world_w, h: world_h }, cell_size);
106        let (grid_w, grid_h) = (grid.grid_w, grid.grid_h);
107        let (cell_w, cell_h) = (grid.cell_w, grid.cell_h);
108        let num_cells = grid_w * grid_h;
109
110        Self {
111            cell_size,
112            cell_w,
113            cell_h,
114            cell_w_inv: 1.0 / cell_w,
115            cell_h_inv: 1.0 / cell_h,
116            grid_w,
117            grid_h,
118            world_w,
119            world_h,
120            agent_cells: Vec::new(),
121            sorted_agents: Vec::new(),
122            cell_start: vec![0; num_cells as usize + 1],
123        }
124    }
125
126    /// Returns the flat index of the cell holding `(x, y)`.
127    ///
128    /// The position wraps, so a position outside the world still falls in a cell.
129    #[inline]
130    pub fn cell_index(&self, x: f32, y: f32) -> u32 {
131        let cx = ((x * self.cell_w_inv).floor() as i32).rem_euclid(self.grid_w as i32) as u32;
132        let cy = ((y * self.cell_h_inv).floor() as i32).rem_euclid(self.grid_h as i32) as u32;
133        cy * self.grid_w + cx
134    }
135
136    /// Sorts the agents at `pos_x` and `pos_y` into their cells.
137    pub fn build(&mut self, pos_x: &[f32], pos_y: &[f32]) {
138        self.build_where(pos_x, pos_y, |_| true);
139    }
140
141    /// Like [`Self::build`], but skips agents that `include` rejects.
142    ///
143    /// Skipped agents are not returned by any query.
144    pub fn build_where(&mut self, pos_x: &[f32], pos_y: &[f32], include: impl Fn(usize) -> bool) {
145        /// Cell index marking a skipped agent.
146        const NO_CELL: u32 = u32::MAX;
147
148        let num_agents = pos_x.len() as u32;
149        let num_cells = self.grid_w * self.grid_h;
150        self.agent_cells.clear();
151        self.sorted_agents.clear();
152        self.cell_start.clear();
153        self.agent_cells.reserve(num_agents as usize);
154        self.sorted_agents.resize(num_agents as usize, 0);
155        self.cell_start.resize((num_cells + 1) as usize, 0);
156
157        // Assigns each agent to its cell and counts the agents per cell.
158        for i in 0..num_agents {
159            let cell = if include(i as usize) {
160                self.cell_index(pos_x[i as usize], pos_y[i as usize])
161            } else {
162                NO_CELL
163            };
164            self.agent_cells.push(cell);
165            if cell != NO_CELL {
166                self.cell_start[cell as usize + 1] += 1;
167            }
168        }
169
170        // A prefix sum gives the start index of each cell.
171        for i in 1..=num_cells {
172            self.cell_start[i as usize] += self.cell_start[i as usize - 1];
173        }
174
175        // Sorts the agents by cell index with a counting sort.
176        let mut write_pos = self.cell_start.clone();
177        for i in 0..num_agents {
178            let cell = self.agent_cells[i as usize];
179            if cell == NO_CELL {
180                continue;
181            }
182            let pos = write_pos[cell as usize];
183            self.sorted_agents[pos as usize] = i;
184            write_pos[cell as usize] += 1;
185        }
186    }
187
188    /// Fills `result` with every agent within `r` of `(x, y)`, clearing it first.
189    ///
190    /// Distances wrap at the world's edges, as for [`Self::for_each_within`].
191    pub fn query_radius(&self, x: f32, y: f32, r: f32, pos_x: &[f32], pos_y: &[f32], result: &mut Vec<u32>) {
192        result.clear();
193        self.for_each_within(x, y, r, pos_x, pos_y, |agent_idx, _dx, _dy, _d2| {
194            result.push(agent_idx);
195        });
196    }
197
198    /// Visits every agent within `r` of `(x, y)`, passing the callback its index, the toroidal
199    /// deltas to it and their squared length.
200    ///
201    /// The deltas come out of the range test either way. A kernel that needs them uses this method
202    /// and computes each delta once, whereas a list of indices makes it recompute them all.
203    pub fn for_each_within<F: FnMut(u32, f32, f32, f32)>(
204        &self,
205        x: f32,
206        y: f32,
207        r: f32,
208        pos_x: &[f32],
209        pos_y: &[f32],
210        mut f: F,
211    ) {
212        let r2 = r * r;
213        let cell_radius_x = (r / self.cell_w).ceil() as i32;
214        let cell_radius_y = (r / self.cell_h).ceil() as i32;
215        let cell_x = ((x * self.cell_w_inv).floor() as i32).rem_euclid(self.grid_w as i32);
216        let cell_y = ((y * self.cell_h_inv).floor() as i32).rem_euclid(self.grid_h as i32);
217        // In case the radius is larger than the world, avoid repeatedly wrapping the grid.
218        let (y_lo, y_hi) = if 2 * cell_radius_y + 1 > self.grid_h as i32 {
219            (0, self.grid_h as i32 - 1)
220        } else {
221            (cell_y - cell_radius_y, cell_y + cell_radius_y)
222        };
223        let (x_lo, x_hi) = if 2 * cell_radius_x + 1 > self.grid_w as i32 {
224            (0, self.grid_w as i32 - 1)
225        } else {
226            (cell_x - cell_radius_x, cell_x + cell_radius_x)
227        };
228
229        for grid_y in y_lo..=y_hi {
230            let wrapped_y = grid_y.rem_euclid(self.grid_h as i32) as u32;
231            for grid_x in x_lo..=x_hi {
232                let wrapped_x = grid_x.rem_euclid(self.grid_w as i32) as u32;
233                let cell_index = wrapped_y * self.grid_w + wrapped_x;
234                let start = self.cell_start[cell_index as usize] as usize;
235                let end = self.cell_start[cell_index as usize + 1] as usize;
236                for &agent_idx in &self.sorted_agents[start..end] {
237                    let dx = axis_delta(x, pos_x[agent_idx as usize], self.world_w, Boundary::Torus);
238                    let dy = axis_delta(y, pos_y[agent_idx as usize], self.world_h, Boundary::Torus);
239                    let d2 = dx * dx + dy * dy;
240                    if d2 <= r2 {
241                        f(agent_idx, dx, dy, d2);
242                    }
243                }
244            }
245        }
246    }
247
248    /// Returns whether this hash was built with `cell_size`.
249    pub fn cell_size_is(&self, cell_size: f32) -> bool {
250        (self.cell_size - cell_size).abs() <= f32::EPSILON
251    }
252
253    /// Returns whether this hash was built for a `world_w` by `world_h` world.
254    pub fn world_is(&self, world_w: f32, world_h: f32) -> bool {
255        (self.world_w - world_w).abs() <= f32::EPSILON && (self.world_h - world_h).abs() <= f32::EPSILON
256    }
257
258    /// Cells along each axis, fitted to the world from the requested cell size.
259    pub fn grid_dims(&self) -> (u32, u32) {
260        (self.grid_w, self.grid_h)
261    }
262
263    /// World distance one cell spans on each axis.
264    pub fn cell_extents(&self) -> (f32, f32) {
265        (self.cell_w, self.cell_h)
266    }
267
268    /// Returns `(cell_start, sorted_agents)`, where cell `c` owns
269    /// `sorted_agents[cell_start[c]..cell_start[c + 1]]`.
270    pub fn buckets(&self) -> (&[u32], &[u32]) {
271        (&self.cell_start, &self.sorted_agents)
272    }
273
274    /// Rebuilds the hash for `new_cell_size` and sorts the agents again, if the cell size changed.
275    pub fn rebuild_with_cell_size(&mut self, new_cell_size: f32, pos_x: &[f32], pos_y: &[f32]) {
276        if (new_cell_size - self.cell_size).abs() > f32::EPSILON {
277            *self = Self::new(new_cell_size, self.world_w, self.world_h);
278            self.build(pos_x, pos_y);
279        }
280    }
281
282    /// Heap memory held by the hash, in bytes.
283    pub fn heap_bytes(&self) -> usize {
284        self.agent_cells.capacity() * 4 + self.sorted_agents.capacity() * 4 + self.cell_start.capacity() * 4
285    }
286}
287
288#[cfg(test)]
289mod tests {
290    use super::*;
291    use crate::authoring::primitives::rng::xorshift64;
292
293    /// A cell size the UI admits would otherwise request a grid of hundreds of millions of cells.
294    #[test]
295    fn the_index_grid_is_capped() {
296        let extent = Extent {
297            w: 10_000.0,
298            h: 10_000.0,
299        };
300        let grid = HashGrid::new(extent, 1.0);
301        assert!(
302            u64::from(grid.grid_w) * u64::from(grid.grid_h) <= MAX_INDEX_CELLS,
303            "{}x{} is over the cap",
304            grid.grid_w,
305            grid.grid_h
306        );
307        // Cells still tile the world exactly, which is what the wrap in a query relies on.
308        assert!((grid.cell_w * grid.grid_w as f32 - extent.w).abs() < 1e-3);
309        assert!((grid.cell_h * grid.grid_h as f32 - extent.h).abs() < 1e-3);
310    }
311
312    /// Both backends walk one geometry, so a query cannot read cells the other sort never wrote.
313    #[test]
314    fn the_sort_and_the_shared_geometry_agree() {
315        for (w, h, cell) in [
316            (1000.0, 1000.0, 50.0),
317            (1000.0, 1000.0, 47.0),
318            (10_000.0, 10_000.0, 1.0),
319        ] {
320            let hash = SpatialHash::new(cell, w, h);
321            let grid = HashGrid::new(Extent { w, h }, cell);
322            assert_eq!(hash.grid_dims(), (grid.grid_w, grid.grid_h), "dims at cell {cell}");
323            assert_eq!(
324                hash.cell_extents(),
325                (grid.cell_w, grid.cell_h),
326                "extents at cell {cell}"
327            );
328        }
329    }
330
331    #[test]
332    fn build_and_query_finds_all_close_agents() {
333        // Three agents near (0, 0), and one far away.
334        let pos_x = vec![0.0, 1.0, -2.0, 50.0];
335        let pos_y = vec![0.0, 2.0, -1.0, 50.0];
336        let mut sh = SpatialHash::new(10.0, 100.0, 100.0);
337        sh.build(&pos_x, &pos_y);
338
339        let mut result = Vec::new();
340        sh.query_radius(0.0, 0.0, 5.0, &pos_x, &pos_y, &mut result);
341
342        result.sort();
343        assert_eq!(result, vec![0, 1, 2]);
344    }
345
346    #[test]
347    fn toroidal_query_finds_wrapped_agent() {
348        // The world is 100 by 100, and the wrap puts an agent at (99, 99) near (1, 1).
349        let pos_x = vec![1.0, 99.0];
350        let pos_y = vec![1.0, 99.0];
351        let mut sh = SpatialHash::new(10.0, 100.0, 100.0);
352        sh.build(&pos_x, &pos_y);
353
354        let mut result = Vec::new();
355        sh.query_radius(1.0, 1.0, 5.0, &pos_x, &pos_y, &mut result);
356
357        result.sort();
358        assert_eq!(result, vec![0, 1]);
359    }
360
361    /// Unsigned toroidal distance on one axis, unlike the signed `space::axis_delta`.
362    fn axis_distance(a: f32, b: f32, world: f32) -> f32 {
363        let d = (a - b).abs();
364        d.min(world - d)
365    }
366
367    #[test]
368    fn matches_brute_force_with_non_divisor_cell_size() {
369        // 47 divides neither world axis, so the hash has to pick its own cell extents.
370        let (world_w, world_h, r) = (1_000.0_f32, 730.0_f32, 47.0_f32);
371        let mut seed = 0x1234_5678_9ABC_DEF0_u64;
372        let mut unit = || {
373            seed = xorshift64(seed);
374            (seed >> 40) as f32 / 16_777_216.0
375        };
376
377        let mut pos_x = Vec::new();
378        let mut pos_y = Vec::new();
379        for _ in 0..500 {
380            pos_x.push(unit() * world_w);
381            pos_y.push(unit() * world_h);
382        }
383
384        let mut sh = SpatialHash::new(r, world_w, world_h);
385        sh.build(&pos_x, &pos_y);
386
387        let mut result = Vec::new();
388        for i in 0..pos_x.len() {
389            sh.query_radius(pos_x[i], pos_y[i], r, &pos_x, &pos_y, &mut result);
390            result.sort();
391
392            let mut expected: Vec<u32> = (0..pos_x.len() as u32)
393                .filter(|&j| {
394                    let dx = axis_distance(pos_x[j as usize], pos_x[i], world_w);
395                    let dy = axis_distance(pos_y[j as usize], pos_y[i], world_h);
396                    dx * dx + dy * dy <= r * r
397                })
398                .collect();
399            expected.sort();
400
401            assert_eq!(result, expected, "neighbors of agent {i} disagree with brute force");
402        }
403    }
404
405    #[test]
406    fn query_wider_than_grid_returns_each_agent_once() {
407        // A radius of 100 in a world 300 wide leaves 3 cells per axis, so the walk spans the grid.
408        let pos_x = vec![0.0, 60.0, 150.0];
409        let pos_y = vec![0.0, 0.0, 0.0];
410        let mut sh = SpatialHash::new(100.0, 300.0, 300.0);
411        sh.build(&pos_x, &pos_y);
412
413        let mut result = Vec::new();
414        sh.query_radius(0.0, 0.0, 100.0, &pos_x, &pos_y, &mut result);
415
416        result.sort();
417        assert_eq!(result, vec![0, 1], "agent 2 is 150 away, and nothing may repeat");
418    }
419
420    #[test]
421    fn single_cell_grid_returns_each_agent_once() {
422        // A radius past half the world collapses the grid to one cell.
423        let pos_x = vec![10.0, 20.0, 60.0];
424        let pos_y = vec![10.0, 20.0, 60.0];
425        let mut sh = SpatialHash::new(200.0, 100.0, 100.0);
426        sh.build(&pos_x, &pos_y);
427
428        let mut result = Vec::new();
429        sh.query_radius(10.0, 10.0, 200.0, &pos_x, &pos_y, &mut result);
430
431        result.sort();
432        assert_eq!(result, vec![0, 1, 2], "one cell means one visit per agent");
433    }
434}