1use 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
21pub(crate) type Dims = crate::shader_bindings::henad::dims::Dims;
23
24struct BufferPair {
26 a: wgpu::Buffer,
27 b: wgpu::Buffer,
28}
29
30struct 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
45pub struct GpuGridState<M: GpuGridModel> {
47 width: u32,
48 height: u32,
49 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 action_seed: u32,
73 params: Vec<ParamValue>,
76
77 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 pub fn new(ctx: &GpuContext, params: &[ParamValue]) -> Self {
101 Self::new_seeded(ctx, params, None)
102 }
103
104 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 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 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 #[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 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 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 queue.write_buffer(&pair.a, 0, bytemuck::cast_slice(seed));
235 pair
236 })
237 .collect();
238
239 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 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 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 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 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 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 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 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 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 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 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 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 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 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 fn view(&self) -> GpuSnapshot {
608 GpuSnapshot {
609 display: Some(Arc::clone(&self.display)),
610 agents: None,
611 }
612 }
613}