Skip to main content

henad_compute/gpu/
grid_engine.rs

1//! The engine that runs a [`GpuGridModel`] as a [`GpuSimState`], the counterpart of [`crate::cpu::grid_engine`].
2
3use std::marker::PhantomData;
4use std::sync::Arc;
5
6use henad_core::authoring::model::binding::{BindingDecl, buffer_target};
7use henad_core::authoring::model::gpu_grid_model::GpuGridModel;
8use henad_core::model::SimState;
9use henad_core::params::ParamValue;
10use henad_core::view::{StatEntry, stat_entries};
11
12use crate::gpu::GpuContext;
13use crate::gpu::capacity::{Demand, layout_entry, storage_bindings};
14use crate::gpu::contracts::{assert_buffer_labels, assert_workgroup_size};
15use crate::gpu::primitives::pipeline::{compute_pipeline, uniform_buffer};
16use crate::gpu::primitives::readback::{CounterReadback, StatsPoll};
17use crate::gpu::sim_thread::GpuSimState;
18use crate::gpu::view::display::{DisplayTarget, GpuDisplay, build_display_target};
19use crate::snapshot::GpuSnapshot;
20
21/// The uniform every display and reduce shader reads, generated from `henad::dims` through `grid_dims.wgsl`.
22pub(crate) type Dims = crate::shader_bindings::henad::dims::Dims;
23
24/// One ping-ponged pair of storage buffers.
25struct BufferPair {
26    a: wgpu::Buffer,
27    b: wgpu::Buffer,
28}
29
30/// A built [`GpuGridAction`], with the uniform kept so a press can reseed it.
31struct ActionPass {
32    label: String,
33    pipeline: wgpu::ComputePipeline,
34    bind_a: wgpu::BindGroup,
35    bind_b: wgpu::BindGroup,
36    uniform: wgpu::Buffer,
37}
38
39impl ActionPass {
40    fn bind(&self, a_is_current: bool) -> &wgpu::BindGroup {
41        if a_is_current { &self.bind_a } else { &self.bind_b }
42    }
43}
44
45/// GPU-resident state for a [`GpuGridModel`], with every buffer, pipeline and bind group its model declares.
46pub struct GpuGridState<M: GpuGridModel> {
47    width: u32,
48    height: u32,
49    /// Display texture size, capped independently of the grid. See [`crate::display_scale`].
50    tex: (u32, u32),
51    tick: u64,
52
53    device: wgpu::Device,
54    queue: wgpu::Queue,
55
56    step_pipeline: wgpu::ComputePipeline,
57    bind_a2b: wgpu::BindGroup,
58    bind_b2a: wgpu::BindGroup,
59
60    display_pipeline: wgpu::ComputePipeline,
61    display_bind_a: wgpu::BindGroup,
62    display_bind_b: wgpu::BindGroup,
63    display: Arc<GpuDisplay>,
64
65    reduce_pipeline: wgpu::ComputePipeline,
66    reduce_bind_a: wgpu::BindGroup,
67    reduce_bind_b: wgpu::BindGroup,
68    readback: CounterReadback,
69
70    actions: Vec<ActionPass>,
71    /// Seed of the next action press, advanced per press so that pressing twice draws twice.
72    action_seed: u32,
73    /// Parameter values, kept to rewrite the action uniforms on each press. Never edited, since a GPU grid model
74    /// declares every parameter reload-only.
75    params: Vec<ParamValue>,
76
77    /// `true` when the `a` side of every buffer holds the current (latest) state.
78    current_is_a: bool,
79
80    _marker: PhantomData<M>,
81}
82
83impl<M: GpuGridModel> std::fmt::Debug for GpuGridState<M> {
84    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
85        f.debug_struct("GpuGridState")
86            .field("model", &M::ID)
87            .field("tick", &self.tick)
88            .field("width", &self.width)
89            .field("height", &self.height)
90            .finish_non_exhaustive()
91    }
92}
93
94impl<M: GpuGridModel> GpuGridState<M> {
95    /// Builds the state on `ctx`, seeded with the model's default seed.
96    ///
97    /// # Panics
98    ///
99    /// Panics as [`Self::new_seeded`] does.
100    pub fn new(ctx: &GpuContext, params: &[ParamValue]) -> Self {
101        Self::new_seeded(ctx, params, None)
102    }
103
104    /// Resources that would be allocated for this model based on `params`.
105    pub fn demand(params: &[ParamValue], limits: &wgpu::Limits) -> Demand {
106        let (width, height) = M::dims(params);
107        let (width, height) = (width.max(1), height.max(1));
108
109        let mut demand = Demand::default();
110        for (k, len) in M::buffer_lens(width, height).into_iter().enumerate() {
111            demand.push_sides(&format!("{}_buffer{k}", M::ID), len, true);
112        }
113        demand.set_display(width, height, limits);
114        for (label, storage) in Self::declared_passes() {
115            demand.push_pass(label, storage);
116        }
117        demand
118    }
119
120    /// Storage buffers each generated pass binds. Read by both [`Self::demand`] and
121    /// [`Self::max_storage_bindings`], so the device that a host requests and the shortfall that the UI
122    /// reports cannot disagree.
123    fn declared_passes() -> Vec<(String, u32)> {
124        let mut passes = vec![
125            (format!("{}_step", M::ID), storage_bindings(M::STEP_BINDINGS)),
126            (format!("{}_display", M::ID), storage_bindings(M::DISPLAY_BINDINGS)),
127            (format!("{}_reduce", M::ID), storage_bindings(M::REDUCE_BINDINGS)),
128        ];
129        for action in M::ACTIONS {
130            passes.push((
131                format!("{}_action_{}", M::ID, action.desc.id),
132                storage_bindings(action.bindings),
133            ));
134        }
135        passes
136    }
137
138    /// Returns the most storage buffers any declared pass binds. It does not depend on params, so a host can query it
139    /// before it has a device.
140    pub fn max_storage_bindings() -> u32 {
141        Self::declared_passes()
142            .into_iter()
143            .map(|(_, storage)| storage)
144            .max()
145            .unwrap_or(0)
146    }
147
148    /// Builds the state on `ctx`, seeded with `seed`, or with the model's fixed default seed when it is `None`.
149    ///
150    /// # Panics
151    ///
152    /// Panics if the device cannot hold the model. A host calls [`Self::demand`] first, and this assert is the
153    /// backstop. Also panics if a shader declares a workgroup size other than
154    /// [`GpuGridModel::WORKGROUP_SIZE`], or a buffer label is reserved or ends in `_in` or `_out`.
155    #[expect(clippy::too_many_lines)]
156    pub fn new_seeded(ctx: &GpuContext, params: &[ParamValue], seed: Option<u64>) -> Self {
157        let device = &ctx.device;
158        let queue = &ctx.queue;
159
160        let (width, height) = M::dims(params);
161        let (width, height) = (width.max(1), height.max(1));
162
163        let shortfalls = Self::demand(params, &device.limits()).shortfalls(&device.limits());
164        assert!(
165            shortfalls.is_empty(),
166            "{} does not fit this device at {width}x{height}: {}",
167            M::ID,
168            shortfalls.join("; ")
169        );
170        assert_buffer_labels(M::ID, M::BUFFERS.iter().copied());
171        // Every pass dispatches square workgroups of `WORKGROUP_SIZE`.
172        let square = [M::WORKGROUP_SIZE, M::WORKGROUP_SIZE, 1];
173        for (pass, shader) in [
174            ("step", M::STEP_SHADER),
175            ("display", M::DISPLAY_SHADER),
176            ("reduce", M::REDUCE_SHADER),
177        ]
178        .into_iter()
179        .chain(M::ACTIONS.iter().map(|action| (action.desc.id, action.shader)))
180        {
181            assert_workgroup_size(M::ID, pass, shader, square);
182        }
183
184        // Ping-ponged storage buffers, seeded from the model.
185        // Buffer lengths come from the model. A bit-packed model holds many cells per u32, and only the
186        // model knows its buffer lengths.
187        let buffer_lens = M::buffer_lens(width, height);
188        assert_eq!(
189            buffer_lens.len(),
190            M::BUFFERS.len(),
191            "{}: buffer_lens must return BUFFER_COUNT ({}) lengths, got {}",
192            M::ID,
193            M::BUFFERS.len(),
194            buffer_lens.len()
195        );
196        let seeds = M::seed_buffers(width, height, params, seed);
197        assert_eq!(
198            seeds.len(),
199            M::BUFFERS.len(),
200            "{}: seed_buffers must return BUFFER_COUNT ({}) vectors, got {}",
201            M::ID,
202            M::BUFFERS.len(),
203            seeds.len()
204        );
205
206        let buffers: Vec<BufferPair> = seeds
207            .iter()
208            .zip(&buffer_lens)
209            .enumerate()
210            .map(|(k, (seed, &len))| {
211                assert_eq!(
212                    seed.len(),
213                    len,
214                    "{}: seed buffer {k} must match buffer_lens[{k}] ({len}) elements, got {}",
215                    M::ID,
216                    seed.len()
217                );
218                let buffer_size = (len * std::mem::size_of::<u32>()) as u64;
219                let make = |side: char| {
220                    device.create_buffer(&wgpu::BufferDescriptor {
221                        label: Some(&format!("{}_buffer{k}_{side}", M::ID)),
222                        size: buffer_size,
223                        usage: wgpu::BufferUsages::STORAGE
224                            | wgpu::BufferUsages::COPY_SRC
225                            | wgpu::BufferUsages::COPY_DST,
226                        mapped_at_creation: false,
227                    })
228                };
229                let pair = BufferPair {
230                    a: make('a'),
231                    b: make('b'),
232                };
233                // Only the `a` side needs seeding since the `b` side is always written during stepping.
234                queue.write_buffer(&pair.a, 0, bytemuck::cast_slice(seed));
235                pair
236            })
237            .collect();
238
239        // Uniforms.
240        // Display and reduce get their own small buffer rather than depending on the model's
241        // layout starting with the dimensions.
242        let step_params = M::step_params_bytes(width, height, params);
243        let step_params_buffer = device.create_buffer(&wgpu::BufferDescriptor {
244            label: Some(&format!("{}_step_params_buffer", M::ID)),
245            size: step_params.len() as u64,
246            usage: wgpu::BufferUsages::UNIFORM | wgpu::BufferUsages::COPY_DST,
247            mapped_at_creation: false,
248        });
249        queue.write_buffer(&step_params_buffer, 0, &step_params);
250
251        let DisplayTarget {
252            view: display_view,
253            dims: tex,
254            display,
255        } = build_display_target(device, ctx.target_format, width, height);
256
257        let dims_buffer = uniform_buffer(
258            device,
259            queue,
260            &format!("{}_dims_buffer", M::ID),
261            bytemuck::bytes_of(&Dims {
262                grid: [width, height],
263                tex: tex.into(),
264            }),
265        );
266
267        let readback = CounterReadback::new(device, &format!("{}_counters", M::ID), M::STATS.len());
268
269        // Pipelines.
270        // Every layout entry and every bind group entry comes from the name its shader gives the
271        // binding, so a slot index cannot disagree with the shader that owns it.
272        // The action uniforms are built before the pipelines, so the resolver below can return an action uniform by
273        // index. Each uniform is rewritten with a fresh seed on every press. The seed is truncated from the
274        // stream the CPU engines use, as the WGSL generator is 32 bit.
275        let action_seed = henad_core::action::action_seed(seed) as u32;
276        let action_uniforms: Vec<wgpu::Buffer> = M::ACTIONS
277            .iter()
278            .enumerate()
279            .map(|(i, action)| {
280                uniform_buffer(
281                    device,
282                    queue,
283                    &format!("{}_action_{}_params", M::ID, action.desc.id),
284                    &M::action_params_bytes(i, width, height, params, action_seed),
285                )
286            })
287            .collect();
288
289        // An action writes the side that already holds the state. Nothing swaps after an action, so
290        // writing the far side would throw the work away.
291        let resolve = |decl: &BindingDecl, a_is_current: bool, action: Option<usize>| -> wgpu::BindingResource<'_> {
292            if let Some((label, writes)) = buffer_target(decl) {
293                let k = M::BUFFERS
294                    .iter()
295                    .position(|l| *l == label)
296                    .unwrap_or_else(|| panic!("{}: no buffer labelled `{label}`, wanted by `{}`", M::ID, decl.name));
297                let pair = &buffers[k];
298                let (read, write) = if a_is_current {
299                    (&pair.a, &pair.b)
300                } else {
301                    (&pair.b, &pair.a)
302                };
303                return if writes && action.is_none() { write } else { read }.as_entire_binding();
304            }
305            match decl.name {
306                "params" => match action {
307                    Some(i) => action_uniforms[i].as_entire_binding(),
308                    None => step_params_buffer.as_entire_binding(),
309                },
310                "dims" => dims_buffer.as_entire_binding(),
311                "counters" => readback.binding(),
312                "output" => wgpu::BindingResource::TextureView(&display_view),
313                other => panic!("{}: `{other}` is reserved but the engine has no resource for it", M::ID),
314            }
315        };
316
317        let build = |label: &str, shader: &str, decls: &[BindingDecl], action: Option<usize>| {
318            let entries: Vec<wgpu::BindGroupLayoutEntry> = decls
319                .iter()
320                .enumerate()
321                .map(|(i, decl)| layout_entry(i as u32, decl))
322                .collect();
323            let layout = device.create_bind_group_layout(&wgpu::BindGroupLayoutDescriptor {
324                label: Some(&format!("{}_{label}_layout", M::ID)),
325                entries: &entries,
326            });
327            let make = |a_is_current: bool, side: &str| {
328                let entries: Vec<wgpu::BindGroupEntry<'_>> = decls
329                    .iter()
330                    .enumerate()
331                    .map(|(i, decl)| wgpu::BindGroupEntry {
332                        binding: i as u32,
333                        resource: resolve(decl, a_is_current, action),
334                    })
335                    .collect();
336                device.create_bind_group(&wgpu::BindGroupDescriptor {
337                    label: Some(&format!("{}_{label}_bind_{side}", M::ID)),
338                    layout: &layout,
339                    entries: &entries,
340                })
341            };
342            let binds = (make(true, "a"), make(false, "b"));
343            let pipeline = compute_pipeline(device, &format!("{}_{label}", M::ID), shader, &layout);
344            (pipeline, binds)
345        };
346
347        let (step_pipeline, (bind_a2b, bind_b2a)) = build("step", M::STEP_SHADER, M::STEP_BINDINGS, None);
348        let (display_pipeline, (display_bind_a, display_bind_b)) =
349            build("display", M::DISPLAY_SHADER, M::DISPLAY_BINDINGS, None);
350        let (reduce_pipeline, (reduce_bind_a, reduce_bind_b)) =
351            build("reduce", M::REDUCE_SHADER, M::REDUCE_BINDINGS, None);
352
353        let actions: Vec<ActionPass> = M::ACTIONS
354            .iter()
355            .enumerate()
356            .map(|(i, action)| {
357                let label = format!("action_{}", action.desc.id);
358                let (pipeline, (bind_a, bind_b)) = build(&label, action.shader, action.bindings, Some(i));
359                ActionPass {
360                    label: format!("{}_{label}", M::ID),
361                    pipeline,
362                    bind_a,
363                    bind_b,
364                    uniform: action_uniforms[i].clone(),
365                }
366            })
367            .collect();
368
369        Self {
370            width,
371            height,
372            tex,
373            tick: 0,
374            device: device.clone(),
375            queue: queue.clone(),
376            step_pipeline,
377            bind_a2b,
378            bind_b2a,
379            display_pipeline,
380            display_bind_a,
381            display_bind_b,
382            display,
383            reduce_pipeline,
384            reduce_bind_a,
385            reduce_bind_b,
386            readback,
387            actions,
388            action_seed,
389            params: params.to_vec(),
390            current_is_a: true,
391            _marker: PhantomData,
392        }
393    }
394
395    /// Returns the workgroups covering the step pass's domain, which a packed model measures in words.
396    fn step_workgroups(&self) -> (u32, u32) {
397        let (x, y) = M::step_dims(self.width, self.height);
398        (x.div_ceil(M::WORKGROUP_SIZE), y.div_ceil(M::WORKGROUP_SIZE))
399    }
400
401    /// Returns the workgroups of the reduce pass, at one invocation per cell.
402    fn cell_workgroups(&self) -> (u32, u32) {
403        (
404            self.width.div_ceil(M::WORKGROUP_SIZE),
405            self.height.div_ceil(M::WORKGROUP_SIZE),
406        )
407    }
408
409    /// Returns the workgroups of the display pass, at one invocation per texel. Equals [`Self::cell_workgroups`]
410    /// until the grid outgrows the texture cap.
411    fn texel_workgroups(&self) -> (u32, u32) {
412        (
413            self.tex.0.div_ceil(M::WORKGROUP_SIZE),
414            self.tex.1.div_ceil(M::WORKGROUP_SIZE),
415        )
416    }
417
418    fn current_display_bind_group(&self) -> &wgpu::BindGroup {
419        if self.current_is_a {
420            &self.display_bind_a
421        } else {
422            &self.display_bind_b
423        }
424    }
425
426    fn current_reduce_bind_group(&self) -> &wgpu::BindGroup {
427        if self.current_is_a {
428            &self.reduce_bind_a
429        } else {
430            &self.reduce_bind_b
431        }
432    }
433}
434
435impl<M: GpuGridModel> SimState for GpuGridState<M> {
436    /// Steps once in its own submission, for a caller that holds only a `SimState`. The GPU runner calls
437    /// `encode_steps` directly and batches the steps.
438    fn step(&mut self) {
439        let mut encoder = self.device.create_command_encoder(&wgpu::CommandEncoderDescriptor {
440            label: Some("gpu_grid_single_step"),
441        });
442        self.encode_steps(&mut encoder, 1, None);
443        self.queue.submit(Some(encoder.finish()));
444    }
445
446    fn tick(&self) -> u64 {
447        self.tick
448    }
449
450    fn stats(&self) -> Vec<StatEntry> {
451        stat_entries(M::STATS, M::stats(self.readback.values()))
452    }
453
454    /// Encodes the action into its own submission.
455    ///
456    /// The GPU runner calls [`GpuSimState::encode_action`] instead, and submits the action before
457    /// its snapshot.
458    fn act(&mut self, index: usize) -> bool {
459        let mut encoder = self.device.create_command_encoder(&wgpu::CommandEncoderDescriptor {
460            label: Some("henad_gpu_grid_action"),
461        });
462        if !GpuSimState::encode_action(self, &mut encoder, index) {
463            return false;
464        }
465        self.queue.submit(Some(encoder.finish()));
466        true
467    }
468
469    /// Rejects every live edit. Its entry declares every parameter reload-only.
470    fn set_param(&mut self, _index: usize, _value: &ParamValue) -> bool {
471        false
472    }
473
474    fn population(&self) -> u64 {
475        u64::from(self.width) * u64::from(self.height)
476    }
477
478    fn heap_bytes(&self) -> usize {
479        // Two ping-ponged sides per buffer, plus the capped RGBA display texture.
480        let buffers: usize = M::buffer_lens(self.width, self.height)
481            .iter()
482            .map(|len| len * std::mem::size_of::<u32>() * 2)
483            .sum();
484        let display_texture = (self.tex.0 as usize) * (self.tex.1 as usize) * 4;
485        buffers + display_texture
486    }
487}
488
489impl<M: GpuGridModel> GpuSimState for GpuGridState<M> {
490    /// Records `count` step dispatches into `encoder`, one compute pass per step.
491    ///
492    /// Each step is a read-after-write hazard on the ping-ponged state buffers, and wgpu inserts
493    /// barriers only between passes. Dispatches looped inside one pass would read stale data.
494    /// Batching happens at the submission level instead, with one encoder for all `count` passes.
495    ///
496    /// If `timestamps` is `Some`, the first pass's beginning and the last pass's end are stamped
497    /// into query indices 0 and 1 so the caller can measure GPU time for the whole batch.
498    fn encode_steps(&mut self, encoder: &mut wgpu::CommandEncoder, count: u32, timestamps: Option<&wgpu::QuerySet>) {
499        if count == 0 {
500            return;
501        }
502        let (wg_x, wg_y) = self.step_workgroups();
503        for i in 0..count {
504            let bind_group = if self.current_is_a {
505                &self.bind_a2b
506            } else {
507                &self.bind_b2a
508            };
509            let is_first = i == 0;
510            let is_last = i == count - 1;
511            // A `ComputePassTimestampWrites` requires at least one of the two indices to be
512            // `Some`, so only the first and last passes of the batch get timestamp writes.
513            let timestamp_writes =
514                timestamps
515                    .filter(|_| is_first || is_last)
516                    .map(|query_set| wgpu::ComputePassTimestampWrites {
517                        query_set,
518                        beginning_of_pass_write_index: is_first.then_some(0),
519                        end_of_pass_write_index: is_last.then_some(1),
520                    });
521            let mut pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
522                label: Some("gpu_grid_step_pass"),
523                timestamp_writes,
524            });
525            pass.set_pipeline(&self.step_pipeline);
526            pass.set_bind_group(0, bind_group, &[]);
527            pass.dispatch_workgroups(wg_x, wg_y, 1);
528            drop(pass);
529            self.current_is_a = !self.current_is_a;
530        }
531        self.tick += u64::from(count);
532    }
533
534    fn encode_action(&mut self, encoder: &mut wgpu::CommandEncoder, index: usize) -> bool {
535        let Some(action) = self.actions.get(index) else {
536            return false;
537        };
538        // A fresh seed per press, so pressing twice draws twice. The write is queued before the
539        // encoder is submitted, so the pass reads the new value.
540        self.action_seed = self.action_seed.wrapping_mul(747_796_405).wrapping_add(2_891_336_453);
541        self.queue.write_buffer(
542            &action.uniform,
543            0,
544            &M::action_params_bytes(index, self.width, self.height, &self.params, self.action_seed),
545        );
546
547        let (groups_x, groups_y) = self.step_workgroups();
548        let mut pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
549            label: Some(&action.label),
550            timestamp_writes: None,
551        });
552        pass.set_pipeline(&action.pipeline);
553        pass.set_bind_group(0, action.bind(self.current_is_a), &[]);
554        pass.dispatch_workgroups(groups_x, groups_y, 1);
555        true
556    }
557
558    fn encode_snapshot_passes(&mut self, encoder: &mut wgpu::CommandEncoder) {
559        {
560            let (wg_x, wg_y) = self.texel_workgroups();
561            let mut pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
562                label: Some("gpu_grid_display_pass"),
563                timestamp_writes: None,
564            });
565            pass.set_pipeline(&self.display_pipeline);
566            pass.set_bind_group(0, self.current_display_bind_group(), &[]);
567            pass.dispatch_workgroups(wg_x, wg_y, 1);
568        }
569        self.encode_stats_passes(encoder);
570    }
571
572    fn encode_stats_passes(&mut self, encoder: &mut wgpu::CommandEncoder) {
573        let (wg_x, wg_y) = self.cell_workgroups();
574
575        // Clear -> accumulate -> copy out. wgpu inserts the barriers between these because they
576        // are separate passes/copies within the one encoder.
577        self.readback.encode_clear(encoder);
578        {
579            let mut pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
580                label: Some("gpu_grid_reduce_pass"),
581                timestamp_writes: None,
582            });
583            pass.set_pipeline(&self.reduce_pipeline);
584            pass.set_bind_group(0, self.current_reduce_bind_group(), &[]);
585            pass.dispatch_workgroups(wg_x, wg_y, 1);
586        }
587        self.readback.encode_copy(encoder);
588    }
589
590    fn begin_stats_readback(&mut self) {
591        self.readback.begin_map();
592    }
593
594    fn poll_stats_readback(&mut self, device: &wgpu::Device, block: bool) -> StatsPoll {
595        if block {
596            self.readback.poll_blocking(device)
597        } else {
598            self.readback.poll(device)
599        }
600    }
601
602    fn stats_readback_pending(&self) -> bool {
603        self.readback.is_pending()
604    }
605
606    /// Returns the display layer alone. A `GpuGridModel` has no agent layer.
607    fn view(&self) -> GpuSnapshot {
608        GpuSnapshot {
609            display: Some(Arc::clone(&self.display)),
610            agents: None,
611        }
612    }
613}