wgpu_primitives/scan/
scanner.rs1use super::pipeline::ScanPipeline;
2use crate::{Error, common, context::Context};
3
4pub struct Scanner {
6 pipeline: ScanPipeline,
7 device: wgpu::Device,
8 queue: wgpu::Queue,
9 scratch_buffer: Option<wgpu::Buffer>,
10 scratch_size_bytes: u64,
11}
12
13impl Scanner {
14 pub fn new(device: &wgpu::Device, queue: &wgpu::Queue) -> Self {
16 Self {
17 pipeline: ScanPipeline::new(device),
18 device: device.clone(),
19 queue: queue.clone(),
20 scratch_buffer: None,
21 scratch_size_bytes: 0,
22 }
23 }
24
25 pub fn from_context(ctx: &Context) -> Self {
27 Self::new(&ctx.device, &ctx.queue)
28 }
29
30 pub async fn scan(&mut self, input: &[u32]) -> Result<Vec<u32>, Error> {
32 if input.is_empty() {
33 return Ok(Vec::new());
34 }
35
36 let num_items = common::math::checked_u32(input.len() as u64)?;
37 let data_buffer = common::buffers::create_storage_buffer(&self.device, input);
38 let dst_buffer =
39 common::buffers::create_empty_storage_buffer(&self.device, data_buffer.size());
40
41 self.scan_gpu_to_gpu(&data_buffer, &dst_buffer, num_items)?;
42
43 let size_bytes = common::math::checked_byte_size(input.len() as u64, 4)?;
44 common::buffers::download_buffer(&self.device, &self.queue, &dst_buffer, size_bytes).await
45 }
46
47 pub fn scan_gpu_to_gpu(
49 &mut self,
50 input_buf: &wgpu::Buffer,
51 output_buf: &wgpu::Buffer,
52 num_items: u32,
53 ) -> Result<(), Error> {
54 let mut encoder = self
55 .device
56 .create_command_encoder(&wgpu::CommandEncoderDescriptor::default());
57 self.record_scan(&mut encoder, input_buf, output_buf, num_items)?;
58 self.queue.submit(Some(encoder.finish()));
59 Ok(())
60 }
61
62 pub fn record_scan(
64 &mut self,
65 encoder: &mut wgpu::CommandEncoder,
66 input_buf: &wgpu::Buffer,
67 output_buf: &wgpu::Buffer,
68 num_items: u32,
69 ) -> Result<(), Error> {
70 if num_items == 0 {
71 return Ok(());
72 }
73
74 let size_bytes = common::math::checked_byte_size(u64::from(num_items), 4)?;
75 common::buffers::validate_buffer(
76 input_buf,
77 "scan input",
78 size_bytes,
79 wgpu::BufferUsages::COPY_SRC,
80 )?;
81 common::buffers::validate_buffer(
82 output_buf,
83 "scan output",
84 size_bytes,
85 wgpu::BufferUsages::COPY_DST | wgpu::BufferUsages::STORAGE,
86 )?;
87
88 encoder.copy_buffer_to_buffer(input_buf, 0, output_buf, 0, size_bytes);
89
90 if num_items == 1 {
91 return Ok(());
92 }
93
94 self.prepare_scratch(num_items);
95
96 let scratch = self
97 .scratch_buffer
98 .as_ref()
99 .expect("scan scratch exists for multi-element inputs");
100
101 struct Level<'a> {
102 buf: &'a wgpu::Buffer,
103 offset: u64,
104 count: u32,
105 }
106
107 let mut levels = Vec::new();
108 levels.push(Level {
109 buf: output_buf,
110 offset: 0,
111 count: num_items,
112 });
113
114 let mut current_scratch_offset = 0u64;
115
116 loop {
117 let current = levels.last().unwrap();
118 if current.count <= 1 {
119 break;
120 }
121
122 let items_per_block = self.pipeline.vt * self.pipeline.block_size;
123
124 let aux_count = current.count.div_ceil(items_per_block);
125 let aux_size = (aux_count * 4) as u64;
126 let aux_offset = crate::common::math::align_to(current_scratch_offset, 256);
127
128 self.pipeline.dispatch(
129 &self.device,
130 encoder,
131 &self.pipeline.scan_pipeline,
132 (current.buf, current.offset),
133 (scratch, aux_offset),
134 current.count,
135 );
136
137 levels.push(Level {
138 buf: scratch,
139 offset: aux_offset,
140 count: aux_count,
141 });
142 current_scratch_offset = aux_offset + aux_size;
143 }
144
145 for i in (0..levels.len() - 1).rev() {
146 let data_level = &levels[i];
147 let aux_level = &levels[i + 1];
148
149 self.pipeline.dispatch(
150 &self.device,
151 encoder,
152 &self.pipeline.add_pipeline,
153 (data_level.buf, data_level.offset),
154 (aux_level.buf, aux_level.offset),
155 data_level.count,
156 );
157 }
158
159 Ok(())
160 }
161
162 fn prepare_scratch(&mut self, num_items: u32) {
163 let needed_bytes = self.pipeline.get_scratch_size(num_items);
164 if self.scratch_buffer.is_none() || needed_bytes > self.scratch_size_bytes {
165 self.scratch_buffer = Some(self.device.create_buffer(&wgpu::BufferDescriptor {
166 label: Some("Scanner Scratch"),
167 size: needed_bytes,
168 usage: wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_SRC,
169 mapped_at_creation: false,
170 }));
171 self.scratch_size_bytes = needed_bytes;
172 }
173 }
174}