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#[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 const SRC_BUFFER_SIZE: NonZeroU64 = NonZeroU64::new(size_of::<u32>() as u64 * 3).unwrap();
80
81 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 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 #[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 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 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 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 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 let binding_size = 2 * alignment + (buffer_size % alignment);
428 binding_size.min(buffer_size)
429}