Skip to main content

dynamis_gpu/
sort.rs

1use crate::buffer::GpuBuffer;
2use crate::{BindingKind, BindingSpec, ComputePipeline};
3use wgpu::{BindGroup, BindGroupEntry, CommandEncoder, Device};
4
5const THREADS: u32 = 256;
6const RADIX_PASSES: usize = 8;
7
8fn shifted_shader(source: &str, shift: u32) -> String {
9    source.replace("__SHIFT__", &shift.to_string())
10}
11
12pub struct GpuSort {
13    histogram_pipelines: [ComputePipeline; RADIX_PASSES],
14    assign_pipelines: [ComputePipeline; RADIX_PASSES],
15    scatter_pipelines: [ComputePipeline; RADIX_PASSES],
16    histogram: GpuBuffer,
17    positions: GpuBuffer,
18}
19
20impl GpuSort {
21    pub fn new(device: &Device, label: &str, capacity: u32) -> Self {
22        let histogram_bindings = [
23            BindingSpec {
24                binding: 0,
25                kind: BindingKind::ReadOnlyStorage,
26            },
27            BindingSpec {
28                binding: 1,
29                kind: BindingKind::ReadOnlyStorage,
30            },
31            BindingSpec {
32                binding: 2,
33                kind: BindingKind::ReadWriteStorage,
34            },
35            BindingSpec {
36                binding: 3,
37                kind: BindingKind::ReadOnlyStorage,
38            },
39        ];
40        let assign_bindings = [
41            BindingSpec {
42                binding: 0,
43                kind: BindingKind::ReadOnlyStorage,
44            },
45            BindingSpec {
46                binding: 1,
47                kind: BindingKind::ReadOnlyStorage,
48            },
49            BindingSpec {
50                binding: 2,
51                kind: BindingKind::ReadWriteStorage,
52            },
53            BindingSpec {
54                binding: 3,
55                kind: BindingKind::ReadOnlyStorage,
56            },
57            BindingSpec {
58                binding: 4,
59                kind: BindingKind::ReadWriteStorage,
60            },
61        ];
62        let scatter_bindings = [
63            BindingSpec {
64                binding: 0,
65                kind: BindingKind::ReadOnlyStorage,
66            },
67            BindingSpec {
68                binding: 1,
69                kind: BindingKind::ReadOnlyStorage,
70            },
71            BindingSpec {
72                binding: 2,
73                kind: BindingKind::ReadOnlyStorage,
74            },
75            BindingSpec {
76                binding: 3,
77                kind: BindingKind::ReadOnlyStorage,
78            },
79            BindingSpec {
80                binding: 4,
81                kind: BindingKind::ReadWriteStorage,
82            },
83            BindingSpec {
84                binding: 5,
85                kind: BindingKind::ReadWriteStorage,
86            },
87            BindingSpec {
88                binding: 6,
89                kind: BindingKind::ReadWriteStorage,
90            },
91            BindingSpec {
92                binding: 7,
93                kind: BindingKind::ReadOnlyStorage,
94            },
95        ];
96        let histogram_shader = include_str!("shaders/sort_histogram.wgsl");
97        let assign_shader = include_str!("shaders/sort_assign.wgsl");
98        let scatter_shader = include_str!("shaders/sort_scatter.wgsl");
99        let histogram_pipelines = std::array::from_fn(|index| {
100            ComputePipeline::new(
101                device,
102                &format!("{label} histogram {index}"),
103                &shifted_shader(histogram_shader, (index * 8) as u32),
104                "main",
105                &histogram_bindings,
106                THREADS,
107            )
108        });
109        let assign_pipelines = std::array::from_fn(|index| {
110            ComputePipeline::new(
111                device,
112                &format!("{label} assign {index}"),
113                &shifted_shader(assign_shader, (index * 8) as u32),
114                "main",
115                &assign_bindings,
116                256,
117            )
118        });
119        let scatter_pipelines = std::array::from_fn(|index| {
120            ComputePipeline::new(
121                device,
122                &format!("{label} scatter {index}"),
123                &shifted_shader(scatter_shader, (index * 8) as u32),
124                "main",
125                &scatter_bindings,
126                THREADS,
127            )
128        });
129        Self {
130            histogram_pipelines,
131            assign_pipelines,
132            scatter_pipelines,
133            histogram: GpuBuffer::new(
134                device,
135                &format!("{label} histogram"),
136                256 * 4,
137                wgpu::BufferUsages::STORAGE
138                    | wgpu::BufferUsages::COPY_DST
139                    | wgpu::BufferUsages::COPY_SRC,
140            ),
141            positions: GpuBuffer::new(
142                device,
143                &format!("{label} positions"),
144                capacity as u64 * 4,
145                wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_SRC,
146            ),
147        }
148    }
149
150    #[expect(
151        clippy::too_many_arguments,
152        reason = "radix sort carries three key channels and their outputs explicitly"
153    )]
154    pub fn sort_pass(
155        &self,
156        device: &Device,
157        encoder: &mut CommandEncoder,
158        count_holder: &GpuBuffer,
159        capacity: u32,
160        pass: usize,
161        keys_lo: &GpuBuffer,
162        keys_hi: &GpuBuffer,
163        values: &GpuBuffer,
164        keys_lo_out: &GpuBuffer,
165        keys_hi_out: &GpuBuffer,
166        values_out: &GpuBuffer,
167    ) {
168        let workgroups = capacity.div_ceil(THREADS);
169        encoder.clear_buffer(self.histogram.buffer(), 0, None);
170        let histogram_group = self.histogram_pipelines[pass].create_bind_group(
171            device,
172            &[
173                BindGroupEntry {
174                    binding: 0,
175                    resource: keys_lo.as_binding(),
176                },
177                BindGroupEntry {
178                    binding: 1,
179                    resource: keys_hi.as_binding(),
180                },
181                BindGroupEntry {
182                    binding: 2,
183                    resource: self.histogram.as_binding(),
184                },
185                BindGroupEntry {
186                    binding: 3,
187                    resource: count_holder.as_binding(),
188                },
189            ],
190        );
191        encode_dispatch(
192            encoder,
193            &self.histogram_pipelines[pass],
194            &histogram_group,
195            workgroups,
196        );
197
198        let assign_group = self.assign_pipelines[pass].create_bind_group(
199            device,
200            &[
201                BindGroupEntry {
202                    binding: 0,
203                    resource: keys_lo.as_binding(),
204                },
205                BindGroupEntry {
206                    binding: 1,
207                    resource: keys_hi.as_binding(),
208                },
209                BindGroupEntry {
210                    binding: 2,
211                    resource: self.histogram.as_binding(),
212                },
213                BindGroupEntry {
214                    binding: 3,
215                    resource: count_holder.as_binding(),
216                },
217                BindGroupEntry {
218                    binding: 4,
219                    resource: self.positions.as_binding(),
220                },
221            ],
222        );
223        encode_dispatch(encoder, &self.assign_pipelines[pass], &assign_group, 1);
224
225        let scatter_group = self.scatter_pipelines[pass].create_bind_group(
226            device,
227            &[
228                BindGroupEntry {
229                    binding: 0,
230                    resource: keys_lo.as_binding(),
231                },
232                BindGroupEntry {
233                    binding: 1,
234                    resource: keys_hi.as_binding(),
235                },
236                BindGroupEntry {
237                    binding: 2,
238                    resource: values.as_binding(),
239                },
240                BindGroupEntry {
241                    binding: 3,
242                    resource: self.positions.as_binding(),
243                },
244                BindGroupEntry {
245                    binding: 4,
246                    resource: keys_lo_out.as_binding(),
247                },
248                BindGroupEntry {
249                    binding: 5,
250                    resource: keys_hi_out.as_binding(),
251                },
252                BindGroupEntry {
253                    binding: 6,
254                    resource: values_out.as_binding(),
255                },
256                BindGroupEntry {
257                    binding: 7,
258                    resource: count_holder.as_binding(),
259                },
260            ],
261        );
262        encode_dispatch(
263            encoder,
264            &self.scatter_pipelines[pass],
265            &scatter_group,
266            workgroups,
267        );
268    }
269
270    #[expect(
271        clippy::too_many_arguments,
272        reason = "radix sort carries three key channels and their outputs explicitly"
273    )]
274    pub fn sort_64(
275        &self,
276        device: &Device,
277        encoder: &mut CommandEncoder,
278        count_holder: &GpuBuffer,
279        capacity: u32,
280        keys_lo: &GpuBuffer,
281        keys_hi: &GpuBuffer,
282        values: &GpuBuffer,
283        keys_lo_out: &GpuBuffer,
284        keys_hi_out: &GpuBuffer,
285        values_out: &GpuBuffer,
286    ) {
287        let workgroups = capacity.div_ceil(THREADS);
288        let mut out_is_current = false;
289        for pass in 0..RADIX_PASSES {
290            encoder.clear_buffer(self.histogram.buffer(), 0, None);
291            let (in_lo, in_hi, in_values) = if out_is_current {
292                (keys_lo_out, keys_hi_out, values_out)
293            } else {
294                (keys_lo, keys_hi, values)
295            };
296
297            let histogram_group = self.histogram_pipelines[pass].create_bind_group(
298                device,
299                &[
300                    BindGroupEntry {
301                        binding: 0,
302                        resource: in_lo.as_binding(),
303                    },
304                    BindGroupEntry {
305                        binding: 1,
306                        resource: in_hi.as_binding(),
307                    },
308                    BindGroupEntry {
309                        binding: 2,
310                        resource: self.histogram.as_binding(),
311                    },
312                    BindGroupEntry {
313                        binding: 3,
314                        resource: count_holder.as_binding(),
315                    },
316                ],
317            );
318            encode_dispatch(
319                encoder,
320                &self.histogram_pipelines[pass],
321                &histogram_group,
322                workgroups,
323            );
324
325            let assign_group = self.assign_pipelines[pass].create_bind_group(
326                device,
327                &[
328                    BindGroupEntry {
329                        binding: 0,
330                        resource: in_lo.as_binding(),
331                    },
332                    BindGroupEntry {
333                        binding: 1,
334                        resource: in_hi.as_binding(),
335                    },
336                    BindGroupEntry {
337                        binding: 2,
338                        resource: self.histogram.as_binding(),
339                    },
340                    BindGroupEntry {
341                        binding: 3,
342                        resource: count_holder.as_binding(),
343                    },
344                    BindGroupEntry {
345                        binding: 4,
346                        resource: self.positions.as_binding(),
347                    },
348                ],
349            );
350            encode_dispatch(encoder, &self.assign_pipelines[pass], &assign_group, 1);
351
352            let (out_lo, out_hi, out_values) = if out_is_current {
353                (keys_lo, keys_hi, values)
354            } else {
355                (keys_lo_out, keys_hi_out, values_out)
356            };
357            let scatter_group = self.scatter_pipelines[pass].create_bind_group(
358                device,
359                &[
360                    BindGroupEntry {
361                        binding: 0,
362                        resource: in_lo.as_binding(),
363                    },
364                    BindGroupEntry {
365                        binding: 1,
366                        resource: in_hi.as_binding(),
367                    },
368                    BindGroupEntry {
369                        binding: 2,
370                        resource: in_values.as_binding(),
371                    },
372                    BindGroupEntry {
373                        binding: 3,
374                        resource: self.positions.as_binding(),
375                    },
376                    BindGroupEntry {
377                        binding: 4,
378                        resource: out_lo.as_binding(),
379                    },
380                    BindGroupEntry {
381                        binding: 5,
382                        resource: out_hi.as_binding(),
383                    },
384                    BindGroupEntry {
385                        binding: 6,
386                        resource: out_values.as_binding(),
387                    },
388                    BindGroupEntry {
389                        binding: 7,
390                        resource: count_holder.as_binding(),
391                    },
392                ],
393            );
394            encode_dispatch(
395                encoder,
396                &self.scatter_pipelines[pass],
397                &scatter_group,
398                workgroups,
399            );
400
401            out_is_current = !out_is_current;
402        }
403    }
404
405    pub fn debug_histogram(&self) -> &GpuBuffer {
406        &self.histogram
407    }
408
409    pub fn debug_positions(&self) -> &GpuBuffer {
410        &self.positions
411    }
412}
413
414fn encode_dispatch(
415    encoder: &mut CommandEncoder,
416    pipeline: &ComputePipeline,
417    group: &BindGroup,
418    workgroups: u32,
419) {
420    let mut pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
421        label: None,
422        timestamp_writes: None,
423    });
424    pass.set_pipeline(pipeline.pipeline());
425    pass.set_bind_group(0, group, &[]);
426    pass.dispatch_workgroups(workgroups, 1, 1);
427}