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

src/compute.rs (10.4K)

  1 //! What a compute job is, apart from the device that runs it: a [`Kernel`]
  2 //! (WGSL source and an entry point), its [`Binding`]s in `@binding(i)` order,
  3 //! and the rules a job is held to before anything reaches a GPU — the
  4 //! binding count, the dispatch, the ping-pong pair, the kernel parsed and
  5 //! validated by naga with its own diagnostics.
  6 //!
  7 //! Two devices run jobs: `vk::ComputeDevice` (a headless Vulkan device,
  8 //! synchronous) and, in the browser, `web::ComputeDevice` (WebGPU, whose
  9 //! readback is a promise, so its `run` is async). Both take these types and
 10 //! answer a job the same way, so a kernel and its bindings are written once.
 11 
 12 /// A compute shader: WGSL source and the `@compute` entry point to run.
 13 #[derive(Clone, Debug, PartialEq, Eq, Hash)]
 14 pub struct Kernel {
 15     pub source: String,
 16     pub entry: String,
 17 }
 18 
 19 impl Kernel {
 20     pub fn new(source: impl Into<String>, entry: impl Into<String>) -> Self {
 21         Kernel { source: source.into(), entry: entry.into() }
 22     }
 23 }
 24 
 25 /// How a binding is declared to the shader, in `@binding(i)` order.
 26 #[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
 27 pub enum BindKind {
 28     /// `var<storage, read>` or `var<storage, read_write>`.
 29     Storage,
 30     /// `var<uniform>`: a small parameter block, 16-byte layout rules apply.
 31     Uniform,
 32 }
 33 
 34 /// One buffer of a job, bound at `@group(0) @binding(i)` for its index in
 35 /// the list handed to a device's `run`.
 36 pub enum Binding<'a> {
 37     /// Read-write storage: uploaded before the dispatch and READ BACK into
 38     /// the same slice after it.
 39     Storage(&'a mut [u8]),
 40     /// Read-only storage: uploaded, never read back.
 41     Input(&'a [u8]),
 42     /// A uniform block: uploaded, never read back.
 43     Uniform(&'a [u8]),
 44 }
 45 
 46 impl<'a> Binding<'a> {
 47     /// A read-write binding over a typed slice (`&mut [f32]`, `&mut [[f32; 3]]`, …).
 48     pub fn rw<T: bytemuck::Pod>(data: &'a mut [T]) -> Self {
 49         Binding::Storage(bytemuck::cast_slice_mut(data))
 50     }
 51 
 52     /// A read-only storage binding over a typed slice.
 53     pub fn input<T: bytemuck::Pod>(data: &'a [T]) -> Self {
 54         Binding::Input(bytemuck::cast_slice(data))
 55     }
 56 
 57     /// A uniform binding over one `Pod` struct.
 58     pub fn uniform<T: bytemuck::Pod>(value: &'a T) -> Self {
 59         Binding::Uniform(bytemuck::bytes_of(value))
 60     }
 61 
 62     pub(crate) fn kind(&self) -> BindKind {
 63         match self {
 64             Binding::Storage(_) | Binding::Input(_) => BindKind::Storage,
 65             Binding::Uniform(_) => BindKind::Uniform,
 66         }
 67     }
 68 
 69     pub(crate) fn bytes(&self) -> &[u8] {
 70         match self {
 71             Binding::Storage(b) => b,
 72             Binding::Input(b) => b,
 73             Binding::Uniform(b) => b,
 74         }
 75     }
 76 }
 77 
 78 /// Workgroups needed to cover `items` at `per_group` invocations each — the
 79 /// `@workgroup_size` of the entry point, which a device's `run_over` reads
 80 /// for you.
 81 pub fn workgroups(items: u32, per_group: u32) -> u32 {
 82     items.div_ceil(per_group.max(1)).max(1)
 83 }
 84 
 85 /// The most bindings one job may carry.
 86 pub const MAX_BINDINGS: usize = 16;
 87 
 88 /// Storage bindings are bound whole, so a buffer's size has to be a multiple
 89 /// of the widest element stride a shader may declare; 16 covers `vec4<f32>`.
 90 pub(crate) const BUFFER_ALIGN: usize = 16;
 91 
 92 /// A binding of `len` bytes as the buffer it is uploaded into: at least one
 93 /// alignment unit, rounded up to one. The padding is uploaded as zeros.
 94 pub(crate) fn padded_len(len: usize) -> usize {
 95     len.max(BUFFER_ALIGN).div_ceil(BUFFER_ALIGN) * BUFFER_ALIGN
 96 }
 97 
 98 /// Hold a job to the rules before any device is asked: the binding count,
 99 /// a dispatch with no zero in it, at least one pass, and a ping-pong pair
100 /// that names a read-only input and a read-write output of one length.
101 pub(crate) fn check_job(
102     bindings: &[Binding<'_>],
103     groups: [u32; 3],
104     passes: u32,
105     ping_pong: Option<(usize, usize)>,
106 ) -> Result<(), String> {
107     if bindings.len() > MAX_BINDINGS {
108         return Err(format!("{} bindings; a job may carry at most {MAX_BINDINGS}", bindings.len()));
109     }
110     if groups.contains(&0) {
111         return Err(format!("workgroup count {groups:?} has a zero"));
112     }
113     if passes == 0 {
114         return Err("a job needs at least one pass".to_string());
115     }
116     if let Some((a, b)) = ping_pong {
117         if a == b || a >= bindings.len() || b >= bindings.len() {
118             return Err(format!("ping-pong pair ({a}, {b}) does not name two distinct bindings of {}", bindings.len()));
119         }
120         if !matches!(bindings[a], Binding::Input(_)) {
121             return Err(format!("ping-pong binding {a} must be a read-only Input: it is where the first pass reads"));
122         }
123         if !matches!(bindings[b], Binding::Storage(_)) {
124             return Err(format!("ping-pong binding {b} must be a read-write Storage: it is where the result lands"));
125         }
126         if bindings[a].bytes().len() != bindings[b].bytes().len() {
127             return Err(format!(
128                 "ping-pong bindings {a} and {b} differ in length ({} vs {} bytes)",
129                 bindings[a].bytes().len(),
130                 bindings[b].bytes().len()
131             ));
132         }
133     }
134     Ok(())
135 }
136 
137 /// The buffer bound at `binding` in a pass: a ping-pong pass with `swapped`
138 /// set (every second one) binds the pair's two buffers the other way round.
139 pub(crate) fn slot_for(binding: usize, swapped: bool, ping_pong: Option<(usize, usize)>) -> usize {
140     match ping_pong {
141         Some((a, b)) if swapped && binding == a => b,
142         Some((a, b)) if swapped && binding == b => a,
143         _ => binding,
144     }
145 }
146 
147 /// The buffer a read-write binding is read back from: the ping-pong output
148 /// from whichever buffer the LAST pass wrote, every other one from its own.
149 pub(crate) fn result_slot(binding: usize, passes: u32, ping_pong: Option<(usize, usize)>) -> usize {
150     match ping_pong {
151         Some((a, b)) if binding == b && passes.is_multiple_of(2) => a,
152         _ => binding,
153     }
154 }
155 
156 /// A kernel parsed and validated, with what a device needs to know of it.
157 pub(crate) struct ParsedKernel {
158     pub module: naga::Module,
159     #[cfg_attr(target_arch = "wasm32", allow(dead_code))]
160     pub info: naga::valid::ModuleInfo,
161     pub workgroup_size: [u32; 3],
162 }
163 
164 impl ParsedKernel {
165     /// Whether the module declares `@group(0) @binding(i)` as read-only
166     /// storage (`var<storage, read>`). WebGPU's layouts tell read-only from
167     /// read-write storage, where Vulkan's do not.
168     #[cfg_attr(not(target_arch = "wasm32"), allow(dead_code))]
169     pub fn read_only_storage(&self, binding: u32) -> bool {
170         self.module.global_variables.iter().any(|(_, g)| {
171             g.binding.as_ref().is_some_and(|b| b.group == 0 && b.binding == binding)
172                 && matches!(g.space, naga::AddressSpace::Storage { access } if !access.contains(naga::StorageAccess::STORE))
173         })
174     }
175 }
176 
177 /// Parse and validate a kernel, every failure reported with naga's own
178 /// diagnostic — a kernel may be a user's, so a bad one is an `Err`, never a
179 /// panic — and find its `@compute` entry point, naming what the module does
180 /// offer when there is none of that name.
181 pub(crate) fn parse_kernel(kernel: &Kernel) -> Result<ParsedKernel, String> {
182     let module = naga::front::wgsl::parse_str(&kernel.source)
183         .map_err(|e| format!("WGSL parse error: {}", e.emit_to_string(&kernel.source).trim_end()))?;
184     let entry = module
185         .entry_points
186         .iter()
187         .find(|ep| ep.name == kernel.entry && ep.stage == naga::ShaderStage::Compute)
188         .ok_or_else(|| {
189             let offered: Vec<&str> = module
190                 .entry_points
191                 .iter()
192                 .filter(|ep| ep.stage == naga::ShaderStage::Compute)
193                 .map(|ep| ep.name.as_str())
194                 .collect();
195             format!(
196                 "no @compute entry point named `{}`; the module offers {}",
197                 kernel.entry,
198                 if offered.is_empty() { "none".to_string() } else { offered.join(", ") }
199             )
200         })?;
201     let workgroup_size = entry.workgroup_size;
202     let info = naga::valid::Validator::new(naga::valid::ValidationFlags::all(), naga::valid::Capabilities::empty())
203         .validate(&module)
204         .map_err(|e| format!("WGSL validation error: {}", e.emit_to_string(&kernel.source).trim_end()))?;
205     Ok(ParsedKernel { module, info, workgroup_size })
206 }
207 
208 #[cfg(test)]
209 mod tests {
210     use super::*;
211 
212     #[test]
213     fn a_job_is_held_to_the_rules_before_a_device_sees_it() {
214         let mut out = [0.0f32; 4];
215         let inp = [0.0f32; 4];
216         let short = [0.0f32; 2];
217         assert!(check_job(&[Binding::rw(&mut out)], [1, 1, 1], 1, None).is_ok());
218         assert!(check_job(&[Binding::rw(&mut out)], [1, 0, 1], 1, None).unwrap_err().contains("zero"));
219         assert!(check_job(&[Binding::rw(&mut out)], [1, 1, 1], 0, None).unwrap_err().contains("pass"));
220         assert!(check_job(&[Binding::input(&inp), Binding::rw(&mut out)], [1, 1, 1], 2, Some((0, 1))).is_ok());
221         assert!(check_job(&[Binding::input(&inp), Binding::rw(&mut out)], [1, 1, 1], 2, Some((1, 0)))
222             .unwrap_err()
223             .contains("read-only Input"));
224         assert!(check_job(&[Binding::input(&short), Binding::rw(&mut out)], [1, 1, 1], 2, Some((0, 1)))
225             .unwrap_err()
226             .contains("differ in length"));
227     }
228 
229     #[test]
230     fn the_ping_pong_result_is_wherever_the_last_pass_wrote() {
231         let pp = Some((0, 1));
232         assert_eq!((slot_for(0, true, pp), slot_for(1, true, pp), slot_for(2, true, pp)), (1, 0, 2));
233         assert_eq!(slot_for(0, false, pp), 0);
234         assert_eq!(result_slot(1, 3, pp), 1);
235         assert_eq!(result_slot(1, 4, pp), 0);
236         assert_eq!(result_slot(1, 4, None), 1);
237     }
238 
239     #[test]
240     fn a_kernel_reports_its_workgroup_size_its_read_only_bindings_and_its_errors() {
241         let src = "@group(0) @binding(0) var<storage, read> a: array<f32>;
242                    @group(0) @binding(1) var<storage, read_write> b: array<f32>;
243                    @compute @workgroup_size(64) fn main(@builtin(global_invocation_id) id: vec3<u32>) {
244                        if id.x < arrayLength(&b) { b[id.x] = a[id.x] * 2.0; }
245                    }";
246         let k = parse_kernel(&Kernel::new(src, "main")).unwrap();
247         assert_eq!(k.workgroup_size, [64, 1, 1]);
248         assert!(k.read_only_storage(0));
249         assert!(!k.read_only_storage(1));
250         assert!(parse_kernel(&Kernel::new(src, "nope")).err().unwrap().contains("offers main"));
251         assert!(parse_kernel(&Kernel::new("fn (", "main")).err().unwrap().contains("parse error"));
252     }
253 }