1use 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
27struct 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 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: 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
85fn 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
102pub 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 current_is_a: bool,
113 ping_pong: bool,
115
116 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 action_seed: u32,
130 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 pub fn new(ctx: &GpuContext, params: &[ParamValue]) -> Self {
154 Self::new_seeded(ctx, params, None)
155 }
156
157 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 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 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 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 #[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 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 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 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 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 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 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 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 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 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 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 pub fn geometry(&self) -> &Geometry {
533 &self.geom
534 }
535}
536
537struct 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 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 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 #[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 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 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 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 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 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 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 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}