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 }