Skip to main content

spatialrust_gpu/kernels/
voxel_sort.rs

1use bytemuck::{Pod, Zeroable};
2use spatialrust_core::SpatialResult;
3use wgpu::util::DeviceExt;
4
5use crate::kernels::gpu_segments::GpuVoxelSegments;
6use crate::kernels::voxel_compact::compact_voxel_segments_gpu_buffers;
7use crate::kernels::voxel_keys::VoxelKeyOutput;
8use crate::kernels::voxel_segments::VoxelSegments;
9use crate::runtime::WgpuRuntime;
10
11const WORKGROUP_SIZE: u32 = 256;
12
13#[repr(C)]
14#[derive(Clone, Copy, Debug, Pod, Zeroable)]
15struct SortParams {
16    padded_count: u32,
17    pair_distance: u32,
18    block_width: u32,
19    _pad: u32,
20}
21
22#[repr(C)]
23#[derive(Clone, Copy, Debug, Default, Pod, Zeroable)]
24pub(crate) struct VoxelSortEntry {
25    ix: i32,
26    iy: i32,
27    iz: i32,
28    point_index: u32,
29}
30
31#[repr(C)]
32#[derive(Clone, Copy, Debug, Pod, Zeroable)]
33struct BuildEntriesParams {
34    point_count: u32,
35    padded_count: u32,
36    _pad0: u32,
37    _pad1: u32,
38}
39
40/// Sorts per-point voxel keys on the GPU and compacts them into segments.
41pub fn build_voxel_segments_gpu(
42    runtime: &WgpuRuntime,
43    keys: &[(i64, i64, i64)],
44) -> SpatialResult<VoxelSegments> {
45    let gpu_segments = build_voxel_segments_gpu_from_keys(runtime, keys)?;
46    gpu_segments.to_voxel_segments(runtime)
47}
48
49/// Builds GPU-resident voxel segments from CPU-side keys.
50pub fn build_voxel_segments_gpu_from_keys(
51    runtime: &WgpuRuntime,
52    keys: &[(i64, i64, i64)],
53) -> SpatialResult<GpuVoxelSegments> {
54    if keys.is_empty() {
55        return empty_gpu_segments(runtime);
56    }
57
58    let point_count = keys.len();
59    let padded_count = point_count.next_power_of_two();
60    let mut key_outputs = vec![VoxelKeyOutput::default(); point_count];
61    for (index, (ix, iy, iz)) in keys.iter().copied().enumerate() {
62        key_outputs[index] = VoxelKeyOutput {
63            ix: ix.clamp(i32::MIN as i64, i32::MAX as i64) as i32,
64            iy: iy.clamp(i32::MIN as i64, i32::MAX as i64) as i32,
65            iz: iz.clamp(i32::MIN as i64, i32::MAX as i64) as i32,
66            _pad: 0,
67        };
68    }
69
70    let device = runtime.device();
71    let keys_buffer = device.create_buffer_init(&wgpu::util::BufferInitDescriptor {
72        label: Some("voxel-sort-keys-input"),
73        contents: bytemuck::cast_slice(&key_outputs),
74        usage: wgpu::BufferUsages::STORAGE,
75    });
76
77    build_voxel_segments_gpu_from_keys_buffer(
78        runtime,
79        &keys_buffer,
80        point_count as u32,
81        padded_count as u32,
82    )
83}
84
85/// Builds GPU-resident voxel segments from a GPU keys buffer.
86pub fn build_voxel_segments_gpu_from_keys_buffer(
87    runtime: &WgpuRuntime,
88    keys_buffer: &wgpu::Buffer,
89    point_count: u32,
90    padded_count: u32,
91) -> SpatialResult<GpuVoxelSegments> {
92    if point_count == 0 {
93        return empty_gpu_segments(runtime);
94    }
95
96    let entries_buffer =
97        build_sort_entries_from_keys_gpu(runtime, keys_buffer, point_count, padded_count)?;
98    let sorted_buffer = sort_entries_gpu(runtime, entries_buffer, padded_count)?;
99    let compact_entries =
100        filter_valid_sorted_entries(runtime, &sorted_buffer, padded_count, point_count)?;
101    compact_voxel_segments_gpu_buffers(runtime, &compact_entries, point_count)
102}
103
104fn build_sort_entries_from_keys_gpu(
105    runtime: &WgpuRuntime,
106    keys_buffer: &wgpu::Buffer,
107    point_count: u32,
108    padded_count: u32,
109) -> SpatialResult<wgpu::Buffer> {
110    let device = runtime.device();
111    let queue = runtime.queue();
112
113    let entries_buffer = device.create_buffer(&wgpu::BufferDescriptor {
114        label: Some("voxel-sort-entries-build"),
115        size: (padded_count as usize * std::mem::size_of::<VoxelSortEntry>()) as u64,
116        usage: wgpu::BufferUsages::STORAGE
117            | wgpu::BufferUsages::COPY_DST
118            | wgpu::BufferUsages::COPY_SRC,
119        mapped_at_creation: false,
120    });
121
122    let params = BuildEntriesParams { point_count, padded_count, _pad0: 0, _pad1: 0 };
123    let params_buffer = device.create_buffer_init(&wgpu::util::BufferInitDescriptor {
124        label: Some("voxel-sort-build-params"),
125        contents: bytemuck::bytes_of(&params),
126        usage: wgpu::BufferUsages::UNIFORM,
127    });
128
129    let pipelines = runtime.pipelines();
130
131    let bind_group = device.create_bind_group(&wgpu::BindGroupDescriptor {
132        label: Some("voxel-sort-build-bind-group"),
133        layout: &pipelines.voxel_sort_build.bind_group_layout,
134        entries: &[
135            wgpu::BindGroupEntry { binding: 0, resource: params_buffer.as_entire_binding() },
136            wgpu::BindGroupEntry { binding: 1, resource: keys_buffer.as_entire_binding() },
137            wgpu::BindGroupEntry { binding: 2, resource: entries_buffer.as_entire_binding() },
138        ],
139    });
140
141    let mut encoder = device.create_command_encoder(&wgpu::CommandEncoderDescriptor {
142        label: Some("voxel-sort-build-encoder"),
143    });
144    {
145        let mut pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
146            label: Some("voxel-sort-build-pass"),
147            timestamp_writes: None,
148        });
149        pass.set_pipeline(&pipelines.voxel_sort_build.pipeline);
150        pass.set_bind_group(0, &bind_group, &[]);
151        pass.dispatch_workgroups(padded_count.div_ceil(WORKGROUP_SIZE), 1, 1);
152    }
153    queue.submit(Some(encoder.finish()));
154
155    Ok(entries_buffer)
156}
157
158fn empty_gpu_segments(runtime: &WgpuRuntime) -> SpatialResult<GpuVoxelSegments> {
159    let device = runtime.device();
160    let make_empty = || {
161        device.create_buffer(&wgpu::BufferDescriptor {
162            label: Some("voxel-sort-empty"),
163            size: 4,
164            usage: wgpu::BufferUsages::STORAGE,
165            mapped_at_creation: false,
166        })
167    };
168    Ok(GpuVoxelSegments::new(0, 0, make_empty(), make_empty(), make_empty()))
169}
170
171fn filter_valid_sorted_entries(
172    runtime: &WgpuRuntime,
173    entries_buffer: &wgpu::Buffer,
174    padded_count: u32,
175    point_count: u32,
176) -> SpatialResult<wgpu::Buffer> {
177    let device = runtime.device();
178    let queue = runtime.queue();
179    let buffer_len = padded_count as u64;
180    let output_len = (point_count as usize * std::mem::size_of::<VoxelSortEntry>()) as u64;
181
182    let flags_buffer = device.create_buffer(&wgpu::BufferDescriptor {
183        label: Some("voxel-sort-filter-flags"),
184        size: buffer_len * std::mem::size_of::<u32>() as u64,
185        usage: wgpu::BufferUsages::STORAGE
186            | wgpu::BufferUsages::COPY_DST
187            | wgpu::BufferUsages::COPY_SRC,
188        mapped_at_creation: false,
189    });
190    let inclusive_buffer = device.create_buffer(&wgpu::BufferDescriptor {
191        label: Some("voxel-sort-filter-inclusive"),
192        size: buffer_len * std::mem::size_of::<u32>() as u64,
193        usage: wgpu::BufferUsages::STORAGE
194            | wgpu::BufferUsages::COPY_SRC
195            | wgpu::BufferUsages::COPY_DST,
196        mapped_at_creation: false,
197    });
198    let scan_scratch_buffer = device.create_buffer(&wgpu::BufferDescriptor {
199        label: Some("voxel-sort-filter-scan-scratch"),
200        size: buffer_len * std::mem::size_of::<u32>() as u64,
201        usage: wgpu::BufferUsages::STORAGE
202            | wgpu::BufferUsages::COPY_DST
203            | wgpu::BufferUsages::COPY_SRC,
204        mapped_at_creation: false,
205    });
206    let output_buffer = device.create_buffer(&wgpu::BufferDescriptor {
207        label: Some("voxel-sort-filter-output"),
208        size: output_len,
209        usage: wgpu::BufferUsages::STORAGE,
210        mapped_at_creation: false,
211    });
212
213    let pipelines = runtime.pipelines();
214    let layout = &pipelines.voxel_sort_filter.bind_group_layout;
215    let mark_params = create_filter_params_buffer(device, point_count, padded_count, 0);
216    let mark_bind_group = create_filter_bind_group(
217        device,
218        layout,
219        &mark_params,
220        entries_buffer,
221        &flags_buffer,
222        &scan_scratch_buffer,
223        &inclusive_buffer,
224        &output_buffer,
225    );
226    let init_params = create_filter_params_buffer(device, point_count, padded_count, 0);
227    let init_bind_group = create_filter_bind_group(
228        device,
229        layout,
230        &init_params,
231        entries_buffer,
232        &flags_buffer,
233        &scan_scratch_buffer,
234        &inclusive_buffer,
235        &output_buffer,
236    );
237
238    let dispatch_padded = padded_count.div_ceil(WORKGROUP_SIZE);
239    let mut encoder = device.create_command_encoder(&wgpu::CommandEncoderDescriptor {
240        label: Some("voxel-sort-filter-batched-encoder"),
241    });
242    encoder.clear_buffer(&flags_buffer, 0, None);
243    encoder.clear_buffer(&inclusive_buffer, 0, None);
244    encoder.clear_buffer(&scan_scratch_buffer, 0, None);
245    {
246        let mut pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
247            label: Some("voxel-sort-filter-mark-pass"),
248            timestamp_writes: None,
249        });
250        pass.set_pipeline(&pipelines.voxel_sort_filter.mark);
251        pass.set_bind_group(0, &mark_bind_group, &[]);
252        pass.dispatch_workgroups(dispatch_padded, 1, 1);
253    }
254
255    {
256        let mut pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
257            label: Some("voxel-sort-filter-init-pass"),
258            timestamp_writes: None,
259        });
260        pass.set_pipeline(&pipelines.voxel_sort_filter.init);
261        pass.set_bind_group(0, &init_bind_group, &[]);
262        pass.dispatch_workgroups(dispatch_padded, 1, 1);
263    }
264
265    let mut scan_read = &inclusive_buffer;
266    let mut scan_write = &scan_scratch_buffer;
267    let mut stride = 1u32;
268    while stride < padded_count {
269        let scan_params = create_filter_params_buffer(device, point_count, padded_count, stride);
270        let scan_bind_group = create_filter_bind_group(
271            device,
272            layout,
273            &scan_params,
274            entries_buffer,
275            &flags_buffer,
276            scan_read,
277            scan_write,
278            &output_buffer,
279        );
280        {
281            let mut pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
282                label: Some("voxel-sort-filter-scan-pass"),
283                timestamp_writes: None,
284            });
285            pass.set_pipeline(&pipelines.voxel_sort_filter.scan);
286            pass.set_bind_group(0, &scan_bind_group, &[]);
287            pass.dispatch_workgroups(dispatch_padded, 1, 1);
288        }
289        std::mem::swap(&mut scan_read, &mut scan_write);
290        stride *= 2;
291    }
292    queue.submit(Some(encoder.finish()));
293
294    let valid_count = read_filter_valid_count(device, queue, scan_read, padded_count)?;
295    if valid_count != point_count as usize {
296        return Err(spatialrust_core::SpatialError::InvalidArgument(format!(
297            "expected {point_count} sorted voxel entries, found {valid_count}"
298        )));
299    }
300
301    let scatter_params = create_filter_params_buffer(device, point_count, padded_count, 0);
302    let scatter_bind_group = create_filter_bind_group(
303        device,
304        layout,
305        &scatter_params,
306        entries_buffer,
307        &flags_buffer,
308        scan_read,
309        scan_write,
310        &output_buffer,
311    );
312    let mut encoder = device.create_command_encoder(&wgpu::CommandEncoderDescriptor {
313        label: Some("voxel-sort-filter-scatter-encoder"),
314    });
315    {
316        let mut pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
317            label: Some("voxel-sort-filter-scatter-pass"),
318            timestamp_writes: None,
319        });
320        pass.set_pipeline(&pipelines.voxel_sort_filter.scatter);
321        pass.set_bind_group(0, &scatter_bind_group, &[]);
322        pass.dispatch_workgroups(dispatch_padded, 1, 1);
323    }
324    queue.submit(Some(encoder.finish()));
325
326    Ok(output_buffer)
327}
328
329fn create_filter_bind_group(
330    device: &wgpu::Device,
331    layout: &wgpu::BindGroupLayout,
332    params_buffer: &wgpu::Buffer,
333    entries_buffer: &wgpu::Buffer,
334    flags_buffer: &wgpu::Buffer,
335    scan_in: &wgpu::Buffer,
336    scan_out: &wgpu::Buffer,
337    output_buffer: &wgpu::Buffer,
338) -> wgpu::BindGroup {
339    device.create_bind_group(&wgpu::BindGroupDescriptor {
340        label: Some("voxel-sort-filter-bind-group"),
341        layout,
342        entries: &[
343            wgpu::BindGroupEntry { binding: 0, resource: params_buffer.as_entire_binding() },
344            wgpu::BindGroupEntry { binding: 1, resource: entries_buffer.as_entire_binding() },
345            wgpu::BindGroupEntry { binding: 2, resource: flags_buffer.as_entire_binding() },
346            wgpu::BindGroupEntry { binding: 3, resource: scan_in.as_entire_binding() },
347            wgpu::BindGroupEntry { binding: 4, resource: scan_out.as_entire_binding() },
348            wgpu::BindGroupEntry { binding: 5, resource: output_buffer.as_entire_binding() },
349        ],
350    })
351}
352
353#[repr(C)]
354#[derive(Clone, Copy, Debug, Pod, Zeroable)]
355struct FilterParams {
356    point_count: u32,
357    padded_count: u32,
358    scan_stride: u32,
359    _pad: u32,
360}
361
362fn create_filter_params_buffer(
363    device: &wgpu::Device,
364    point_count: u32,
365    padded_count: u32,
366    scan_stride: u32,
367) -> wgpu::Buffer {
368    let params = FilterParams { point_count, padded_count, scan_stride, _pad: 0 };
369    device.create_buffer_init(&wgpu::util::BufferInitDescriptor {
370        label: Some("voxel-sort-filter-params"),
371        contents: bytemuck::bytes_of(&params),
372        usage: wgpu::BufferUsages::UNIFORM,
373    })
374}
375
376fn read_filter_valid_count(
377    device: &wgpu::Device,
378    queue: &wgpu::Queue,
379    inclusive_buffer: &wgpu::Buffer,
380    padded_count: u32,
381) -> SpatialResult<usize> {
382    let offset = ((padded_count as u64).saturating_sub(1)) * std::mem::size_of::<u32>() as u64;
383    let staging = device.create_buffer(&wgpu::BufferDescriptor {
384        label: Some("voxel-sort-filter-count-staging"),
385        size: std::mem::size_of::<u32>() as u64,
386        usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
387        mapped_at_creation: false,
388    });
389
390    let mut encoder = device.create_command_encoder(&wgpu::CommandEncoderDescriptor {
391        label: Some("voxel-sort-filter-count-encoder"),
392    });
393    encoder.copy_buffer_to_buffer(inclusive_buffer, offset, &staging, 0, staging.size());
394    queue.submit(Some(encoder.finish()));
395
396    let slice = staging.slice(..);
397    let (sender, receiver) = std::sync::mpsc::channel();
398    slice.map_async(wgpu::MapMode::Read, move |result| {
399        let _ = sender.send(result);
400    });
401    device.poll(wgpu::Maintain::Wait);
402    receiver
403        .recv()
404        .map_err(|_| {
405            spatialrust_core::SpatialError::InvalidArgument(
406                "failed to receive wgpu map result".to_owned(),
407            )
408        })?
409        .map_err(|error| {
410            spatialrust_core::SpatialError::InvalidArgument(format!(
411                "failed to map wgpu buffer: {error}"
412            ))
413        })?;
414
415    let data = slice.get_mapped_range();
416    let count = bytemuck::cast_slice::<u8, u32>(&data)[0] as usize;
417    drop(data);
418    staging.unmap();
419    Ok(count)
420}
421
422fn sort_entries_gpu(
423    runtime: &WgpuRuntime,
424    entries_buffer: wgpu::Buffer,
425    padded_count: u32,
426) -> SpatialResult<wgpu::Buffer> {
427    let device = runtime.device();
428    let queue = runtime.queue();
429    let pipelines = runtime.pipelines();
430    let mut encoder = device.create_command_encoder(&wgpu::CommandEncoderDescriptor {
431        label: Some("voxel-sort-batched-encoder"),
432    });
433
434    let mut k = 2u32;
435    while k <= padded_count {
436        let mut j = k / 2;
437        while j >= 1 {
438            let params = SortParams { padded_count, pair_distance: j, block_width: k, _pad: 0 };
439            let params_buffer = device.create_buffer_init(&wgpu::util::BufferInitDescriptor {
440                label: Some("voxel-sort-params"),
441                contents: bytemuck::bytes_of(&params),
442                usage: wgpu::BufferUsages::UNIFORM,
443            });
444            let bind_group = device.create_bind_group(&wgpu::BindGroupDescriptor {
445                label: Some("voxel-sort-bind-group"),
446                layout: &pipelines.voxel_sort.bind_group_layout,
447                entries: &[
448                    wgpu::BindGroupEntry {
449                        binding: 0,
450                        resource: params_buffer.as_entire_binding(),
451                    },
452                    wgpu::BindGroupEntry {
453                        binding: 1,
454                        resource: entries_buffer.as_entire_binding(),
455                    },
456                ],
457            });
458            {
459                let mut pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
460                    label: Some("voxel-sort-pass"),
461                    timestamp_writes: None,
462                });
463                pass.set_pipeline(&pipelines.voxel_sort.pipeline);
464                pass.set_bind_group(0, &bind_group, &[]);
465                pass.dispatch_workgroups(padded_count.div_ceil(WORKGROUP_SIZE), 1, 1);
466            }
467            j /= 2;
468        }
469        k *= 2;
470    }
471    queue.submit(Some(encoder.finish()));
472
473    Ok(entries_buffer)
474}
475
476#[cfg(test)]
477mod tests {
478    use super::build_voxel_segments_gpu;
479    use crate::kernels::voxel_segments::build_voxel_segments;
480    use crate::runtime::WgpuRuntime;
481
482    #[test]
483    fn gpu_segment_build_matches_cpu_reference() {
484        let runtime = WgpuRuntime::new_headless().expect("wgpu runtime");
485        let keys = vec![(0, 0, 0), (1, 0, 0), (0, 0, 0), (1, 0, 0), (2, 1, 0)];
486        let cpu = build_voxel_segments(&keys);
487        let gpu = build_voxel_segments_gpu(&runtime, &keys).expect("gpu segments");
488        assert_eq!(cpu, gpu);
489    }
490}