Skip to main content

henad_compute/gpu/
agent_engine.rs

1//! The engine that runs a [`GpuAgentModel`] as a [`GpuSimState`], the counterpart of [`crate::cpu::agent_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::field::Extent;
8use henad_core::authoring::model::gpu_agent_model::{Geometry, GpuAgentModel, PassCtx, PassId};
9use henad_core::model::SimState;
10use henad_core::params::ParamValue;
11use henad_core::view::{StatEntry, stat_entries};
12
13use crate::display_scale::display_dims;
14use crate::gpu::capacity::{Demand, layout_entry, storage_bindings};
15use crate::gpu::contracts::{assert_buffer_labels, assert_workgroup_size};
16use crate::gpu::primitives::dispatch::{WORKGROUP, linear_dispatch};
17use crate::gpu::primitives::pipeline::{compute_pipeline, lane_buffer, storage_buffer, uniform_buffer};
18use crate::gpu::primitives::readback::{CounterReadback, StatsPoll};
19use crate::gpu::primitives::reduce::GpuLaneReduce;
20use crate::gpu::primitives::spatial_hash::{GpuSpatialHash, HashGrid};
21use crate::gpu::sim_thread::GpuSimState;
22use crate::gpu::view::agents::GpuAgents;
23use crate::gpu::view::display::{DisplayTarget, GpuDisplay, build_display_target};
24use crate::gpu::{GpuContext, MAX_STEPS_PER_SUBMISSION};
25use crate::snapshot::GpuSnapshot;
26
27/// A resource that exists once, or once per ping-ponged side.
28struct Sides<T> {
29    a: T,
30    b: Option<T>,
31}
32
33impl<T> Sides<T> {
34    fn pick(&self, a_is_current: bool) -> &T {
35        if a_is_current {
36            &self.a
37        } else {
38            self.b.as_ref().unwrap_or(&self.a)
39        }
40    }
41}
42
43struct BufferSides {
44    a: wgpu::Buffer,
45    b: Option<wgpu::Buffer>,
46}
47
48impl BufferSides {
49    /// Returns `(current, next)`, the same buffer twice when this one is written in place.
50    fn sides(&self, a_is_current: bool) -> (&wgpu::Buffer, &wgpu::Buffer) {
51        match &self.b {
52            None => (&self.a, &self.a),
53            Some(b) if a_is_current => (&self.a, b),
54            Some(b) => (b, &self.a),
55        }
56    }
57}
58
59struct EncodedPass {
60    label: String,
61    pipeline: wgpu::ComputePipeline,
62    binds: Sides<wgpu::BindGroup>,
63    groups: (u32, u32),
64    /// Uniform of the pass, kept so an action can reseed it before a press. Unused for a pass that is not an action.
65    uniform: wgpu::Buffer,
66}
67
68impl EncodedPass {
69    fn encode(
70        &self,
71        encoder: &mut wgpu::CommandEncoder,
72        a_is_current: bool,
73        timestamps: Option<wgpu::ComputePassTimestampWrites<'_>>,
74    ) {
75        let mut pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
76            label: Some(&self.label),
77            timestamp_writes: timestamps,
78        });
79        pass.set_pipeline(&self.pipeline);
80        pass.set_bind_group(0, self.binds.pick(a_is_current), &[]);
81        pass.dispatch_workgroups(self.groups.0, self.groups.1, 1);
82    }
83}
84
85/// Returns the timestamp writes of a pass, `None` for a pass that neither opens nor closes a batch.
86///
87/// A `ComputePassTimestampWrites` needs at least one index set.
88fn stamps(
89    query_set: Option<&wgpu::QuerySet>,
90    opening: bool,
91    closing: bool,
92) -> Option<wgpu::ComputePassTimestampWrites<'_>> {
93    query_set
94        .filter(|_| opening || closing)
95        .map(|query_set| wgpu::ComputePassTimestampWrites {
96            query_set,
97            beginning_of_pass_write_index: opening.then_some(0),
98            end_of_pass_write_index: closing.then_some(1),
99        })
100}
101
102/// GPU-resident state for a [`GpuAgentModel`], with every buffer, pass and bind group its model declares.
103pub struct GpuAgentState<M: GpuAgentModel> {
104    geom: Geometry,
105    tick: u64,
106
107    device: wgpu::Device,
108    queue: wgpu::Queue,
109
110    buffers: Vec<BufferSides>,
111    /// `true` when the `a` side of every double buffered buffer holds the current state.
112    current_is_a: bool,
113    /// Set when some buffer requests double buffering. Nothing flips otherwise.
114    ping_pong: bool,
115
116    /// Spatial hash, if the model declares one.
117    index: Option<GpuSpatialHash>,
118    index_binds: Option<Sides<wgpu::BindGroup>>,
119
120    steps: Vec<EncodedPass>,
121    display: Option<(EncodedPass, Arc<GpuDisplay>)>,
122
123    reduce: GpuLaneReduce,
124    reduce_pass: EncodedPass,
125    counters: Option<CounterReadback>,
126
127    actions: Vec<EncodedPass>,
128    /// Seed of the next action press, advanced per press so that pressing twice draws twice.
129    action_seed: u32,
130    /// Parameter values, kept to rewrite the action uniforms on each press.
131    params: Vec<ParamValue>,
132
133    agents: Sides<Arc<GpuAgents>>,
134
135    _marker: PhantomData<M>,
136}
137
138impl<M: GpuAgentModel> std::fmt::Debug for GpuAgentState<M> {
139    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
140        f.debug_struct("GpuAgentState")
141            .field("model", &M::ID)
142            .field("tick", &self.tick)
143            .finish_non_exhaustive()
144    }
145}
146
147impl<M: GpuAgentModel> GpuAgentState<M> {
148    /// Builds the state on `ctx`, seeded with the model's default seed.
149    ///
150    /// # Panics
151    ///
152    /// Panics as [`Self::new_seeded`] does.
153    pub fn new(ctx: &GpuContext, params: &[ParamValue]) -> Self {
154        Self::new_seeded(ctx, params, None)
155    }
156
157    /// Returns the geometry for `params`, without touching a device. `limits` sets only the display cap.
158    pub fn geometry_for(params: &[ParamValue], limits: &wgpu::Limits) -> Geometry {
159        let (num_agents, extent) = M::dims(params);
160        let num_agents = num_agents.max(1);
161        let extent = Extent {
162            w: extent.w.max(1.0),
163            h: extent.h.max(1.0),
164        };
165        let (width, height) = extent.cells();
166        Geometry {
167            num_agents,
168            extent,
169            width,
170            height,
171            n_cells: width * height,
172            display: display_dims(width, height, limits.max_texture_dimension_2d),
173            index: M::INDEX.then(|| HashGrid::new(extent, M::index_cell_size(params))),
174        }
175    }
176
177    /// Resources that would be allocated for this model based on `params`.
178    pub fn demand(params: &[ParamValue], limits: &wgpu::Limits) -> Demand {
179        let geom = Self::geometry_for(params, limits);
180        let mut demand = Demand::default();
181        for (spec, len) in M::BUFFERS.iter().zip(M::buffer_lens(&geom)) {
182            demand.push_sides(&format!("{}_{}", M::ID, spec.label), len, spec.double_buffered);
183        }
184        if let Some(grid) = geom.index {
185            demand.push_index(M::ID, grid.num_cells(), geom.num_agents);
186        }
187        if M::DISPLAY.is_some() {
188            demand.set_display(geom.width, geom.height, limits);
189        }
190
191        for (label, storage) in Self::declared_passes() {
192            demand.push_pass(label, storage);
193        }
194        demand
195    }
196
197    /// Storage buffers each declared pass binds. Read by both [`Self::demand`] and
198    /// [`Self::max_storage_bindings`], so the device that a host requests and the shortfall that the UI
199    /// reports cannot disagree.
200    fn declared_passes() -> Vec<(String, u32)> {
201        let mut passes: Vec<(String, u32)> = M::STEP_PASSES
202            .iter()
203            .map(|spec| (format!("{}_{}", M::ID, spec.label), storage_bindings(spec.bindings)))
204            .collect();
205        if let Some(spec) = &M::DISPLAY {
206            passes.push((format!("{}_display", M::ID), storage_bindings(spec.bindings)));
207        }
208        passes.push((format!("{}_reduce_leaf", M::ID), storage_bindings(M::REDUCE.bindings)));
209        for action in M::ACTIONS {
210            passes.push((
211                format!("{}_action_{}", M::ID, action.desc.id),
212                storage_bindings(action.pass.bindings),
213            ));
214        }
215        passes
216    }
217
218    /// Returns the most storage buffers any declared pass binds. It does not depend on params, so a host can query it
219    /// before it has a device.
220    pub fn max_storage_bindings() -> u32 {
221        Self::declared_passes()
222            .into_iter()
223            .map(|(_, storage)| storage)
224            .max()
225            .unwrap_or(0)
226    }
227
228    /// Builds the state on `ctx`, seeded with `seed`, or with the model's own default seed when it is `None`.
229    ///
230    /// # Panics
231    ///
232    /// Panics if the device cannot hold the model. A host calls [`Self::demand`] first, and this assert is the
233    /// backstop. Also panics if a shader declares a workgroup size other than the one its pass dispatches, or a
234    /// buffer label is reserved or ends in `_in` or `_out`.
235    #[expect(clippy::too_many_lines, reason = "one linear construction of every wgpu object")]
236    #[cfg_attr(
237        all(target_arch = "wasm32", target_feature = "atomics"),
238        expect(
239            clippy::arc_with_non_send_sync,
240            reason = "the agent layers hold wgpu buffers, which atomics leave unsendable"
241        )
242    )]
243    pub fn new_seeded(ctx: &GpuContext, params: &[ParamValue], seed: Option<u64>) -> Self {
244        let device = &ctx.device;
245        let queue = &ctx.queue;
246
247        let limits = device.limits();
248        let shortfalls = Self::demand(params, &limits).shortfalls(&limits);
249        assert!(
250            shortfalls.is_empty(),
251            "{} does not fit this device at these params: {}",
252            M::ID,
253            shortfalls.join("; ")
254        );
255        assert_buffer_labels(M::ID, M::BUFFERS.iter().map(|spec| spec.label));
256        // Every pass but the display folds a linear domain onto workgroups of `WORKGROUP`.
257        let linear = [WORKGROUP, 1, 1];
258        for (pass, shader) in M::STEP_PASSES
259            .iter()
260            .map(|spec| (spec.label, spec.shader))
261            .chain([("reduce leaf", M::REDUCE.shader)])
262            .chain(M::ACTIONS.iter().map(|action| (action.desc.id, action.pass.shader)))
263        {
264            assert_workgroup_size(M::ID, pass, shader, linear);
265        }
266        if let Some(spec) = &M::DISPLAY {
267            assert_workgroup_size(M::ID, "display", spec.shader, [spec.workgroup, spec.workgroup, 1]);
268        }
269
270        let mut geom = Self::geometry_for(params, &limits);
271        let (num_agents, extent) = (geom.num_agents, geom.extent);
272        let (width, height) = (geom.width, geom.height);
273
274        // The hash fits its own grid to the index cell size. That size can differ from the field's cell size.
275        let index =
276            M::INDEX.then(|| GpuSpatialHash::new(device, queue, M::ID, extent, M::index_cell_size(params), num_agents));
277        geom.index = index.as_ref().map(GpuSpatialHash::grid);
278
279        // Buffers.
280        let lens = M::buffer_lens(&geom);
281        assert_eq!(
282            lens.len(),
283            M::BUFFERS.len(),
284            "{}: buffer_lens must return one length per BUFFERS entry",
285            M::ID
286        );
287        let seeds = M::seed_buffers(&geom, params, seed);
288        assert_eq!(
289            seeds.len(),
290            M::BUFFERS.len(),
291            "{}: seed_buffers must return one vector per BUFFERS entry",
292            M::ID
293        );
294
295        let buffers: Vec<BufferSides> = M::BUFFERS
296            .iter()
297            .zip(&lens)
298            .map(|(spec, &len)| {
299                let make_buffer = |side: &str| {
300                    let label = format!("{}_{}_{side}", M::ID, spec.label);
301                    if spec.drawable {
302                        lane_buffer(device, &label, len)
303                    } else {
304                        storage_buffer(device, &label, len)
305                    }
306                };
307                BufferSides {
308                    a: make_buffer("a"),
309                    // Only the current side is seeded, so the other side is written by the first step.
310                    b: spec.double_buffered.then(|| make_buffer("b")),
311                }
312            })
313            .collect();
314
315        let mut clear = device.create_command_encoder(&wgpu::CommandEncoderDescriptor {
316            label: Some(&format!("{}_seed", M::ID)),
317        });
318        for (k, (buffer, bytes)) in buffers.iter().zip(&seeds).enumerate() {
319            if bytes.is_empty() {
320                clear.clear_buffer(&buffer.a, 0, None);
321                continue;
322            }
323            assert_eq!(
324                bytes.len(),
325                lens[k] * std::mem::size_of::<u32>(),
326                "{}: seed for buffer '{}' must be buffer_lens[{k}] words of bytes",
327                M::ID,
328                M::BUFFERS[k].label
329            );
330            queue.write_buffer(&buffer.a, 0, bytes);
331        }
332        queue.submit(Some(clear.finish()));
333
334        let ping_pong = M::BUFFERS.iter().any(|spec| spec.double_buffered);
335        let index_binds = index.as_ref().map(|hash| {
336            let pos = &buffers[M::POS_BUFFER];
337            Sides {
338                a: hash.bind_positions(device, &format!("{}_hash_bind_a", M::ID), &pos.a),
339                b: pos
340                    .b
341                    .as_ref()
342                    .map(|b| hash.bind_positions(device, &format!("{}_hash_bind_b", M::ID), b)),
343            }
344        });
345
346        let counters =
347            (M::COUNTERS > 0).then(|| CounterReadback::new(device, &format!("{}_counters", M::ID), M::COUNTERS));
348
349        let display_spec = M::DISPLAY;
350        // The display pass binds the view, and the snapshot carries the handle.
351        let (display_view, display_handle) = match display_spec
352            .as_ref()
353            .map(|_| build_display_target(device, ctx.target_format, width, height))
354        {
355            Some(DisplayTarget { view, display, .. }) => (Some(view), Some(display)),
356            None => (None, None),
357        };
358
359        let reduce_domain = M::REDUCE.domain.invocations(&geom);
360        let reduce = GpuLaneReduce::new(device, queue, M::ID, M::REDUCE.lanes, reduce_domain);
361
362        let build = PassBuilder {
363            device,
364            queue,
365            geom: &geom,
366            params,
367            buffers: &buffers,
368            ping_pong,
369            index: index.as_ref(),
370            counters: counters.as_ref(),
371            display_view: display_view.as_ref(),
372            reduce: &reduce,
373        };
374
375        let steps: Vec<EncodedPass> = M::STEP_PASSES
376            .iter()
377            .enumerate()
378            .map(|(i, spec)| {
379                let invocations = spec.domain.invocations(&geom);
380                build.linear_pass::<M>(PassId::Step(i), spec.label, spec.shader, spec.bindings, invocations)
381            })
382            .collect();
383
384        // One invocation per texel, as the grid engine's display pass does.
385        let display = display_spec.as_ref().zip(display_handle).map(|(spec, handle)| {
386            let (tex_w, tex_h) = geom.display;
387            let groups = (tex_w.div_ceil(spec.workgroup), tex_h.div_ceil(spec.workgroup));
388            let pass = build.pass::<M>(
389                PassId::Display,
390                "display",
391                spec.shader,
392                spec.bindings,
393                groups,
394                tex_w * tex_h,
395            );
396            (pass, handle)
397        });
398
399        let reduce_groups = reduce.agent_groups();
400        let reduce_pass = build.pass::<M>(
401            PassId::Reduce,
402            "reduce_leaf",
403            M::REDUCE.shader,
404            M::REDUCE.bindings,
405            reduce_groups,
406            reduce_domain,
407        );
408
409        // Truncated from the same stream the CPU engines use, since the WGSL generator is 32 bit.
410        let action_seed = henad_core::action::action_seed(seed) as u32;
411        let actions: Vec<EncodedPass> = M::ACTIONS
412            .iter()
413            .enumerate()
414            .map(|(i, action)| {
415                let spec = &action.pass;
416                let invocations = spec.domain.invocations(&geom);
417                build.pass_in_place::<M>(
418                    PassId::Action(i),
419                    &format!("action_{}", action.desc.id),
420                    spec.shader,
421                    spec.bindings,
422                    linear_dispatch(invocations),
423                    invocations,
424                    true,
425                    action_seed,
426                )
427            })
428            .collect();
429
430        let make_agents = |a_is_current: bool| {
431            let (pos, _) = buffers[M::POS_BUFFER].sides(a_is_current);
432            let (color, _) = buffers[M::COLOR_BUFFER].sides(a_is_current);
433            Arc::new(GpuAgents {
434                pos: pos.clone(),
435                color: color.clone(),
436                count: num_agents,
437                world_w: extent.w,
438                world_h: extent.h,
439            })
440        };
441        let agents = Sides {
442            a: make_agents(true),
443            b: ping_pong.then(|| make_agents(false)),
444        };
445
446        Self {
447            geom,
448            tick: 0,
449            device: device.clone(),
450            queue: queue.clone(),
451            buffers,
452            current_is_a: true,
453            ping_pong,
454            index,
455            index_binds,
456            steps,
457            display,
458            reduce,
459            reduce_pass,
460            counters,
461            actions,
462            action_seed,
463            params: params.to_vec(),
464            agents,
465            _marker: PhantomData,
466        }
467    }
468
469    /// Submits `steps` steps, at most [`MAX_STEPS_PER_SUBMISSION`] per command buffer, as the runner does. Nothing
470    /// waits on the GPU.
471    pub fn run_batched(&mut self, steps: u32) {
472        let mut remaining = steps;
473        while remaining > 0 {
474            let batch = remaining.min(MAX_STEPS_PER_SUBMISSION);
475            let mut encoder = self.device.create_command_encoder(&wgpu::CommandEncoderDescriptor {
476                label: Some("henad_gpu_agent_batch"),
477            });
478            self.encode_steps(&mut encoder, batch, None);
479            self.queue.submit(Some(encoder.finish()));
480            remaining -= batch;
481        }
482    }
483
484    /// Runs the snapshot passes and waits for the readback, so `stats()` reports the current tick.
485    pub fn refresh_stats(&mut self) {
486        let mut encoder = self.device.create_command_encoder(&wgpu::CommandEncoderDescriptor {
487            label: Some("henad_gpu_agent_snapshot"),
488        });
489        self.encode_snapshot_passes(&mut encoder);
490        let device = self.device.clone();
491        self.queue.submit(Some(encoder.finish()));
492        self.begin_stats_readback();
493        self.poll_stats_readback(&device, true);
494    }
495
496    /// Returns the current side of buffer `index` as raw words, blocking on the GPU.
497    ///
498    /// # Panics
499    ///
500    /// Panics if the readback fails.
501    pub fn read_buffer(&self, index: usize) -> Vec<u32> {
502        let (buffer, _) = self.buffers[index].sides(self.current_is_a);
503        let size = buffer.size();
504        let staging = self.device.create_buffer(&wgpu::BufferDescriptor {
505            label: Some("henad_gpu_agent_readback"),
506            size,
507            usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
508            mapped_at_creation: false,
509        });
510        let mut encoder = self.device.create_command_encoder(&wgpu::CommandEncoderDescriptor {
511            label: Some("henad_gpu_agent_readback"),
512        });
513        encoder.copy_buffer_to_buffer(buffer, 0, &staging, 0, size);
514        self.queue.submit(Some(encoder.finish()));
515
516        let (tx, rx) = flume::bounded(1);
517        staging
518            .slice(..)
519            .map_async(wgpu::MapMode::Read, move |r| drop(tx.send(r)));
520        self.device
521            .poll(wgpu::PollType::wait_indefinitely())
522            .expect("readback poll");
523        rx.recv().expect("readback channel").expect("readback map");
524        let data = staging.slice(..).get_mapped_range().expect("readback range");
525        let out = bytemuck::cast_slice::<u8, u32>(&data).to_vec();
526        drop(data);
527        staging.unmap();
528        out
529    }
530
531    /// Geometry the state was built at.
532    pub fn geometry(&self) -> &Geometry {
533        &self.geom
534    }
535}
536
537/// Carries what every pass needs to resolve its bindings into wgpu objects.
538struct PassBuilder<'a> {
539    device: &'a wgpu::Device,
540    queue: &'a wgpu::Queue,
541    geom: &'a Geometry,
542    params: &'a [ParamValue],
543    buffers: &'a [BufferSides],
544    ping_pong: bool,
545    index: Option<&'a GpuSpatialHash>,
546    counters: Option<&'a CounterReadback>,
547    display_view: Option<&'a wgpu::TextureView>,
548    reduce: &'a GpuLaneReduce,
549}
550
551impl PassBuilder<'_> {
552    /// A pass over a linear invocation domain, folded onto the 2D workgroup grid.
553    fn linear_pass<M: GpuAgentModel>(
554        &self,
555        id: PassId,
556        label: &str,
557        shader: &str,
558        bindings: &[BindingDecl],
559        invocations: u32,
560    ) -> EncodedPass {
561        self.pass::<M>(id, label, shader, bindings, linear_dispatch(invocations), invocations)
562    }
563
564    /// Builds a pass over `groups` workgroups.
565    ///
566    /// `groups.0` is the fold width that `henad::dispatch::linear_index` expects, and goes in the uniform as
567    /// `groups_x`.
568    fn pass<M: GpuAgentModel>(
569        &self,
570        id: PassId,
571        label: &str,
572        shader: &str,
573        bindings: &[BindingDecl],
574        groups: (u32, u32),
575        invocations: u32,
576    ) -> EncodedPass {
577        self.pass_in_place::<M>(id, label, shader, bindings, groups, invocations, false, 0)
578    }
579
580    /// Builds a pass as [`Self::pass`] does, with `in_place` binding writes to the side that already holds the state.
581    ///
582    /// Nothing swaps after an action, so writing the far side would throw the work away.
583    #[expect(
584        clippy::too_many_arguments,
585        reason = "one call site, and every argument is a pass fact"
586    )]
587    fn pass_in_place<M: GpuAgentModel>(
588        &self,
589        id: PassId,
590        label: &str,
591        shader: &str,
592        bindings: &[BindingDecl],
593        groups: (u32, u32),
594        invocations: u32,
595        in_place: bool,
596        seed: u32,
597    ) -> EncodedPass {
598        let label = format!("{}_{label}", M::ID);
599
600        let uniform = uniform_buffer(
601            self.device,
602            self.queue,
603            &format!("{label}_params"),
604            &M::pass_params_bytes(
605                id,
606                PassCtx {
607                    geom: self.geom,
608                    invocations,
609                    groups_x: groups.0,
610                    seed,
611                },
612                self.params,
613            ),
614        );
615
616        let entries: Vec<wgpu::BindGroupLayoutEntry> = bindings
617            .iter()
618            .enumerate()
619            .map(|(i, decl)| layout_entry(i as u32, decl))
620            .collect();
621        let layout = self.device.create_bind_group_layout(&wgpu::BindGroupLayoutDescriptor {
622            label: Some(&format!("{label}_layout")),
623            entries: &entries,
624        });
625
626        let make_bind = |a_is_current: bool, side: &str| {
627            let entries: Vec<wgpu::BindGroupEntry<'_>> = bindings
628                .iter()
629                .enumerate()
630                .map(|(i, decl)| wgpu::BindGroupEntry {
631                    binding: i as u32,
632                    resource: self.resource::<M>(decl, a_is_current, in_place, &uniform),
633                })
634                .collect();
635            self.device.create_bind_group(&wgpu::BindGroupDescriptor {
636                label: Some(&format!("{label}_bind_{side}")),
637                layout: &layout,
638                entries: &entries,
639            })
640        };
641
642        let binds = Sides {
643            a: make_bind(true, "a"),
644            b: self.ping_pong.then(|| make_bind(false, "b")),
645        };
646
647        let pipeline = compute_pipeline(self.device, &label, shader, &layout);
648        EncodedPass {
649            label,
650            pipeline,
651            binds,
652            groups,
653            uniform,
654        }
655    }
656
657    /// Returns the resource that a shader binding's name refers to. The name alone links the binding to its resource.
658    ///
659    /// Panics rather than returning an error, since a name that resolves to nothing is a shader
660    /// and model that disagree, and no run can be correct after that.
661    fn resource<'r, M: GpuAgentModel>(
662        &'r self,
663        decl: &BindingDecl,
664        a_is_current: bool,
665        in_place: bool,
666        uniform: &'r wgpu::Buffer,
667    ) -> wgpu::BindingResource<'r> {
668        if let Some((label, writes)) = buffer_target(decl) {
669            let k = M::BUFFERS
670                .iter()
671                .position(|spec| spec.label == label)
672                .unwrap_or_else(|| panic!("{}: no buffer labelled `{label}`, wanted by `{}`", M::ID, decl.name));
673            let (read, write) = self.buffers[k].sides(a_is_current);
674            return if writes && !in_place { write } else { read }.as_entire_binding();
675        }
676        match decl.name {
677            "params" => uniform.as_entire_binding(),
678            "cell_start" => self.index.expect("INDEX is declared").cell_start_binding(),
679            "sorted" => self.index.expect("INDEX is declared").sorted_binding(),
680            "counters" => self.counters.expect("COUNTERS is non-zero").binding(),
681            "partials" => self.reduce.partials_binding(),
682            "output" => wgpu::BindingResource::TextureView(self.display_view.expect("DISPLAY is declared")),
683            other => panic!("{}: `{other}` is reserved but the engine has no resource for it", M::ID),
684        }
685    }
686}
687
688impl<M: GpuAgentModel> SimState for GpuAgentState<M> {
689    /// Steps once in its own submission, for a caller that holds only a `SimState`. The GPU runner batches
690    /// instead.
691    fn step(&mut self) {
692        let mut encoder = self.device.create_command_encoder(&wgpu::CommandEncoderDescriptor {
693            label: Some("gpu_agent_single_step"),
694        });
695        self.encode_steps(&mut encoder, 1, None);
696        self.queue.submit(Some(encoder.finish()));
697    }
698
699    fn tick(&self) -> u64 {
700        self.tick
701    }
702
703    fn stats(&self) -> Vec<StatEntry> {
704        let counters = self.counters.as_ref().map_or(&[][..], CounterReadback::values);
705        stat_entries(M::STATS, M::stats(&self.reduce.sums(), counters, &self.geom))
706    }
707
708    /// Encodes the action into its own submission.
709    ///
710    /// The GPU runner calls [`GpuSimState::encode_action`] instead, and submits the action before
711    /// its snapshot.
712    fn act(&mut self, index: usize) -> bool {
713        let mut encoder = self.device.create_command_encoder(&wgpu::CommandEncoderDescriptor {
714            label: Some("henad_gpu_agent_action"),
715        });
716        if !GpuSimState::encode_action(self, &mut encoder, index) {
717            return false;
718        }
719        self.queue.submit(Some(encoder.finish()));
720        true
721    }
722
723    /// Rejects every live edit. Its entry declares every parameter reload-only.
724    fn set_param(&mut self, _index: usize, _value: &ParamValue) -> bool {
725        false
726    }
727
728    fn population(&self) -> u64 {
729        u64::from(self.geom.num_agents)
730    }
731
732    fn heap_bytes(&self) -> usize {
733        let buffers: usize = self
734            .buffers
735            .iter()
736            .map(|sides| (sides.a.size() + sides.b.as_ref().map_or(0, wgpu::Buffer::size)) as usize)
737            .sum();
738        let display = self.display.as_ref().map_or(0, |_| {
739            (self.geom.display.0 as usize) * (self.geom.display.1 as usize) * 4
740        });
741        buffers
742            + display
743            + self.index.as_ref().map_or(0, GpuSpatialHash::heap_bytes)
744            + self.reduce.heap_bytes()
745            + M::COUNTERS * std::mem::size_of::<u32>()
746    }
747}
748
749impl<M: GpuAgentModel> GpuSimState for GpuAgentState<M> {
750    /// Records one compute pass per declared pass per step, all into one encoder.
751    ///
752    /// A pass is the synchronization boundary wgpu inserts barriers at, so a step's passes
753    /// cannot be collapsed into dispatches inside one pass. The later ones would read stale
754    /// data.
755    fn encode_steps(&mut self, encoder: &mut wgpu::CommandEncoder, count: u32, timestamps: Option<&wgpu::QuerySet>) {
756        if count == 0 {
757            return;
758        }
759        let last_pass = self.steps.len().saturating_sub(1);
760
761        for i in 0..count {
762            let is_first = i == 0;
763            let is_last = i == count - 1;
764
765            // A batch begins with the index rebuild, so the opening stamp goes on its counting
766            // pass. A stamp on an empty pass is silently never written.
767            if let (Some(hash), Some(binds)) = (&self.index, &self.index_binds) {
768                hash.encode_build(
769                    encoder,
770                    binds.pick(self.current_is_a),
771                    timestamps.filter(|_| is_first).map(|query_set| (query_set, 0)),
772                );
773            }
774
775            for (j, pass) in self.steps.iter().enumerate() {
776                let opening = is_first && j == 0 && self.index.is_none();
777                let closing = is_last && j == last_pass;
778                pass.encode(encoder, self.current_is_a, stamps(timestamps, opening, closing));
779            }
780
781            if self.ping_pong {
782                self.current_is_a = !self.current_is_a;
783            }
784        }
785
786        self.tick += u64::from(count);
787    }
788
789    fn encode_action(&mut self, encoder: &mut wgpu::CommandEncoder, index: usize) -> bool {
790        let Some(action) = self.actions.get(index) else {
791            return false;
792        };
793        // A fresh seed per press, so pressing twice draws twice. The write is queued before the
794        // encoder is submitted, so the pass reads the new value.
795        self.action_seed = self.action_seed.wrapping_mul(747_796_405).wrapping_add(2_891_336_453);
796        self.queue.write_buffer(
797            &action.uniform,
798            0,
799            &M::pass_params_bytes(
800                PassId::Action(index),
801                PassCtx {
802                    geom: &self.geom,
803                    invocations: M::ACTIONS[index].pass.domain.invocations(&self.geom),
804                    groups_x: action.groups.0,
805                    seed: self.action_seed,
806                },
807                &self.params,
808            ),
809        );
810        action.encode(encoder, self.current_is_a, None);
811        true
812    }
813
814    fn encode_snapshot_passes(&mut self, encoder: &mut wgpu::CommandEncoder) {
815        if let Some((pass, _)) = &self.display {
816            pass.encode(encoder, self.current_is_a, None);
817        }
818        self.encode_stats_passes(encoder);
819    }
820
821    fn encode_stats_passes(&mut self, encoder: &mut wgpu::CommandEncoder) {
822        self.reduce_pass.encode(encoder, self.current_is_a, None);
823        self.reduce.encode(encoder);
824        if let Some(counters) = &mut self.counters {
825            counters.encode_copy(encoder);
826        }
827    }
828
829    fn begin_stats_readback(&mut self) {
830        self.reduce.begin_readback();
831        if let Some(counters) = &mut self.counters {
832            counters.begin_map();
833        }
834    }
835
836    fn poll_stats_readback(&mut self, device: &wgpu::Device, block: bool) -> StatsPoll {
837        let sums_poll = self.reduce.poll_readback(device, block);
838        let counters_poll = match &mut self.counters {
839            Some(readback) if block => readback.poll_blocking(device),
840            Some(readback) => readback.poll(device),
841            None => StatsPoll::Landed,
842        };
843        match (sums_poll, counters_poll) {
844            (StatsPoll::Pending, _) | (_, StatsPoll::Pending) => StatsPoll::Pending,
845            (StatsPoll::Failed, _) | (_, StatsPoll::Failed) => StatsPoll::Failed,
846            (StatsPoll::Landed, StatsPoll::Landed) => StatsPoll::Landed,
847        }
848    }
849
850    fn stats_readback_pending(&self) -> bool {
851        self.reduce.readback_pending() || self.counters.as_ref().is_some_and(CounterReadback::is_pending)
852    }
853
854    fn view(&self) -> GpuSnapshot {
855        GpuSnapshot {
856            display: self.display.as_ref().map(|(_, display)| Arc::clone(display)),
857            agents: Some(Arc::clone(self.agents.pick(self.current_is_a))),
858        }
859    }
860}