wgpu_core/indirect_validation/
dispatch.rs

1use super::CreateIndirectValidationPipelineError;
2use crate::{
3    device::DeviceError,
4    hal_label,
5    pipeline::{CreateComputePipelineError, CreateShaderModuleError},
6};
7use alloc::{boxed::Box, format, string::ToString as _};
8use core::num::NonZeroU64;
9use scopeguard::{guard, ScopeGuard};
10
11/// This machinery requires the following limits:
12///
13/// - max_bind_groups: 2,
14/// - max_dynamic_storage_buffers_per_pipeline_layout: 1,
15/// - max_storage_buffers_per_shader_stage: 2,
16/// - max_storage_buffer_binding_size: 3 * min_storage_buffer_offset_alignment,
17/// - max_immediate_size: 4,
18/// - max_compute_invocations_per_workgroup 1
19///
20/// These are all indirectly satisfied by `DownlevelFlags::INDIRECT_EXECUTION`, which is also
21/// required for this module's functionality to work.
22#[derive(Debug)]
23pub(crate) struct Dispatch {
24    module: Box<dyn hal::DynShaderModule>,
25    dst_bind_group_layout: Box<dyn hal::DynBindGroupLayout>,
26    src_bind_group_layout: Box<dyn hal::DynBindGroupLayout>,
27    pipeline_layout: Box<dyn hal::DynPipelineLayout>,
28    pipeline: Box<dyn hal::DynComputePipeline>,
29    dst_buffer: Box<dyn hal::DynBuffer>,
30    dst_bind_group: Box<dyn hal::DynBindGroup>,
31}
32
33pub struct Params<'a> {
34    pub pipeline_layout: &'a dyn hal::DynPipelineLayout,
35    pub pipeline: &'a dyn hal::DynComputePipeline,
36    pub dst_buffer: &'a dyn hal::DynBuffer,
37    pub dst_bind_group: &'a dyn hal::DynBindGroup,
38    pub aligned_offset: u64,
39    pub offset_remainder: u64,
40}
41
42impl Dispatch {
43    pub(super) fn new(
44        device: &dyn hal::DynDevice,
45        instance_flags: wgt::InstanceFlags,
46        limits: &wgt::Limits,
47    ) -> Result<Self, CreateIndirectValidationPipelineError> {
48        let max_compute_workgroups_per_dimension = limits.max_compute_workgroups_per_dimension;
49
50        let src = format!(
51            "
52            @group(0) @binding(0)
53            var<storage, read_write> dst: array<u32, 6>;
54            @group(1) @binding(0)
55            var<storage, read> src: array<u32>;
56            struct OffsetPc {{
57                inner: u32,
58            }}
59            var<immediate> offset: OffsetPc;
60
61            @compute @workgroup_size(1)
62            fn main() {{
63                let src = vec3(src[offset.inner], src[offset.inner + 1], src[offset.inner + 2]);
64                let max_compute_workgroups_per_dimension = {max_compute_workgroups_per_dimension}u;
65                if (
66                    src.x > max_compute_workgroups_per_dimension ||
67                    src.y > max_compute_workgroups_per_dimension ||
68                    src.z > max_compute_workgroups_per_dimension
69                ) {{
70                    dst = array(0u, 0u, 0u, 0u, 0u, 0u);
71                }} else {{
72                    dst = array(src.x, src.y, src.z, src.x, src.y, src.z);
73                }}
74            }}
75        "
76        );
77
78        // SAFETY: The value we are passing to `new_unchecked` is not zero, so this is safe.
79        const SRC_BUFFER_SIZE: NonZeroU64 = NonZeroU64::new(size_of::<u32>() as u64 * 3).unwrap();
80
81        // SAFETY: The value we are passing to `new_unchecked` is not zero, so this is safe.
82        const DST_BUFFER_SIZE: NonZeroU64 = NonZeroU64::new(SRC_BUFFER_SIZE.get() * 2).unwrap();
83
84        #[cfg(feature = "wgsl")]
85        let module = naga::front::wgsl::parse_str(&src).map_err(|inner| {
86            CreateShaderModuleError::Parsing(naga::error::ShaderError {
87                source: src.clone(),
88                label: None,
89                inner: Box::new(inner),
90            })
91        })?;
92        #[cfg(not(feature = "wgsl"))]
93        #[allow(clippy::diverging_sub_expression)]
94        let module = panic!("Indirect validation requires the wgsl feature flag to be enabled!");
95
96        let info = crate::device::create_validator(
97            wgt::Features::IMMEDIATES,
98            wgt::DownlevelFlags::empty(),
99            naga::valid::ValidationFlags::all(),
100        )
101        .validate(&module)
102        .map_err(|inner| {
103            CreateShaderModuleError::Validation(naga::error::ShaderError {
104                source: src,
105                label: None,
106                inner,
107            })
108        })?;
109        let hal_shader = hal::ShaderInput::Naga(hal::NagaShader {
110            module: alloc::borrow::Cow::Owned(module),
111            info,
112            debug_source: None,
113        });
114        let hal_desc = hal::ShaderModuleDescriptor {
115            label: hal_label(
116                Some("(wgpu internal) Indirect dispatch validation shader module"),
117                instance_flags,
118            ),
119            runtime_checks: wgt::ShaderRuntimeChecks::unchecked(),
120        };
121        let module =
122            unsafe { device.create_shader_module(&hal_desc, hal_shader) }.map_err(|error| {
123                match error {
124                    hal::ShaderError::Device(error) => {
125                        CreateShaderModuleError::Device(DeviceError::from_hal(error))
126                    }
127                    hal::ShaderError::Compilation(ref msg) => {
128                        log::error!("Shader error: {msg}");
129                        CreateShaderModuleError::Generation
130                    }
131                }
132            })?;
133        let module = guard(module, |module| unsafe {
134            device.destroy_shader_module(module)
135        });
136
137        let dst_bind_group_layout_desc = hal::BindGroupLayoutDescriptor {
138            label: hal_label(
139                Some("(wgpu internal) Indirect dispatch validation destination bind group layout"),
140                instance_flags,
141            ),
142            flags: hal::BindGroupLayoutFlags::empty(),
143            entries: &[wgt::BindGroupLayoutEntry {
144                binding: 0,
145                visibility: wgt::ShaderStages::COMPUTE,
146                ty: wgt::BindingType::Buffer {
147                    ty: wgt::BufferBindingType::Storage { read_only: false },
148                    has_dynamic_offset: false,
149                    min_binding_size: Some(DST_BUFFER_SIZE),
150                },
151                count: None,
152            }],
153        };
154        let dst_bind_group_layout = unsafe {
155            device
156                .create_bind_group_layout(&dst_bind_group_layout_desc)
157                .map_err(DeviceError::from_hal)?
158        };
159        let dst_bind_group_layout = guard(dst_bind_group_layout, |bgl| unsafe {
160            device.destroy_bind_group_layout(bgl)
161        });
162
163        let src_bind_group_layout_desc = hal::BindGroupLayoutDescriptor {
164            label: hal_label(
165                Some("(wgpu internal) Indirect dispatch validation source bind group layout"),
166                instance_flags,
167            ),
168            flags: hal::BindGroupLayoutFlags::empty(),
169            entries: &[wgt::BindGroupLayoutEntry {
170                binding: 0,
171                visibility: wgt::ShaderStages::COMPUTE,
172                ty: wgt::BindingType::Buffer {
173                    ty: wgt::BufferBindingType::Storage { read_only: true },
174                    has_dynamic_offset: true,
175                    min_binding_size: Some(SRC_BUFFER_SIZE),
176                },
177                count: None,
178            }],
179        };
180        let src_bind_group_layout = unsafe {
181            device
182                .create_bind_group_layout(&src_bind_group_layout_desc)
183                .map_err(DeviceError::from_hal)?
184        };
185        let src_bind_group_layout = guard(src_bind_group_layout, |bgl| unsafe {
186            device.destroy_bind_group_layout(bgl)
187        });
188
189        let pipeline_layout_desc = hal::PipelineLayoutDescriptor {
190            label: hal_label(
191                Some("(wgpu internal) Indirect dispatch validation pipeline layout"),
192                instance_flags,
193            ),
194            flags: hal::PipelineLayoutFlags::empty(),
195            bind_group_layouts: &[
196                Some(dst_bind_group_layout.as_ref()),
197                Some(src_bind_group_layout.as_ref()),
198            ],
199            immediate_size: 4,
200        };
201        let pipeline_layout = unsafe {
202            device
203                .create_pipeline_layout(&pipeline_layout_desc)
204                .map_err(DeviceError::from_hal)?
205        };
206        let pipeline_layout = guard(pipeline_layout, |pipeline_layout| unsafe {
207            device.destroy_pipeline_layout(pipeline_layout)
208        });
209
210        let pipeline_desc = hal::ComputePipelineDescriptor {
211            label: hal_label(
212                Some("(wgpu internal) Indirect dispatch validation pipeline"),
213                instance_flags,
214            ),
215            layout: pipeline_layout.as_ref(),
216            stage: hal::ProgrammableStage {
217                module: module.as_ref(),
218                entry_point: "main",
219                constants: &Default::default(),
220                zero_initialize_workgroup_memory: false,
221            },
222            cache: None,
223        };
224        let pipeline =
225            unsafe { device.create_compute_pipeline(&pipeline_desc) }.map_err(|err| match err {
226                hal::PipelineError::Device(error) => {
227                    CreateComputePipelineError::Device(DeviceError::from_hal(error))
228                }
229                hal::PipelineError::Linkage(_stages, msg) => {
230                    CreateComputePipelineError::Internal(msg)
231                }
232                hal::PipelineError::EntryPoint(_stage) => CreateComputePipelineError::Internal(
233                    crate::device::ENTRYPOINT_FAILURE_ERROR.to_string(),
234                ),
235                hal::PipelineError::PipelineConstants(_, error) => {
236                    CreateComputePipelineError::PipelineConstants(error)
237                }
238            })?;
239        let pipeline = guard(pipeline, |pipeline| unsafe {
240            device.destroy_compute_pipeline(pipeline)
241        });
242
243        let dst_buffer_desc = hal::BufferDescriptor {
244            label: hal_label(
245                Some("(wgpu internal) Indirect dispatch validation destination buffer"),
246                instance_flags,
247            ),
248            size: DST_BUFFER_SIZE.get(),
249            usage: wgt::BufferUses::INDIRECT | wgt::BufferUses::STORAGE_READ_WRITE,
250            memory_flags: hal::MemoryFlags::empty(),
251        };
252        let dst_buffer =
253            unsafe { device.create_buffer(&dst_buffer_desc) }.map_err(DeviceError::from_hal)?;
254        let dst_buffer = guard(dst_buffer, |buffer| unsafe {
255            device.destroy_buffer(buffer)
256        });
257
258        let dst_bind_group_desc = hal::BindGroupDescriptor {
259            label: hal_label(
260                Some("(wgpu internal) Indirect dispatch validation destination bind group"),
261                instance_flags,
262            ),
263            layout: dst_bind_group_layout.as_ref(),
264            entries: &[hal::BindGroupEntry {
265                binding: 0,
266                resource_index: 0,
267                count: 1,
268            }],
269            // SAFETY: We just created the buffer with this size.
270            buffers: &[hal::BufferBinding::new_unchecked(
271                dst_buffer.as_ref(),
272                0,
273                Some(DST_BUFFER_SIZE),
274            )],
275            samplers: &[],
276            textures: &[],
277            acceleration_structures: &[],
278            external_textures: &[],
279        };
280        let dst_bind_group = unsafe {
281            device
282                .create_bind_group(&dst_bind_group_desc)
283                .map_err(DeviceError::from_hal)
284        }?;
285
286        // Error returns after we start consuming guards could bypass resource cleanup.
287        #[deny(clippy::question_mark_used)]
288        Ok(Self {
289            module: ScopeGuard::into_inner(module),
290            dst_bind_group_layout: ScopeGuard::into_inner(dst_bind_group_layout),
291            src_bind_group_layout: ScopeGuard::into_inner(src_bind_group_layout),
292            pipeline_layout: ScopeGuard::into_inner(pipeline_layout),
293            pipeline: ScopeGuard::into_inner(pipeline),
294            dst_buffer: ScopeGuard::into_inner(dst_buffer),
295            dst_bind_group,
296        })
297    }
298
299    /// `Ok(None)` will only be returned if `buffer_size` is `0`.
300    pub(super) fn create_src_bind_group(
301        &self,
302        device: &dyn hal::DynDevice,
303        limits: &wgt::Limits,
304        buffer_size: u64,
305        buffer: &dyn hal::DynBuffer,
306        instance_flags: wgt::InstanceFlags,
307    ) -> Result<Option<Box<dyn hal::DynBindGroup>>, DeviceError> {
308        let binding_size = calculate_src_buffer_binding_size(buffer_size, limits);
309        let Some(binding_size) = NonZeroU64::new(binding_size) else {
310            return Ok(None);
311        };
312        let hal_desc = hal::BindGroupDescriptor {
313            label: hal_label(
314                Some("(wgpu internal) Indirect dispatch validation source bind group"),
315                instance_flags,
316            ),
317            layout: self.src_bind_group_layout.as_ref(),
318            entries: &[hal::BindGroupEntry {
319                binding: 0,
320                resource_index: 0,
321                count: 1,
322            }],
323            // SAFETY: We calculated the binding size to fit within the buffer.
324            buffers: &[hal::BufferBinding::new_unchecked(buffer, 0, binding_size)],
325            samplers: &[],
326            textures: &[],
327            acceleration_structures: &[],
328            external_textures: &[],
329        };
330        unsafe {
331            device
332                .create_bind_group(&hal_desc)
333                .map(Some)
334                .map_err(DeviceError::from_hal)
335        }
336    }
337
338    pub fn params<'a>(&'a self, limits: &wgt::Limits, offset: u64, buffer_size: u64) -> Params<'a> {
339        // The offset we receive is only required to be aligned to 4 bytes.
340        //
341        // Binding offsets and dynamic offsets are required to be aligned to
342        // min_storage_buffer_offset_alignment (256 bytes by default).
343        //
344        // So, we work around this limitation by calculating an aligned offset
345        // and pass the remainder through a immediate data.
346        //
347        // We could bind the whole buffer and only have to pass the offset
348        // through a immediate data but we might run into the
349        // max_storage_buffer_binding_size limit.
350        //
351        // See the inner docs of `calculate_src_buffer_binding_size` to
352        // see how we get the appropriate `binding_size`.
353        let alignment = limits.min_storage_buffer_offset_alignment as u64;
354        let binding_size = calculate_src_buffer_binding_size(buffer_size, limits);
355        let aligned_offset = offset - offset % alignment;
356        // This works because `binding_size` is either `buffer_size` or `alignment * 2 + buffer_size % alignment`.
357        let max_aligned_offset = buffer_size - binding_size;
358        let aligned_offset = aligned_offset.min(max_aligned_offset);
359        let offset_remainder = offset - aligned_offset;
360
361        Params {
362            pipeline_layout: self.pipeline_layout.as_ref(),
363            pipeline: self.pipeline.as_ref(),
364            dst_buffer: self.dst_buffer.as_ref(),
365            dst_bind_group: self.dst_bind_group.as_ref(),
366            aligned_offset,
367            offset_remainder,
368        }
369    }
370
371    pub(super) fn dispose(self, device: &dyn hal::DynDevice) {
372        let Dispatch {
373            module,
374            dst_bind_group_layout,
375            src_bind_group_layout,
376            pipeline_layout,
377            pipeline,
378            dst_buffer,
379            dst_bind_group,
380        } = self;
381
382        unsafe {
383            device.destroy_bind_group(dst_bind_group);
384            device.destroy_buffer(dst_buffer);
385            device.destroy_compute_pipeline(pipeline);
386            device.destroy_pipeline_layout(pipeline_layout);
387            device.destroy_bind_group_layout(src_bind_group_layout);
388            device.destroy_bind_group_layout(dst_bind_group_layout);
389            device.destroy_shader_module(module);
390        }
391    }
392}
393
394fn calculate_src_buffer_binding_size(buffer_size: u64, limits: &wgt::Limits) -> u64 {
395    let alignment = limits.min_storage_buffer_offset_alignment as u64;
396
397    // We need to choose a binding size that can address all possible sets of 12 contiguous bytes in the buffer taking
398    // into account that the dynamic offset needs to be a multiple of `min_storage_buffer_offset_alignment`.
399
400    // Given the know variables: `offset`, `buffer_size`, `alignment` and the rule `offset + 12 <= buffer_size`.
401
402    // Let `chunks = floor(buffer_size / alignment)`.
403    // Let `chunk` be the interval `[0, chunks]`.
404    // Let `offset = alignment * chunk + r` where `r` is the interval [0, alignment - 4].
405    // Let `binding` be the interval `[offset, offset + 12]`.
406    // Let `aligned_offset = alignment * chunk`.
407    // Let `aligned_binding` be the interval `[aligned_offset, aligned_offset + r + 12]`.
408    // Let `aligned_binding_size = r + 12 = [12, alignment + 8]`.
409    // Let `min_aligned_binding_size = alignment + 8`.
410
411    // `min_aligned_binding_size` is the minimum binding size required to address all 12 contiguous bytes in the buffer
412    // but the last aligned_offset + min_aligned_binding_size might overflow the buffer. In order to avoid this we must
413    // pick a larger `binding_size` that satisfies: `last_aligned_offset + binding_size = buffer_size` and
414    // `binding_size >= min_aligned_binding_size`.
415
416    // Let `buffer_size = alignment * chunks + sr` where `sr` is the interval [0, alignment - 4].
417    // Let `last_aligned_offset = alignment * (chunks - u)` where `u` is the interval [0, chunks].
418    // => `binding_size = buffer_size - last_aligned_offset`
419    // => `binding_size = alignment * chunks + sr - alignment * (chunks - u)`
420    // => `binding_size = alignment * chunks + sr - alignment * chunks + alignment * u`
421    // => `binding_size = sr + alignment * u`
422    // => `min_aligned_binding_size <= sr + alignment * u`
423    // => `alignment + 8 <= sr + alignment * u`
424    // => `u` must be at least 2
425    // => `binding_size = sr + alignment * 2`
426
427    let binding_size = 2 * alignment + (buffer_size % alignment);
428    binding_size.min(buffer_size)
429}