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(¶ms), 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 }