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
40pub 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
49pub 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
85pub 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(¶ms),
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(¶ms),
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(¶ms),
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}