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

src/web/compute.rs (14.4K)

  1 //! `crate::compute` jobs on WebGPU: the browser's `vk::ComputeDevice`.
  2 //!
  3 //! The same kernels, bindings and rules (`crate::compute`), answered the same
  4 //! way — a ping-pong job's result lands in its output binding, every
  5 //! read-write binding is read back into its slice — with one difference
  6 //! the platform makes: reading a buffer back is a promise, so [`run`] and its
  7 //! siblings are `async`. A job is still one submission, its passes in one
  8 //! compute pass; WebGPU orders the dispatches of a pass by itself.
  9 //!
 10 //! A kernel is parsed and validated by naga before WebGPU sees it, so a
 11 //! user's bad WGSL is an `Err` with naga's diagnostic, as natively; what
 12 //! WebGPU still rejects (a limit, a layout) is caught in a validation error
 13 //! scope and returned too. The device asks for the adapter's own limits on
 14 //! storage buffers and workgroups: WebGPU's defaults allow eight storage
 15 //! buffers a stage, where a job may carry sixteen. What the adapter offers
 16 //! is the ceiling — SwiftShader's is ten — and a job past it is an `Err`
 17 //! naming the limit. (cce-designer's kernels bind at most seven.)
 18 //!
 19 //! [`run`]: ComputeDevice::run
 20 
 21 use std::collections::HashMap;
 22 
 23 use wasm_bindgen::{JsCast, JsValue};
 24 use web_sys::{
 25     gpu_buffer_usage as buffer_usage, gpu_map_mode as map_mode, gpu_shader_stage as shader_stage, GpuBindGroupDescriptor,
 26     GpuBindGroupEntry, GpuBindGroupLayout, GpuBindGroupLayoutDescriptor, GpuBindGroupLayoutEntry, GpuBuffer,
 27     GpuBufferBinding, GpuBufferBindingLayout, GpuBufferBindingType, GpuBufferDescriptor, GpuComputePassDescriptor,
 28     GpuComputePipeline, GpuComputePipelineDescriptor, GpuDevice, GpuErrorFilter, GpuPipelineLayoutDescriptor,
 29     GpuProgrammableStage, GpuQueue, GpuShaderModuleDescriptor,
 30 };
 31 
 32 use crate::compute::{
 33     check_job, padded_len, parse_kernel, result_slot, slot_for, workgroups, BindKind, Binding, Kernel,
 34 };
 35 
 36 /// The limits a compute device asks the adapter for in full.
 37 const LIMITS: &[&str] = &[
 38     "maxStorageBuffersPerShaderStage",
 39     "maxUniformBuffersPerShaderStage",
 40     "maxStorageBufferBindingSize",
 41     "maxUniformBufferBindingSize",
 42     "maxBufferSize",
 43     "maxComputeWorkgroupStorageSize",
 44     "maxComputeInvocationsPerWorkgroup",
 45     "maxComputeWorkgroupSizeX",
 46     "maxComputeWorkgroupSizeY",
 47     "maxComputeWorkgroupSizeZ",
 48     "maxComputeWorkgroupsPerDimension",
 49 ];
 50 
 51 #[derive(Clone, PartialEq, Eq, Hash)]
 52 struct PipelineKey {
 53     kernel: Kernel,
 54     kinds: Vec<BindKind>,
 55 }
 56 
 57 struct Pipeline {
 58     pipeline: GpuComputePipeline,
 59     layout: GpuBindGroupLayout,
 60     workgroup_size: [u32; 3],
 61 }
 62 
 63 /// A buffer and its size in bytes.
 64 struct Slot {
 65     buffer: GpuBuffer,
 66     size: usize,
 67 }
 68 
 69 /// A WebGPU device that runs compute jobs. See the module docs.
 70 pub struct ComputeDevice {
 71     _gpu: web_sys::Gpu,
 72     adapter: web_sys::GpuAdapter,
 73     device: GpuDevice,
 74     queue: GpuQueue,
 75     pipelines: HashMap<PipelineKey, Pipeline>,
 76     /// One buffer per binding index, grown when a job needs more room.
 77     slots: Vec<Option<Slot>>,
 78     /// The mappable copies read-write bindings are read back through.
 79     readback: Vec<Option<Slot>>,
 80 }
 81 
 82 fn js_err(e: JsValue) -> String {
 83     e.as_string()
 84         .or_else(|| js_sys::Reflect::get(&e, &"message".into()).ok().and_then(|m| m.as_string()))
 85         .unwrap_or_else(|| format!("{e:?}"))
 86 }
 87 
 88 impl ComputeDevice {
 89     /// A device on the browser's WebGPU adapter. `Err` when the browser
 90     /// offers none.
 91     pub async fn new() -> Result<Self, String> {
 92         let (gpu, adapter, device) =
 93             super::request_device(LIMITS).await.map_err(|e| format!("no WebGPU compute device: {}", js_err(e)))?;
 94         let queue = device.queue();
 95         Ok(Self { _gpu: gpu, adapter, device, queue, pipelines: HashMap::new(), slots: Vec::new(), readback: Vec::new() })
 96     }
 97 
 98     /// The adapter as the browser describes it, for a log line or a status
 99     /// readout. A browser may say little: it guards what fingerprints a machine.
100     pub fn device_name(&self) -> String {
101         let info = self.adapter.info();
102         let parts: Vec<String> =
103             [info.vendor(), info.architecture(), info.device(), info.description()].into_iter().filter(|s| !s.is_empty()).collect();
104         if parts.is_empty() { "WebGPU".into() } else { format!("WebGPU {}", parts.join(" ")) }
105     }
106 
107     /// The entry point's `@workgroup_size`, compiling the kernel if needed.
108     pub async fn workgroup_size(&mut self, kernel: &Kernel, kinds: &[BindKind]) -> Result<[u32; 3], String> {
109         let key = PipelineKey { kernel: kernel.clone(), kinds: kinds.to_vec() };
110         Ok(self.pipeline(&key).await?.workgroup_size)
111     }
112 
113     /// Run the kernel over `items` invocations along x, the workgroup count
114     /// from the entry point's own `@workgroup_size` (see `vk::ComputeDevice::run_over`).
115     pub async fn run_over(&mut self, kernel: &Kernel, bindings: &mut [Binding<'_>], items: u32) -> Result<(), String> {
116         let kinds: Vec<BindKind> = bindings.iter().map(Binding::kind).collect();
117         let wg = self.workgroup_size(kernel, &kinds).await?;
118         self.execute(kernel, bindings, [workgroups(items, wg[0]), 1, 1], 1, None).await
119     }
120 
121     /// Upload every binding, dispatch `groups` workgroups of the kernel, and
122     /// read every [`Binding::Storage`] back into its slice once the GPU is done.
123     pub async fn run(&mut self, kernel: &Kernel, bindings: &mut [Binding<'_>], groups: [u32; 3]) -> Result<(), String> {
124         self.execute(kernel, bindings, groups, 1, None).await
125     }
126 
127     /// [`run_passes`](Self::run_passes) sized over `items`, like [`run_over`](Self::run_over).
128     pub async fn run_passes_over(
129         &mut self,
130         kernel: &Kernel,
131         bindings: &mut [Binding<'_>],
132         items: u32,
133         passes: u32,
134         ping_pong: Option<(usize, usize)>,
135     ) -> Result<(), String> {
136         let kinds: Vec<BindKind> = bindings.iter().map(Binding::kind).collect();
137         let wg = self.workgroup_size(kernel, &kinds).await?;
138         self.execute(kernel, bindings, [workgroups(items, wg[0]), 1, 1], passes, ping_pong).await
139     }
140 
141     /// `passes` dispatches of the kernel in one submission, with an optional
142     /// ping-pong pair — the semantics of `vk::ComputeDevice::run_passes`.
143     pub async fn run_passes(
144         &mut self,
145         kernel: &Kernel,
146         bindings: &mut [Binding<'_>],
147         groups: [u32; 3],
148         passes: u32,
149         ping_pong: Option<(usize, usize)>,
150     ) -> Result<(), String> {
151         self.execute(kernel, bindings, groups, passes, ping_pong).await
152     }
153 
154     async fn execute(
155         &mut self,
156         kernel: &Kernel,
157         bindings: &mut [Binding<'_>],
158         groups: [u32; 3],
159         passes: u32,
160         ping_pong: Option<(usize, usize)>,
161     ) -> Result<(), String> {
162         check_job(bindings, groups, passes, ping_pong)?;
163         let kinds: Vec<BindKind> = bindings.iter().map(Binding::kind).collect();
164         let key = PipelineKey { kernel: kernel.clone(), kinds };
165         let (pipeline, layout) = {
166             let p = self.pipeline(&key).await?;
167             (p.pipeline.clone(), p.layout.clone())
168         };
169 
170         // Buffers: one per binding index, reused when big enough. The
171         // padding is uploaded too, as zeros, so an `arrayLength` that counts
172         // it reads what it would natively.
173         let uniform_cap = self.device.limits().max_uniform_buffer_binding_size() as usize;
174         let mut sizes = Vec::with_capacity(bindings.len());
175         for (i, b) in bindings.iter().enumerate() {
176             let bytes = b.bytes();
177             let padded = padded_len(bytes.len());
178             if b.kind() == BindKind::Uniform && padded > uniform_cap {
179                 return Err(format!("binding {i}: a uniform block of {} bytes exceeds the device's {uniform_cap}", bytes.len()));
180             }
181             let usage = buffer_usage::STORAGE | buffer_usage::UNIFORM | buffer_usage::COPY_SRC | buffer_usage::COPY_DST;
182             ensure(&self.device, &mut self.slots, i, padded, usage, "compute-binding")?;
183             let mut upload = bytes.to_vec();
184             upload.resize(padded, 0);
185             let slot = self.slots[i].as_ref().unwrap();
186             self.queue.write_buffer_with_u32_and_u8_slice(&slot.buffer, 0, &upload).map_err(js_err)?;
187             sizes.push(padded);
188         }
189 
190         self.device.push_error_scope(GpuErrorFilter::Validation);
191         // One bind group, or two with the ping-pong pair swapped in the
192         // second, so alternate passes bind the buffers the other way round.
193         let group_count = if ping_pong.is_some() { 2 } else { 1 };
194         let mut groups_bound = Vec::with_capacity(group_count);
195         for g in 0..group_count {
196             let entries: Vec<GpuBindGroupEntry> = (0..bindings.len())
197                 .map(|i| {
198                     let slot = slot_for(i, g == 1, ping_pong);
199                     let binding = GpuBufferBinding::new(&self.slots[slot].as_ref().unwrap().buffer);
200                     binding.set_size(sizes[slot] as u32);
201                     GpuBindGroupEntry::new_with_gpu_buffer_binding(i as u32, &binding)
202                 })
203                 .collect();
204             groups_bound.push(self.device.create_bind_group(&GpuBindGroupDescriptor::new(&entries, &layout)));
205         }
206         let encoder = self.device.create_command_encoder();
207         let pass = encoder.begin_compute_pass_with_descriptor(&GpuComputePassDescriptor::new());
208         pass.set_pipeline(&pipeline);
209         for p in 0..passes {
210             let g = if ping_pong.is_some() && p % 2 == 1 { 1 } else { 0 };
211             pass.set_bind_group(0, Some(&groups_bound[g]));
212             pass.dispatch_workgroups_with_workgroup_count_y_and_workgroup_count_z(groups[0], groups[1], groups[2]);
213         }
214         pass.end();
215 
216         // Copy every read-write binding's result into a mappable buffer.
217         let mut reads = Vec::new();
218         for (i, b) in bindings.iter().enumerate() {
219             if let Binding::Storage(out) = b {
220                 let from = result_slot(i, passes, ping_pong);
221                 let len = padded_len(out.len());
222                 ensure(&self.device, &mut self.readback, i, len, buffer_usage::MAP_READ | buffer_usage::COPY_DST, "compute-readback")?;
223                 encoder
224                     .copy_buffer_to_buffer_with_u32_and_u32_and_u32(
225                         &self.slots[from].as_ref().unwrap().buffer,
226                         0,
227                         &self.readback[i].as_ref().unwrap().buffer,
228                         0,
229                         len as u32,
230                     )
231                     .map_err(js_err)?;
232                 reads.push(i);
233             }
234         }
235         self.queue.submit(&[encoder.finish()]);
236         if let Some(err) = wasm_bindgen_futures::JsFuture::from(self.device.pop_error_scope()).await.map_err(js_err)?.dyn_ref::<web_sys::GpuError>()
237         {
238             return Err(format!("WebGPU rejected the job: {}", err.message()));
239         }
240 
241         for i in reads {
242             let Binding::Storage(out) = &mut bindings[i] else { unreachable!() };
243             let buffer = &self.readback[i].as_ref().unwrap().buffer;
244             wasm_bindgen_futures::JsFuture::from(buffer.map_async_with_u32_and_u32(map_mode::READ, 0, padded_len(out.len()) as u32))
245                 .await
246                 .map_err(|e| format!("binding {i}: readback: {}", js_err(e)))?;
247             let mapped = js_sys::Uint8Array::new(&buffer.get_mapped_range().map_err(js_err)?.into());
248             mapped.subarray(0, out.len() as u32).copy_to(out);
249             buffer.unmap();
250         }
251         Ok(())
252     }
253 
254     async fn pipeline(&mut self, key: &PipelineKey) -> Result<&Pipeline, String> {
255         if !self.pipelines.contains_key(key) {
256             let built = self.build_pipeline(key).await?;
257             self.pipelines.insert(key.clone(), built);
258         }
259         Ok(&self.pipelines[key])
260     }
261 
262     async fn build_pipeline(&self, key: &PipelineKey) -> Result<Pipeline, String> {
263         let parsed = parse_kernel(&key.kernel)?;
264         // Read-only storage is its own binding type here (WebGPU tells it
265         // from read-write, which Vulkan does not): the module says which.
266         let entries: Vec<GpuBindGroupLayoutEntry> = key
267             .kinds
268             .iter()
269             .enumerate()
270             .map(|(i, k)| {
271                 let ty = match k {
272                     BindKind::Uniform => GpuBufferBindingType::Uniform,
273                     BindKind::Storage if parsed.read_only_storage(i as u32) => GpuBufferBindingType::ReadOnlyStorage,
274                     BindKind::Storage => GpuBufferBindingType::Storage,
275                 };
276                 let buffer = GpuBufferBindingLayout::new();
277                 buffer.set_type(ty);
278                 let entry = GpuBindGroupLayoutEntry::new(i as u32, shader_stage::COMPUTE);
279                 entry.set_buffer(&buffer);
280                 entry
281             })
282             .collect();
283         self.device.push_error_scope(GpuErrorFilter::Validation);
284         let layout = self.device.create_bind_group_layout(&GpuBindGroupLayoutDescriptor::new(&entries)).map_err(js_err)?;
285         let pipeline_layout = self.device.create_pipeline_layout(&GpuPipelineLayoutDescriptor::new(&[js_sys::JsOption::wrap(layout.clone())]));
286         let module = self.device.create_shader_module(&GpuShaderModuleDescriptor::new(&key.kernel.source));
287         let stage = GpuProgrammableStage::new(&module);
288         stage.set_entry_point(&key.kernel.entry);
289         let pipeline = self.device.create_compute_pipeline(&GpuComputePipelineDescriptor::new(&pipeline_layout, &stage));
290         if let Some(err) = wasm_bindgen_futures::JsFuture::from(self.device.pop_error_scope()).await.map_err(js_err)?.dyn_ref::<web_sys::GpuError>()
291         {
292             return Err(format!("compute pipeline: {}", err.message()));
293         }
294         Ok(Pipeline { pipeline, layout, workgroup_size: parsed.workgroup_size })
295     }
296 }
297 
298 /// `slots[i]` holds a buffer of at least `size` bytes with `usage`.
299 fn ensure(device: &GpuDevice, slots: &mut Vec<Option<Slot>>, i: usize, size: usize, usage: u32, label: &str) -> Result<(), String> {
300     if slots.len() <= i {
301         slots.resize_with(i + 1, || None);
302     }
303     if slots[i].as_ref().is_some_and(|s| s.size >= size) {
304         return Ok(());
305     }
306     if let Some(old) = slots[i].take() {
307         old.buffer.destroy();
308     }
309     let desc = GpuBufferDescriptor::new(size as u32, usage);
310     desc.set_label(label);
311     slots[i] = Some(Slot { buffer: device.create_buffer(&desc).map_err(js_err)?, size });
312     Ok(())
313 }
314 
315 impl Drop for ComputeDevice {
316     fn drop(&mut self) {
317         for s in self.slots.iter().chain(self.readback.iter()).flatten() {
318             s.buffer.destroy();
319         }
320     }
321 }