git.lucas.co / cce-ui
GPU-accelerated UI toolkit (Vulkan)
git clone https://git.lucas.co/cce-ui.git

src/vk/compute.rs (27.9K)

  1 //! The compute-job API: upload buffers, dispatch a WGSL kernel, read back.
  2 //!
  3 //! The renderer already had everything a compute job needs — naga compiles
  4 //! WGSL to SPIR-V at runtime, the path tracer builds compute pipelines over
  5 //! storage buffers, and `RtOffscreen` runs a headless device with no window
  6 //! — but all of it was internal to the RT pass and read back only an image.
  7 //! This module is that machinery with a general face: a [`ComputeDevice`]
  8 //! owns a headless [`VkCore`], and [`ComputeDevice::run`] takes a
  9 //! [`Kernel`] and a list of [`Binding`]s, uploads them, dispatches, waits,
 10 //! and copies every read-write binding back into the caller's slice.
 11 //!
 12 //! Designed for cce-designer's Phase 7 step 4 (see its `shapeshifter.md`):
 13 //! the per-point solver operators — relax, diffuse, collide — written once
 14 //! in WGSL over the columnar attribute arrays a `Detail` already keeps, with
 15 //! the CPU evaluator as the reference each is held to. The shape of the API
 16 //! follows from that use: the caller has arrays in memory and wants them
 17 //! transformed, so buffers are HOST-VISIBLE and mapped, upload and readback
 18 //! are memcpys through the mapping, and there is no staging copy. On an
 19 //! integrated GPU that is the fastest path there is; on a discrete one it is
 20 //! correct and simple, and a device-local tier can be added behind the same
 21 //! API if a workload ever asks for it.
 22 //!
 23 //! Every failure is an `Err(String)`, never a panic, because a kernel may be
 24 //! user-authored: a WGSL error comes back with naga's own diagnostic, a
 25 //! missing entry point names what the module does offer, and a device
 26 //! without Vulkan reports as such from [`ComputeDevice::new`].
 27 //!
 28 //! Not `Send`: it owns a device and a command buffer. Make one per thread
 29 //! that computes, and keep it — pipelines cache by source and entry point,
 30 //! and buffers are reused across runs when they fit.
 31 
 32 use super::core::VkCore;
 33 use super::renderer::{create_cpu_buffer, destroy_cpu_buffer, AllocatedBuffer};
 34 use ash::vk;
 35 use std::collections::HashMap;
 36 use std::ffi::CString;
 37 
 38 pub use crate::compute::{workgroups, BindKind, Binding, Kernel, MAX_BINDINGS};
 39 use crate::compute::{check_job, padded_len, parse_kernel, result_slot, slot_for};
 40 
 41 #[derive(Clone, PartialEq, Eq, Hash)]
 42 struct PipelineKey {
 43     kernel: Kernel,
 44     kinds: Vec<BindKind>,
 45 }
 46 
 47 struct Pipeline {
 48     set_layout: vk::DescriptorSetLayout,
 49     layout: vk::PipelineLayout,
 50     module: vk::ShaderModule,
 51     pipeline: vk::Pipeline,
 52     workgroup_size: [u32; 3],
 53 }
 54 
 55 /// A headless device that runs compute jobs. See the module docs.
 56 pub struct ComputeDevice {
 57     pipelines: HashMap<PipelineKey, Pipeline>,
 58     /// One buffer per binding index, grown when a job needs more room.
 59     slots: Vec<AllocatedBuffer>,
 60     descriptor_pool: vk::DescriptorPool,
 61     cmd: vk::CommandBuffer,
 62     fence: vk::Fence,
 63     /// Declared last: everything above is destroyed before the device.
 64     core: VkCore,
 65 }
 66 
 67 impl ComputeDevice {
 68     /// A device on the machine's preferred GPU (`CCE_VK_DEVICE` steers it,
 69     /// as for every renderer). `Err` when there is no usable Vulkan at all.
 70     pub fn new() -> Result<Self, String> {
 71         // VkCore reports an absent driver by panicking, as a renderer with no
 72         // window to draw into has nothing better to do; a compute consumer
 73         // has a CPU path to fall back to, so the panic is caught here.
 74         let core = std::panic::catch_unwind(VkCore::new_headless).map_err(|e| {
 75             let msg = e
 76                 .downcast_ref::<String>()
 77                 .cloned()
 78                 .or_else(|| e.downcast_ref::<&str>().map(|s| s.to_string()))
 79                 .unwrap_or_else(|| "unknown".to_string());
 80             format!("no Vulkan compute device: {msg}")
 81         })?;
 82         let device = core.device.clone();
 83         unsafe {
 84             let cmd = device
 85                 .allocate_command_buffers(
 86                     &vk::CommandBufferAllocateInfo::default()
 87                         .command_pool(core.command_pool)
 88                         .level(vk::CommandBufferLevel::PRIMARY)
 89                         .command_buffer_count(1),
 90                 )
 91                 .map_err(|e| format!("command buffer: {e}"))?[0];
 92             let fence = device
 93                 .create_fence(&vk::FenceCreateInfo::default(), None)
 94                 .map_err(|e| format!("fence: {e}"))?;
 95             let pool_sizes = [
 96                 vk::DescriptorPoolSize::default()
 97                     .ty(vk::DescriptorType::STORAGE_BUFFER)
 98                     .descriptor_count(2 * MAX_BINDINGS as u32),
 99                 vk::DescriptorPoolSize::default()
100                     .ty(vk::DescriptorType::UNIFORM_BUFFER)
101                     .descriptor_count(2 * MAX_BINDINGS as u32),
102             ];
103             let descriptor_pool = device
104                 .create_descriptor_pool(
105                     &vk::DescriptorPoolCreateInfo::default().max_sets(2).pool_sizes(&pool_sizes),
106                     None,
107                 )
108                 .map_err(|e| format!("descriptor pool: {e}"))?;
109             Ok(ComputeDevice {
110                 pipelines: HashMap::new(),
111                 slots: Vec::new(),
112                 descriptor_pool,
113                 cmd,
114                 fence,
115                 core,
116             })
117         }
118     }
119 
120     /// The physical device's name, for a log line or a status readout.
121     pub fn device_name(&self) -> String {
122         unsafe {
123             let props = self.core.instance.get_physical_device_properties(self.core.physical_device);
124             std::ffi::CStr::from_ptr(props.device_name.as_ptr()).to_string_lossy().into_owned()
125         }
126     }
127 
128     /// The entry point's `@workgroup_size`, compiling the kernel if needed.
129     pub fn workgroup_size(&mut self, kernel: &Kernel, kinds: &[BindKind]) -> Result<[u32; 3], String> {
130         let key = PipelineKey { kernel: kernel.clone(), kinds: kinds.to_vec() };
131         Ok(self.pipeline(&key)?.workgroup_size)
132     }
133 
134     /// Run the kernel over `items` invocations along x — the common case, a
135     /// job that is one invocation per element — computing the workgroup
136     /// count from the entry point's own `@workgroup_size`. A kernel should
137     /// still guard `id.x < arrayLength(...)`: the last group is padded.
138     pub fn run_over(&mut self, kernel: &Kernel, bindings: &mut [Binding<'_>], items: u32) -> Result<(), String> {
139         let kinds: Vec<BindKind> = bindings.iter().map(Binding::kind).collect();
140         let wg = self.workgroup_size(kernel, &kinds)?;
141         self.run(kernel, bindings, [workgroups(items, wg[0]), 1, 1])
142     }
143 
144     /// Upload every binding, dispatch `groups` workgroups of the kernel, wait
145     /// for the GPU, and read every [`Binding::Storage`] back into its slice.
146     pub fn run(&mut self, kernel: &Kernel, bindings: &mut [Binding<'_>], groups: [u32; 3]) -> Result<(), String> {
147         self.execute(kernel, bindings, groups, 1, None)
148     }
149 
150     /// [`run_passes`](Self::run_passes) with the dispatch sized from the entry
151     /// point's `@workgroup_size` over `items`, like [`run_over`](Self::run_over).
152     pub fn run_passes_over(
153         &mut self,
154         kernel: &Kernel,
155         bindings: &mut [Binding<'_>],
156         items: u32,
157         passes: u32,
158         ping_pong: Option<(usize, usize)>,
159     ) -> Result<(), String> {
160         let kinds: Vec<BindKind> = bindings.iter().map(Binding::kind).collect();
161         let wg = self.workgroup_size(kernel, &kinds)?;
162         self.execute(kernel, bindings, [workgroups(items, wg[0]), 1, 1], passes, ping_pong)
163     }
164 
165     /// `passes` dispatches of the kernel in ONE submission — uploaded once,
166     /// a memory barrier between passes, waited on once, read back once —
167     /// which is what an iterative solve needs: measured on an integrated
168     /// GPU, a pass submitted on its own costs about half a millisecond of
169     /// round trip whatever its size, and sixteen of those lose to the CPU
170     /// at every mesh size a designer works at.
171     ///
172     /// `ping_pong = Some((a, b))` makes passes alternate the roles of two
173     /// bindings: `a` must be a [`Binding::Input`] (the first pass reads it)
174     /// and `b` a [`Binding::Storage`] of the same length (the first pass
175     /// writes it); the second pass reads `b` and writes `a`'s buffer, and so
176     /// on. Whichever buffer the LAST pass wrote is read back into `b`'s
177     /// slice, so the caller always finds the result where it bound the
178     /// output. A Jacobi solve is exactly this shape.
179     pub fn run_passes(
180         &mut self,
181         kernel: &Kernel,
182         bindings: &mut [Binding<'_>],
183         groups: [u32; 3],
184         passes: u32,
185         ping_pong: Option<(usize, usize)>,
186     ) -> Result<(), String> {
187         self.execute(kernel, bindings, groups, passes, ping_pong)
188     }
189 
190     fn execute(
191         &mut self,
192         kernel: &Kernel,
193         bindings: &mut [Binding<'_>],
194         groups: [u32; 3],
195         passes: u32,
196         ping_pong: Option<(usize, usize)>,
197     ) -> Result<(), String> {
198         check_job(bindings, groups, passes, ping_pong)?;
199         let kinds: Vec<BindKind> = bindings.iter().map(Binding::kind).collect();
200         let key = PipelineKey { kernel: kernel.clone(), kinds };
201         let (pipeline, layout, set_layout) = {
202             let p = self.pipeline(&key)?;
203             (p.pipeline, p.layout, p.set_layout)
204         };
205 
206         // Buffers: one per binding index, reused when big enough. Uploads are
207         // memcpys through the persistent mapping.
208         let device = self.core.device.clone();
209         let mut sizes = Vec::with_capacity(bindings.len());
210         for (i, b) in bindings.iter().enumerate() {
211             let bytes = b.bytes();
212             let padded = padded_len(bytes.len());
213             if b.kind() == BindKind::Uniform {
214                 let cap = unsafe {
215                     self.core.instance.get_physical_device_properties(self.core.physical_device).limits.max_uniform_buffer_range
216                 } as usize;
217                 if padded > cap {
218                     return Err(format!("binding {i}: a uniform block of {} bytes exceeds the device's {cap}", bytes.len()));
219                 }
220             }
221             if i >= self.slots.len() {
222                 self.slots.push(AllocatedBuffer::null());
223             }
224             if (self.slots[i].size as usize) < padded {
225                 let allocator = self.core.allocator.as_mut().unwrap();
226                 destroy_cpu_buffer(&device, allocator, &mut self.slots[i]);
227                 self.slots[i] = create_cpu_buffer(
228                     &device,
229                     allocator,
230                     padded as vk::DeviceSize,
231                     vk::BufferUsageFlags::STORAGE_BUFFER | vk::BufferUsageFlags::UNIFORM_BUFFER,
232                     "compute-binding",
233                 );
234             }
235             let mapped = self.slots[i]
236                 .allocation
237                 .as_mut()
238                 .and_then(|a| a.mapped_slice_mut())
239                 .ok_or_else(|| format!("binding {i}: buffer memory is not host-visible"))?;
240             mapped[..bytes.len()].copy_from_slice(bytes);
241             // The padding is defined too, so an `arrayLength` that counts it
242             // reads zeros rather than whatever the last job left there.
243             for b in &mut mapped[bytes.len()..padded] {
244                 *b = 0;
245             }
246             sizes.push(padded);
247         }
248 
249         unsafe {
250             // One descriptor set per run, from a pool reset each time.
251             device
252                 .reset_descriptor_pool(self.descriptor_pool, vk::DescriptorPoolResetFlags::empty())
253                 .map_err(|e| format!("descriptor pool reset: {e}"))?;
254             // One descriptor set, or two with the ping-pong pair swapped in
255             // the second, so alternate passes bind the buffers the other
256             // way round without a write between dispatches.
257             let set_count = if ping_pong.is_some() { 2 } else { 1 };
258             let set_layouts = vec![set_layout; set_count];
259             let sets = device
260                 .allocate_descriptor_sets(
261                     &vk::DescriptorSetAllocateInfo::default()
262                         .descriptor_pool(self.descriptor_pool)
263                         .set_layouts(&set_layouts),
264                 )
265                 .map_err(|e| format!("descriptor set: {e}"))?;
266             let mut infos: Vec<[vk::DescriptorBufferInfo; 1]> = Vec::with_capacity(set_count * bindings.len());
267             for (si, _) in sets.iter().enumerate() {
268                 for i in 0..bindings.len() {
269                     let slot = slot_for(i, si == 1, ping_pong);
270                     infos.push([vk::DescriptorBufferInfo::default()
271                         .buffer(self.slots[slot].buffer)
272                         .offset(0)
273                         .range(sizes[slot] as vk::DeviceSize)]);
274                 }
275             }
276             let mut writes: Vec<vk::WriteDescriptorSet> = Vec::with_capacity(infos.len());
277             for (si, set) in sets.iter().enumerate() {
278                 for (i, b) in bindings.iter().enumerate() {
279                     writes.push(
280                         vk::WriteDescriptorSet::default()
281                             .dst_set(*set)
282                             .dst_binding(i as u32)
283                             .descriptor_type(match b.kind() {
284                                 BindKind::Storage => vk::DescriptorType::STORAGE_BUFFER,
285                                 BindKind::Uniform => vk::DescriptorType::UNIFORM_BUFFER,
286                             })
287                             .buffer_info(&infos[si * bindings.len() + i]),
288                     );
289                 }
290             }
291             device.update_descriptor_sets(&writes, &[]);
292 
293             device
294                 .begin_command_buffer(
295                     self.cmd,
296                     &vk::CommandBufferBeginInfo::default().flags(vk::CommandBufferUsageFlags::ONE_TIME_SUBMIT),
297                 )
298                 .map_err(|e| format!("begin: {e}"))?;
299             device.cmd_bind_pipeline(self.cmd, vk::PipelineBindPoint::COMPUTE, pipeline);
300             for pass in 0..passes {
301                 let set = sets[if ping_pong.is_some() && pass % 2 == 1 { 1 } else { 0 }];
302                 device.cmd_bind_descriptor_sets(self.cmd, vk::PipelineBindPoint::COMPUTE, layout, 0, &[set], &[]);
303                 device.cmd_dispatch(self.cmd, groups[0], groups[1], groups[2]);
304                 if pass + 1 < passes {
305                     // The next pass reads what this one wrote.
306                     device.cmd_pipeline_barrier(
307                         self.cmd,
308                         vk::PipelineStageFlags::COMPUTE_SHADER,
309                         vk::PipelineStageFlags::COMPUTE_SHADER,
310                         vk::DependencyFlags::empty(),
311                         &[vk::MemoryBarrier::default()
312                             .src_access_mask(vk::AccessFlags::SHADER_WRITE)
313                             .dst_access_mask(vk::AccessFlags::SHADER_READ | vk::AccessFlags::SHADER_WRITE)],
314                         &[],
315                         &[],
316                     );
317                 }
318             }
319             // Shader writes become host-visible before the fence is signalled.
320             device.cmd_pipeline_barrier(
321                 self.cmd,
322                 vk::PipelineStageFlags::COMPUTE_SHADER,
323                 vk::PipelineStageFlags::HOST,
324                 vk::DependencyFlags::empty(),
325                 &[vk::MemoryBarrier::default()
326                     .src_access_mask(vk::AccessFlags::SHADER_WRITE)
327                     .dst_access_mask(vk::AccessFlags::HOST_READ)],
328                 &[],
329                 &[],
330             );
331             device.end_command_buffer(self.cmd).map_err(|e| format!("end: {e}"))?;
332 
333             let cmds = [self.cmd];
334             device
335                 .queue_submit(self.core.queue, &[vk::SubmitInfo::default().command_buffers(&cmds)], self.fence)
336                 .map_err(|e| format!("submit: {e}"))?;
337             device
338                 .wait_for_fences(&[self.fence], true, u64::MAX)
339                 .map_err(|e| format!("fence wait: {e}"))?;
340             device.reset_fences(&[self.fence]).map_err(|e| format!("fence reset: {e}"))?;
341         }
342 
343         // Read back the read-write bindings — the ping-pong output from
344         // whichever buffer the last pass wrote.
345         for (i, b) in bindings.iter_mut().enumerate() {
346             if let Binding::Storage(out) = b {
347                 let mapped = self.slots[result_slot(i, passes, ping_pong)]
348                     .allocation
349                     .as_ref()
350                     .and_then(|a| a.mapped_slice())
351                     .ok_or_else(|| format!("binding {i}: buffer memory is not host-visible"))?;
352                 out.copy_from_slice(&mapped[..out.len()]);
353             }
354         }
355         Ok(())
356     }
357 
358     fn pipeline(&mut self, key: &PipelineKey) -> Result<&Pipeline, String> {
359         if !self.pipelines.contains_key(key) {
360             let built = self.build_pipeline(key)?;
361             self.pipelines.insert(key.clone(), built);
362         }
363         Ok(&self.pipelines[key])
364     }
365 
366     fn build_pipeline(&self, key: &PipelineKey) -> Result<Pipeline, String> {
367         let (spirv, workgroup_size) = compile_kernel(&key.kernel)?;
368         let device = &self.core.device;
369         unsafe {
370             let bindings: Vec<vk::DescriptorSetLayoutBinding> = key
371                 .kinds
372                 .iter()
373                 .enumerate()
374                 .map(|(i, k)| {
375                     vk::DescriptorSetLayoutBinding::default()
376                         .binding(i as u32)
377                         .descriptor_type(match k {
378                             BindKind::Storage => vk::DescriptorType::STORAGE_BUFFER,
379                             BindKind::Uniform => vk::DescriptorType::UNIFORM_BUFFER,
380                         })
381                         .descriptor_count(1)
382                         .stage_flags(vk::ShaderStageFlags::COMPUTE)
383                 })
384                 .collect();
385             let set_layout = device
386                 .create_descriptor_set_layout(&vk::DescriptorSetLayoutCreateInfo::default().bindings(&bindings), None)
387                 .map_err(|e| format!("descriptor set layout: {e}"))?;
388             let set_layouts = [set_layout];
389             let layout = match device.create_pipeline_layout(&vk::PipelineLayoutCreateInfo::default().set_layouts(&set_layouts), None) {
390                 Ok(l) => l,
391                 Err(e) => {
392                     device.destroy_descriptor_set_layout(set_layout, None);
393                     return Err(format!("pipeline layout: {e}"));
394                 }
395             };
396             let module = match device.create_shader_module(&vk::ShaderModuleCreateInfo::default().code(&spirv), None) {
397                 Ok(m) => m,
398                 Err(e) => {
399                     device.destroy_pipeline_layout(layout, None);
400                     device.destroy_descriptor_set_layout(set_layout, None);
401                     return Err(format!("shader module: {e}"));
402                 }
403             };
404             let entry = CString::new(key.kernel.entry.as_str()).map_err(|e| format!("entry point name: {e}"))?;
405             let pipeline = device.create_compute_pipelines(
406                 vk::PipelineCache::null(),
407                 &[vk::ComputePipelineCreateInfo::default()
408                     .stage(
409                         vk::PipelineShaderStageCreateInfo::default()
410                             .stage(vk::ShaderStageFlags::COMPUTE)
411                             .module(module)
412                             .name(&entry),
413                     )
414                     .layout(layout)],
415                 None,
416             );
417             let pipeline = match pipeline {
418                 Ok(p) => p[0],
419                 Err((_, e)) => {
420                     device.destroy_shader_module(module, None);
421                     device.destroy_pipeline_layout(layout, None);
422                     device.destroy_descriptor_set_layout(set_layout, None);
423                     return Err(format!("compute pipeline: {e}"));
424                 }
425             };
426             Ok(Pipeline { set_layout, layout, module, pipeline, workgroup_size })
427         }
428     }
429 }
430 
431 impl Drop for ComputeDevice {
432     fn drop(&mut self) {
433         let device = self.core.device.clone();
434         unsafe {
435             let _ = device.device_wait_idle();
436             for (_, p) in self.pipelines.drain() {
437                 device.destroy_pipeline(p.pipeline, None);
438                 device.destroy_shader_module(p.module, None);
439                 device.destroy_pipeline_layout(p.layout, None);
440                 device.destroy_descriptor_set_layout(p.set_layout, None);
441             }
442             device.destroy_descriptor_pool(self.descriptor_pool, None);
443             device.destroy_fence(self.fence, None);
444             // The command buffer dies with the pool in VkCore's Drop.
445         }
446         let allocator = self.core.allocator.as_mut().unwrap();
447         for slot in &mut self.slots {
448             destroy_cpu_buffer(&device, allocator, slot);
449         }
450     }
451 }
452 
453 /// WGSL to SPIR-V with every failure reported, plus the entry point's
454 /// workgroup size. `compile_wgsl` in the renderer panics on a bad shader,
455 /// which is right for the toolkit's own; a kernel here may be a user's.
456 fn compile_kernel(kernel: &Kernel) -> Result<(Vec<u32>, [u32; 3]), String> {
457     let parsed = parse_kernel(kernel)?;
458     let options = naga::back::spv::Options {
459         lang_version: (1, 0),
460         flags: naga::back::spv::WriterFlags::LABEL_VARYINGS,
461         ..Default::default()
462     };
463     let spirv =
464         naga::back::spv::write_vec(&parsed.module, &parsed.info, &options, None).map_err(|e| format!("SPIR-V: {e}"))?;
465     Ok((spirv, parsed.workgroup_size))
466 }
467 
468 #[cfg(test)]
469 mod tests {
470     use super::*;
471 
472     /// A device, or None with a note: these tests run on whatever Vulkan
473     /// the machine has (lavapipe counts) and skip where there is none.
474     fn device() -> Option<ComputeDevice> {
475         match ComputeDevice::new() {
476             Ok(d) => Some(d),
477             Err(e) => {
478                 println!("skipping compute test: {e}");
479                 None
480             }
481         }
482     }
483 
484     const DOUBLE: &str = r#"
485 @group(0) @binding(0) var<storage, read_write> data: array<f32>;
486 @compute @workgroup_size(64)
487 fn main(@builtin(global_invocation_id) id: vec3<u32>) {
488     let i = id.x;
489     if (i < arrayLength(&data)) {
490         data[i] = data[i] * 2.0;
491     }
492 }"#;
493 
494     #[test]
495     fn a_storage_binding_round_trips_through_the_kernel() {
496         let Some(mut dev) = device() else { return };
497         println!("compute on {}", dev.device_name());
498         // 1001 floats: not a multiple of the 16-byte padding, so the last
499         // element sits beside padding and must still come back doubled.
500         let mut data: Vec<f32> = (0..1001).map(|i| i as f32 * 0.5).collect();
501         let kernel = Kernel::new(DOUBLE, "main");
502         dev.run_over(&kernel, &mut [Binding::rw(&mut data)], 1001).unwrap();
503         for (i, v) in data.iter().enumerate() {
504             assert_eq!(*v, i as f32, "element {i}");
505         }
506         // Again, on the same device: the pipeline and the buffer are reused.
507         dev.run_over(&kernel, &mut [Binding::rw(&mut data)], 1001).unwrap();
508         assert_eq!(data[1000], 2000.0);
509         assert_eq!(dev.pipelines.len(), 1, "one pipeline for one kernel");
510         assert_eq!(dev.slots.len(), 1);
511         assert_eq!(dev.workgroup_size(&kernel, &[BindKind::Storage]).unwrap(), [64, 1, 1]);
512     }
513 
514     #[test]
515     fn inputs_and_uniforms_bind_beside_the_output() {
516         let Some(mut dev) = device() else { return };
517         #[repr(C)]
518         #[derive(Clone, Copy, bytemuck::Pod, bytemuck::Zeroable)]
519         struct Params {
520             scale: f32,
521             offset: f32,
522             _pad: [f32; 2],
523         }
524         const AXPY: &str = r#"
525 struct Params { scale: f32, offset: f32, pad: vec2<f32> }
526 @group(0) @binding(0) var<storage, read> a: array<vec4<f32>>;
527 @group(0) @binding(1) var<uniform> params: Params;
528 @group(0) @binding(2) var<storage, read_write> out: array<vec4<f32>>;
529 @compute @workgroup_size(32)
530 fn axpy(@builtin(global_invocation_id) id: vec3<u32>) {
531     let i = id.x;
532     if (i < arrayLength(&out)) {
533         out[i] = a[i] * params.scale + vec4<f32>(params.offset);
534     }
535 }"#;
536         let a: Vec<[f32; 4]> = (0..300).map(|i| [i as f32; 4]).collect();
537         let mut out = vec![[0.0f32; 4]; 300];
538         let params = Params { scale: 3.0, offset: 1.0, _pad: [0.0; 2] };
539         dev.run_over(
540             &Kernel::new(AXPY, "axpy"),
541             &mut [Binding::input(&a), Binding::uniform(&params), Binding::rw(&mut out)],
542             300,
543         )
544         .unwrap();
545         for (i, v) in out.iter().enumerate() {
546             assert_eq!(*v, [i as f32 * 3.0 + 1.0; 4], "element {i}");
547         }
548         assert!(a.iter().enumerate().all(|(i, v)| *v == [i as f32; 4]), "an input is never written back");
549     }
550 
551     #[test]
552     fn a_bad_kernel_is_an_error_not_a_panic() {
553         let Some(mut dev) = device() else { return };
554         let mut data = vec![1.0f32; 4];
555         let err = dev.run_over(&Kernel::new("fn main( {", "main"), &mut [Binding::rw(&mut data)], 4).unwrap_err();
556         assert!(err.starts_with("WGSL parse error"), "{err}");
557         let err = dev.run_over(&Kernel::new(DOUBLE, "nope"), &mut [Binding::rw(&mut data)], 4).unwrap_err();
558         assert!(err.contains("`nope`") && err.contains("main"), "names the missing entry and the offer: {err}");
559         let typed = "@group(0) @binding(0) var<storage, read_write> d: array<f32>;\n@compute @workgroup_size(1) fn main() { d[0] = 1u; }";
560         let err = dev.run_over(&Kernel::new(typed, "main"), &mut [Binding::rw(&mut data)], 1).unwrap_err();
561         // naga's front end types as it parses, so a type error is a parse
562         // error; what matters is that it is a WGSL diagnostic, not a panic.
563         assert!(err.starts_with("WGSL") && err.contains("u32"), "{err}");
564         assert_eq!(data, vec![1.0; 4], "nothing ran");
565         // The device is still good after every failure.
566         dev.run_over(&Kernel::new(DOUBLE, "main"), &mut [Binding::rw(&mut data)], 4).unwrap();
567         assert_eq!(data, vec![2.0; 4]);
568     }
569 
570     /// Passes chain inside one submission, and a ping-pong pair alternates
571     /// so the result lands in the output slice whether the count is odd or
572     /// even.
573     #[test]
574     fn passes_chain_and_ping_pong_lands_in_the_output() {
575         let Some(mut dev) = device() else { return };
576         // In place: three doublings are one octupling.
577         let mut data: Vec<f32> = (0..500).map(|i| i as f32).collect();
578         dev.run_passes_over(&Kernel::new(DOUBLE, "main"), &mut [Binding::rw(&mut data)], 500, 3, None).unwrap();
579         assert!(data.iter().enumerate().all(|(i, v)| *v == i as f32 * 8.0));
580 
581         const COPY_DOUBLE: &str = r#"
582 @group(0) @binding(0) var<storage, read> src: array<f32>;
583 @group(0) @binding(1) var<storage, read_write> dst: array<f32>;
584 @compute @workgroup_size(64)
585 fn main(@builtin(global_invocation_id) id: vec3<u32>) {
586     let i = id.x;
587     if (i < arrayLength(&dst)) { dst[i] = src[i] * 2.0; }
588 }"#;
589         let src: Vec<f32> = (0..500).map(|i| i as f32).collect();
590         let kernel = Kernel::new(COPY_DOUBLE, "main");
591         for passes in [1u32, 2, 3, 4] {
592             let mut dst = vec![0.0f32; 500];
593             dev.run_passes_over(&kernel, &mut [Binding::input(&src), Binding::rw(&mut dst)], 500, passes, Some((0, 1))).unwrap();
594             let factor = 2f32.powi(passes as i32);
595             assert!(dst.iter().enumerate().all(|(i, v)| *v == i as f32 * factor), "{passes} passes give x{factor}");
596         }
597         assert!(src.iter().enumerate().all(|(i, v)| *v == i as f32), "the input slice is never written");
598 
599         let mut dst = vec![0.0f32; 500];
600         let err = dev.run_passes_over(&kernel, &mut [Binding::input(&src), Binding::rw(&mut dst)], 500, 2, Some((1, 0))).unwrap_err();
601         assert!(err.contains("must be a read-only Input"), "{err}");
602         let err = dev.run_passes_over(&kernel, &mut [Binding::input(&src), Binding::rw(&mut dst)], 500, 0, None).unwrap_err();
603         assert!(err.contains("at least one pass"), "{err}");
604     }
605 
606     #[test]
607     fn workgroup_arithmetic() {
608         assert_eq!(workgroups(0, 64), 1, "a dispatch of zero groups is invalid");
609         assert_eq!(workgroups(1, 64), 1);
610         assert_eq!(workgroups(64, 64), 1);
611         assert_eq!(workgroups(65, 64), 2);
612         assert_eq!(workgroups(1000, 0), 1000, "a zero group size is treated as one");
613     }
614 }