1use 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
23pub 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 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, ¶ms_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}