Skip to main content

wde_renderer/
compute.rs

1//! On-demand compute dispatch for callers outside the render world's own per-frame systems (e.g.
2//! background tasks). Unlike [`crate::core::PipelineManager`], which builds pipelines
3//! asynchronously across frames and is only usable from render-app ECS systems, this builds
4//! synchronously and caches by label, so it's safe to call from any thread and cheap to call
5//! repeatedly.
6
7use std::collections::HashMap;
8use std::sync::{Arc, Mutex};
9
10use wde_wgpu::bind_group::{BindGroupBuilder, BindGroupLayout, WgpuBindGroupLayout};
11use wde_wgpu::buffer::{Buffer, BufferBindingType, BufferUsage};
12use wde_wgpu::command_buffer::CommandBuffer;
13use wde_wgpu::compute_pipeline::ComputePipeline;
14use wde_wgpu::render_pipeline::ShaderStages;
15
16use crate::core::RenderInstance;
17
18struct CachedCompute {
19    pipeline: ComputePipeline,
20    layout: WgpuBindGroupLayout
21}
22
23/// Caches a compiled pipeline + bind group layout per label, so repeat dispatches skip WGSL
24/// parsing/validation. Blocks the calling thread until the result is read back.
25pub struct ComputeDispatcher {
26    instance: RenderInstance,
27    cache: Mutex<HashMap<&'static str, Arc<CachedCompute>>>
28}
29impl ComputeDispatcher {
30    pub fn new(instance: &RenderInstance) -> Self {
31        Self {
32            instance: RenderInstance(instance.0.clone()),
33            cache: Mutex::new(HashMap::new())
34        }
35    }
36
37    /// Dispatches `shader`'s `main` over `output_len` `f32`s (binding 1), with `params` as a
38    /// uniform buffer (binding 0) and `inputs` as read-only `f32` storage buffers at bindings
39    /// `2, 3, ...`. `inputs.len()` must stay the same across calls sharing `label`, since the
40    /// bind group layout is built once and cached under it.
41    pub fn dispatch_f32<P: bytemuck::NoUninit>(
42        &self,
43        label: &'static str,
44        shader: &str,
45        params: &P,
46        inputs: &[&[f32]],
47        output_len: usize,
48        workgroups: (u32, u32, u32)
49    ) -> Result<Vec<f32>, String> {
50        let data = self
51            .instance
52            .0
53            .read()
54            .map_err(|_| "render instance lock poisoned".to_string())?;
55
56        let cached = {
57            let mut cache = self.cache.lock().unwrap();
58            if let Some(cached) = cache.get(label) {
59                cached.clone()
60            } else {
61                let layout = BindGroupLayout::new(label, |b| {
62                    b.add_buffer(0, ShaderStages::COMPUTE, BufferBindingType::Uniform);
63                    b.add_buffer(
64                        1,
65                        ShaderStages::COMPUTE,
66                        BufferBindingType::Storage { read_only: false }
67                    );
68                    for i in 0..inputs.len() {
69                        b.add_buffer(
70                            2 + i as u32,
71                            ShaderStages::COMPUTE,
72                            BufferBindingType::Storage { read_only: true }
73                        );
74                    }
75                });
76                let wgpu_layout = layout.build(&data).map_err(|e| format!("{label}: {e:?}"))?;
77                let mut pipeline = ComputePipeline::new(label);
78                pipeline
79                    .set_shader(shader)
80                    .set_bind_groups(vec![wgpu_layout.clone()]);
81                pipeline.init(&data).map_err(|e| format!("{label}: {e:?}"))?;
82                let cached = Arc::new(CachedCompute {
83                    pipeline,
84                    layout: wgpu_layout
85                });
86                cache.insert(label, cached.clone());
87                cached
88            }
89        };
90
91        let params_buf = Buffer::new(
92            &data,
93            &format!("{label}-params"),
94            std::mem::size_of::<P>(),
95            BufferUsage::UNIFORM | BufferUsage::COPY_DST,
96            Some(bytemuck::bytes_of(params))
97        );
98        let output_size = output_len * std::mem::size_of::<f32>();
99        let output_buf = Buffer::new(
100            &data,
101            &format!("{label}-output"),
102            output_size,
103            BufferUsage::STORAGE | BufferUsage::COPY_SRC,
104            None
105        );
106        let input_bufs: Vec<Buffer> = inputs
107            .iter()
108            .enumerate()
109            .map(|(i, slice)| {
110                Buffer::new(
111                    &data,
112                    &format!("{label}-input{i}"),
113                    std::mem::size_of_val(*slice),
114                    BufferUsage::STORAGE | BufferUsage::COPY_DST,
115                    Some(bytemuck::cast_slice(slice))
116                )
117            })
118            .collect();
119
120        let mut entries = vec![
121            BindGroupBuilder::buffer(0, &params_buf),
122            BindGroupBuilder::buffer(1, &output_buf),
123        ];
124        for (i, buf) in input_bufs.iter().enumerate() {
125            entries.push(BindGroupBuilder::buffer(2 + i as u32, buf));
126        }
127        let bind_group = BindGroupBuilder::build(label, &data, &cached.layout, &entries)
128            .map_err(|e| format!("{label}: {e:?}"))?;
129
130        let mut cmd = CommandBuffer::new(&data, label);
131        {
132            let mut pass = cmd.create_compute_pass(label);
133            pass.set_pipeline(&cached.pipeline)
134                .map_err(|e| format!("{label}: {e:?}"))?
135                .set_bind_group(0, &bind_group);
136            pass.dispatch(workgroups.0, workgroups.1, workgroups.2)
137                .map_err(|e| format!("{label}: {e:?}"))?;
138        }
139        cmd.submit(&data);
140
141        let staging = Buffer::new(
142            &data,
143            &format!("{label}-staging"),
144            output_size,
145            BufferUsage::MAP_READ | BufferUsage::COPY_DST,
146            None
147        );
148        staging.copy_from_buffer(&data, &output_buf);
149
150        let mut result = vec![0.0f32; output_len];
151        staging.map_read(&data, |view| {
152            result.copy_from_slice(bytemuck::cast_slice(&view));
153        });
154        Ok(result)
155    }
156}