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}