Skip to main content

wgpu_primitives/scan/
scanner.rs

1use super::pipeline::ScanPipeline;
2use crate::{Error, common, context::Context};
3
4/// Performs an inclusive unsigned 32-bit prefix scan on a wgpu device.
5pub 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    /// Creates a scanner that submits work through an existing wgpu device and queue.
15    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    /// Creates a scanner from the crate's optional convenience context.
26    pub fn from_context(ctx: &Context) -> Self {
27        Self::new(&ctx.device, &ctx.queue)
28    }
29
30    /// Uploads values, scans them on the GPU, and downloads the inclusive prefixes.
31    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    /// Scans caller-owned GPU buffers and submits the work immediately.
48    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    /// Records a GPU prefix scan without submitting or waiting for the work.
63    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}