Writing a GPU grid model¶
In this tutorial we'll build Conway's Game of Life (Life) on the GPU from scratch, where the grid lives in a storage buffer, every step is a compute dispatch, and the cells never leave the GPU.
Henad 0.3
This page describes Henad 0.3.
This page assumes you have worked through our CPU Grid Model tutorial first.
We will reuse its palette, and use it as a cross-check for correctness.
Every file on this page goes under src/ in the same project, made from the template as Your own project describes.
You'll also need a machine with a GPU that wgpu can drive with compute support, which it reaches through Vulkan, Metal or DirectX 12.
What we will write¶
A GpuGridModel is a trait that describes a grid stepped by compute shaders.
Recall that the CPU trait needed one function, step_cell.
This one needs 3 WGSL shaders and a mod.rs that declares some metadata and the initial state.
flowchart LR
I["<code>fn seed_buffers</code><br>configure initial state"] --> A
A["state buffer<br>two sides, swapped by the engine"] --> S["<code>step.wgsl</code><br>once per word"]
S -->|each tick| A
A -.->|on publish| D["<code>display.wgsl</code><br>once per texel"]
A -.->|on publish| R["<code>reduce.wgsl</code><br>once per cell"]
step.wgsl, display.wgsl, reduce.wgsl and the declarations in mod.rs are the pieces we need to write ourselves.
The three shaders differ in how they are dispatched:
| Shader | Dispatched | Runs |
|---|---|---|
step.wgsl |
one invocation per unit of state | every step |
display.wgsl |
one invocation per display texel | on publish, up to about 60 times a second |
reduce.wgsl |
one invocation per cell | on publish, alongside display |
The Henad engine handles the rest of the simulation, such as allocating both sides of the state buffer and swapping them after every step, building every pipeline and bind group from the shaders, batching steps into submissions, the display texture, and the snapshot the UI draws.
Let's get started by creating src/gpu_life/, containing mod.rs, step.wgsl, display.wgsl and reduce.wgsl.
Update rule step.wgsl¶
Let's write down the update rule first, as we did on the CPU. See the CPU Grid Model tutorial for details on the rules.
We will need a different representation of the grid, due to limits on storage buffer sizes, i.e. how much data we can easily store on the GPU.
Cells as bits¶
The trait's default is one u32 per cell and one step invocation per cell, and for a first model that default might be fine on smaller scales.
You can accept it and skip ahead to the Bindings, if you do not want to pack the cells into bits (yet).
However, for practical reasons, it is best to have them packed into bits any way.
A 100M-cell grid at one u32 per cell is 400 MB per side, which blows the 128 MiB a storage binding is guaranteed on a baseline device (see wgpu default limits).
The same grid at one bit per cell is 12.5 MB per side.
So we pack 32 cells into each u32, and pad every row (with 0s) up to a whole number of words1.
Cell x of row y is just bit x % 32 of word y * words_per_row + x / 32.
CELLS 0 ...... 31 │ 32 ..... 63 │ ... │ ... width-1 [.... padding .....]
WORDS ── word 0 ─ │ ── word 1 ─ │ ... │ ────────── last word ───────────
BITS 0 ...... 31 │ 0 ...... 31 │ ... │ 0 ........................... 31
One invocation per word
Packing changes data ownership. If the step still dispatched one invocation per cell, 32 invocations would share each output word, and every one of them would read-modify-write it, racing with the other 31. Our step therefore dispatches one invocation per word, so each word has exactly one owner and is written with one plain store.
Bindings¶
The shader binds the two sides of the state buffer and a small uniform:
@group(0) @binding(0) var<storage, read> state_in: array<u32>; // (1)!
@group(0) @binding(1) var<storage, read_write> state_out: array<u32>;
@group(0) @binding(2) var<uniform> params: vec2<u32>; // (2)!
- The engine binds each declaration by its name.
state_inandstate_outare the two sides of the buffer labelledstate, andparamsis the uniform. The order is up to the shader, and a read and write pair per buffer followed by the uniform is the shape every shipped step shader uses. - The uniform holds the content
mod.rsdecides to send, and Life needs nothing but the grid dimensions.
SWAR counting¶
A u32 consists of 32 independent one-bit lanes, so a bitwise operation on one word is equivalent to operating on 32 cells at once.
This computational style is called SWAR.
Instead of one 4-bit count per cell, we keep the count bit-sliced.
Picture the 32 neighbour counts of a word written out in binary, one count per column.
A bit-sliced representation stores that table by rows rather than by columns.
Word sb0 holds the lowest bit of every count, sb1 the next bit up, sb2 the one above that, and bit j of each word belongs to cell j.
Reading one cell's count back means picking bit j out of each word and reading those bits as a number, and no such read ever happens in the shader.
cell j 0 1 2 3 ... 31
count 3 0 8 2 ... 5
─── ─── ─── ─── ───
sb0 (weight 1) 1 0 0 0 ... 1
sb1 (weight 2) 1 0 0 1 ... 0
sb2 (weight 4) 0 0 0 0 ... 1
sb3 (weight 8) 0 0 1 0 ... 0
A count of 8 is 1000 in binary and needs that fourth row, so a complete count would take four words.
We will keep three and drop sb3 on purpose, because the rule of Life never needs it.
Without its top bit a count of 8 becomes 000, the same as a count of 0, and both of those kill the cell.
Every other count keeps its exact value in sb0 to sb2.
Summing one-bit inputs into a bit-sliced count is a job for a carry-save adder, and an adder is nothing but XOR and AND:
// One column of the adder tree. `sum` is the weight-w result, and `carry` feeds weight 2w.
struct Adder {
sum: u32,
carry: u32,
}
fn full_add(a: u32, b: u32, c: u32) -> Adder { // (1)!
let t = a ^ b;
return Adder(t ^ c, (a & b) | (c & t));
}
fn half_add(a: u32, b: u32) -> Adder { // (2)!
return Adder(a ^ b, a & b);
}
- Three one-bit inputs in, a sum bit and a carry bit out, for all 32 lanes at once.
- The same with two inputs.
Before any adding, each invocation gathers its neighbourhood. For a word of cells, the west neighbour of every cell is the same word shifted left by one bit, with bit 0 filled in from the previous word, and similarly for the east. A small struct carries the three words of one row:
// Preloaded row window, with west and east being the cells shifted by 1 bit left and right, respectively.
struct Row {
cells: u32, // bit j = cell (word*32 + j)
west: u32, // bit j = its west neighbour
east: u32, // bit j = its east neighbour
}
Loading a row¶
fn load_row(row: u32, word: u32, stride: u32, width: u32) -> Row {
let base = row * stride;
let mid = state_in[base + word];
let left = state_in[base + (word + stride - 1u) % stride]; // (1)!
let right = state_in[base + (word + 1u) % stride];
var r: Row;
r.cells = mid;
r.west = (mid << 1u) | (left >> 31u); // bit 0 comes from the previous word's bit 31
r.east = (mid >> 1u) | (right << 31u); // bit 31 comes from the next word's bit 0
// Those two shifts assume the grid's x-wrap lands on a word edge, which holds only when
// width % 32 == 0. When the last word is ragged, exactly two bits are wrong, and need to be fixed.
// When it isn't ragged, both patches rewrite the value that's already there.
let last = width - 1u;
if word == 0u { // (2)!
r.west = (r.west & ~1u) | ((left >> (last % 32u)) & 1u);
}
if word == last / 32u {
let b = last % 32u;
r.east = (r.east & ~(1u << b)) | ((right & 1u) << b);
}
return r;
}
- Neighbouring words wrap within the row through
% stride, so a row's first and last words see each other. That is the x half of the torus. - The two
ifpatches finish the job. A width that divides by 32 puts the wrap on a word edge, and the shifts above are already right. Any other width leaves a ragged last word, and exactly two bits come out wrong, one at each end of the row. The patches rewrite those two bits from the true wrap positions, and when nothing was wrong they rewrite the value already there, so there is no branch on raggedness.
The entry point¶
@compute
@workgroup_size(16, 16) // (1)!
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
let width = params.x;
let height = params.y;
let stride = (width + 31u) / 32u;
let word = global_id.x;
let y = global_id.y;
if word >= stride || y >= height { // (2)!
return;
}
let up = (y + height - 1u) % height; // (3)!
let down = (y + 1u) % height;
let r_up = load_row(up, word, stride, width);
let r_mid = load_row(y, word, stride, width);
let r_down = load_row(down, word, stride, width);
// Compress the 8 neighbours into weight-1 sums and weight-2 carries.
let a = full_add(r_up.west, r_up.cells, r_up.east);
let b = full_add(r_down.west, r_down.cells, r_down.east);
let c = half_add(r_mid.west, r_mid.east);
// Weight 1 adds the three sums into one bit, with a carry into weight 2.
let d = full_add(a.sum, b.sum, c.sum);
let sb0 = d.sum;
// Weight 2 adds four terms, the three stage-1 carries and d.carry.
let e = full_add(a.carry, b.carry, c.carry);
let f = half_add(e.sum, d.carry);
let sb1 = f.sum;
// Weight 4 adds two terms, and the weight-8 carry is dropped. Only n == 8 sets it, and n == 8
// has sb1 == 0, so the rule below already excludes it.
let sb2 = e.carry ^ f.carry; // (4)!
// A cell lives on a count of 3, and on 2 if it is alive already. Bit-sliced, 3 is 011 and 2 is
// 010. Both counts need sb2 == 0 and sb1 == 1, and the term (sb0 | cells) covers the bit that differs.
let alive = ~sb2 & sb1 & (sb0 | r_mid.cells); // (5)!
// Trailing bits of a ragged last word hold no cell, and nothing reads them. load_row's patches
// keep real cells off them, and display and reduce stop at the width. The layout still requires
// them to be zero, and the mask below clears them.
let cells_here = min(width - word * 32u, 32u);
var mask = 0xFFFFFFFFu;
if cells_here < 32u {
mask = (1u << cells_here) - 1u;
}
state_out[y * stride + word] = alive & mask; // (6)!
}
- Every shader of a grid model declares the same square workgroup, and the trait's
WORKGROUP_SIZEdefaults to 16 to match it. - The dispatch is rounded up to whole workgroups, so some invocations hang off the right and bottom edges of the grid. They leave before touching anything.
- The y half of the torus. Rows wrap with a plain modulo, since a row is a whole number of words.
- Here
sb3goes missing. The weight-8 carry out of this column would be the fourth row of the table above, and the tree never computes it. Dropping it is sound because 8 is1000in binary, so a count of 8 leavessb1clear and the rule below already rejects it. - The rule collapses beautifully in bit-sliced form. A cell survives on a count of 2, binary
010, and is born on 3, binary011. Both needsb2 == 0andsb1 == 1, and they differ only insb0, where a set bit means born regardless and a clear bit needs the cell already alive. One expression for all 32 cells. - One plain store, by the word's one owner. This is the line the ownership rule above protects.
That is the whole rule. It compiles to a few dozen bitwise instructions per word, with no loop and no branch over the cells, and all 32 lanes are resolved at once.
Display display.wgsl¶
On the CPU the engine built our display texture for us, indexing PALETTE by the cell value.
On the GPU we need to draw our texture instead, because only the model knows how a word of bits maps to colours.
#import henad::dims::{Dims, cell_at} // (1)!
@group(0) @binding(0) var<storage, read> state: array<u32>;
@group(0) @binding(1) var output: texture_storage_2d<rgba8unorm, write>;
@group(0) @binding(2) var<uniform> dims: Dims; // (2)!
// Palette matches the CPU model's `PALETTE`: dead = 0x15/0x15/0x15, alive = 0x00/0xE6/0x76.
const DEAD_COLOR: vec4<f32> = vec4<f32>(21.0 / 255.0, 21.0 / 255.0, 21.0 / 255.0, 1.0); // (3)!
const ALIVE_COLOR: vec4<f32> = vec4<f32>(0.0 / 255.0, 230.0 / 255.0, 118.0 / 255.0, 1.0);
@compute
@workgroup_size(16, 16)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
if (global_id.x >= dims.tex.x || global_id.y >= dims.tex.y) {
return;
}
let width = dims.grid.x;
let cell_xy = cell_at(global_id.xy, dims); // (4)!
let x = cell_xy.x;
let y = cell_xy.y;
// Read the containing word and extract this cell's bit, unlike the per-word step pass.
let words_per_row = (width + 31u) / 32u;
let word = state[y * words_per_row + (x / 32u)];
let cell = (word >> (x % 32u)) & 1u;
let color = select(DEAD_COLOR, ALIVE_COLOR, cell == 1u);
textureStore(output, vec2<i32>(global_id.xy), color);
}
- Shared WGSL ships with henad-core and can be reached with
#import henad::<module>, resolved at build time.henad::dimsholds theDimsstruct that every grid model's display and reduce shaders read. - Display and reduce bind
stateand their ownDimsuniform, carrying the grid size and the texture size. Our step uniform never reaches them. - The shader writes RGBA directly, so it carries its own copy of the two palette colours as WGSL constants. It is recommended to maintain consistency with the CPU palette.
- The pass dispatches one invocation per texel, never per cell. The texture is capped at 4096 a side, so a big grid is sampled, and
cell_atreads the cell attexel * grid / tex. Nothing special happens if this cap is not reached.
Why sample?
One texel per cell would cap the grid at the device's maximum texture dimension and cost four bytes per cell, which at 163842 is over a gigabyte of RGBA for something drawn into a panel roughly a thousand pixels wide. The GPU grid models page has the details.
Statistics reduce.wgsl¶
We would like to display a count of cells alive in the statistics.
On the CPU we counted with reduce_chunks at publish time, and the GPU equivalent is a reduction pass that runs at the same snapshot cadence:
#import henad::dims::Dims
@group(0) @binding(0) var<storage, read> state: array<u32>;
@group(0) @binding(1) var<storage, read_write> counters: atomic<u32>; // (1)!
@group(0) @binding(2) var<uniform> dims: Dims;
var<workgroup> partial: atomic<u32>; // (2)!
@compute
@workgroup_size(16, 16)
fn main(
@builtin(global_invocation_id) global_id: vec3<u32>,
@builtin(local_invocation_index) local_index: u32,
) {
if (local_index == 0u) {
atomicStore(&partial, 0u);
}
workgroupBarrier();
// The bounds check wraps the work in an `if`. Every invocation in the workgroup has to reach
// both barriers, and an early `return` in a partial grid tile would skip them.
let width = dims.grid.x;
let height = dims.grid.y;
if (global_id.x < width && global_id.y < height) { // (3)!
// Each invocation reads the word holding its cell and extracts the cell's bit. A per-word
// countOneBits would have to dispatch over words and mask off the padding bits of each
// row's last word, and this pass runs only when the stats are sampled.
let words_per_row = (width + 31u) / 32u;
let word = state[global_id.y * words_per_row + (global_id.x / 32u)];
if (((word >> (global_id.x % 32u)) & 1u) == 1u) {
atomicAdd(&partial, 1u);
}
}
workgroupBarrier();
if (local_index == 0u) {
atomicAdd(&counters, atomicLoad(&partial)); // (4)!
}
}
- One
u32counter per series inSTATS, which we declare in a moment. Life has one series, so this is a single atomic rather than an array. - Workgroup memory, shared by the 256 invocations of one workgroup and nobody else.
- The bounds check is an
ifaround the work rather than an earlyreturn. Every invocation in a workgroup has to reach both barriers, including the invocations hanging off the grid's ragged edge, and an early return would leave them stranded. - Folding locally first means one global atomic per 256 cells instead of one per cell, which keeps the pass negligible at the grid sizes the engine targets.
Implementing GpuGridModel¶
With the three shaders written, let's start on mod.rs.
use henad::authoring::prelude::*;
pub struct GpuLifeModel;
impl GpuGridModel for GpuLifeModel {}
The struct is empty, similar to the CPU model. A GPU model is const metadata with a few pure functions, and every buffer lives with the engine.
This won't compile yet.
Cargo compiles only the files src/lib.rs reaches, so first declare the module there, next to mod life;:
Now let's run cargo check and see what the compiler says is missing:
error[E0046]: not all trait items implemented, missing: `NAME`, `ID`, `DESCRIPTION`, `PALETTE`, `STATS`,
`BUFFERS`, `STEP_BINDINGS`, `DISPLAY_BINDINGS`, `REDUCE_BINDINGS`, `STEP_SHADER`,
`DISPLAY_SHADER`, `REDUCE_SHADER`, `param_descriptors`, `dims`, `seed_buffers`,
`step_params_bytes`, `stats`
--> src/gpu_life/mod.rs:5:1
|
5 | impl GpuGridModel for GpuLifeModel {}
| ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ missing 17 items in implementation
|
= help: implement the missing item: `const NAME: &'static str = "";`
= help: implement the missing item: `const BUFFERS: &'static [&'static str] = &[];`
= help: implement the missing item: `fn seed_buffers(_: u32, _: u32, _: &[ParamValue], _: Option<u64>) -> Vec<Vec<u32>> { todo!() }`
17 items is more than the CPU trait required, but 12 of them are one-line consts. We'll work down the list for the rest of this tutorial.
Identity¶
The impl starts with the NAME, ID and DESCRIPTION of the model, exactly as on the CPU.
impl GpuGridModel for GpuLifeModel {
const NAME: &'static str = "Game of Life (GPU)";
const ID: &'static str = "gpu_life"; // (1)!
const DESCRIPTION: &'static str = "Conway's Game of Life on a toroidal grid, stepped entirely on the GPU";
}
- The example port already uses the ID
gpu_game_of_life. A model set holds each ID once, and our ID differs so that both models can sit in one set, such as one that also holdshenad::models::example_models().
Colours¶
The stats UI still reads PALETTE, even though the display shader carries its own colours.
Back on the CPU page we left the palette outside the impl block for exactly this moment.
Make it pub in life.rs,
pub const PALETTE: [[u8; 4]; 2] = [
[0x15, 0x15, 0x15, 0xFF], // Dead
[0x00, 0xE6, 0x76, 0xFF], // Alive
];
and point the trait at it, so the chart shows the same colours on both backends:
Buffers and shaders¶
Next come the declarations with no CPU counterpart, the buffers the step ping-pongs and the three shaders we wrote:
const BUFFERS: &'static [&'static str] = &["state"]; // (1)!
const STEP_SHADER: &'static str = crate::shader_bindings::gpu_life::step::SHADER_STRING; // (2)!
const DISPLAY_SHADER: &'static str = crate::shader_bindings::gpu_life::display::SHADER_STRING;
const REDUCE_SHADER: &'static str = crate::shader_bindings::gpu_life::reduce::SHADER_STRING;
const STEP_BINDINGS: &'static [BindingDecl] = crate::binding_decls::bindings::GPU_LIFE_STEP; // (3)!
const DISPLAY_BINDINGS: &'static [BindingDecl] = crate::binding_decls::bindings::GPU_LIFE_DISPLAY;
const REDUCE_BINDINGS: &'static [BindingDecl] = crate::binding_decls::bindings::GPU_LIFE_REDUCE;
- One label per ping-ponged buffer. Life needs one buffer, and a shader's binding names refer to it,
state_inandstate_outabove. - The WGSL source, embedded as a string at build time.
- Each shader's
@group(0)declarations in@bindingorder, read off the source at build time, so the Rust side cannot disagree with the WGSL about what is bound where.
Both modules are generated when the crate builds.
The template's build.rs is already in the project and runs henad-build over the shaders, and henad::include_shaders!() at the top of src/lib.rs brings the output in:
//! Generates bindings for the shaders under `src`, and stamps the build of this crate's models.
fn main() -> Result<(), henad_build::ShaderBuildError> {
henad_build::stamp_commit();
henad_build::ShaderBuild::discover("src")?.generate()?;
Ok(())
}
It finds every .wgsl file under src itself, so a new shader needs no edit to the build file.
Our three shaders are already part of it, with nothing to list.
Each shader's path decides its names, so gpu_life/step.wgsl becomes shader_bindings::gpu_life::step and GPU_LIFE_STEP.
The shaders page has the rules a path follows.
BindingDecl comes from the prelude.
A model with a second buffer
Life keeps everything in one buffer. A model whose cells carry more than a step can recompute declares more buffers, and all of them swap sides together. The shipped GPU SIR keeps a per-cell random number generator in a second buffer, and its step shader binds two interleaved pairs before the uniform:
@group(0) @binding(0) var<storage, read> state_in: array<u32>;
@group(0) @binding(1) var<storage, read_write> state_out: array<u32>;
@group(0) @binding(2) var<storage, read> rng_in: array<u32>;
@group(0) @binding(3) var<storage, read_write> rng_out: array<u32>;
@group(0) @binding(4) var<uniform> params: Params;
The shipped display and reduce shaders bind only state in either case.
Sizes¶
Unlike a CPU grid model, nothing is prepended to a GPU model's parameter list, so width and height are ours to declare:
henad::params! {
const GRID_WIDTH = u32_param("grid_width", "Grid Width", 1024, 1, 16_384);
const GRID_HEIGHT = u32_param("grid_height", "Grid Height", 1024, 1, 16_384);
}
u32_param takes the ID, the label the UI shows, then the default, the minimum and the maximum.
The range goes up to 16384 a side, well past the CPU model's 10000, because the packed layout makes such a grid affordable.
Three functions then tell the engine how big everything is:
fn param_descriptors() -> Vec<ParamDescriptor> {
descriptors()
}
fn dims(params: &[ParamValue]) -> (u32, u32) { // (1)!
(
extract_u32(params, GRID_WIDTH, 1024),
extract_u32(params, GRID_HEIGHT, 1024),
)
}
fn buffer_lens(width: u32, height: u32) -> Vec<usize> { // (2)!
vec![words_per_row(width) * (height as usize)]
}
fn step_dims(width: u32, height: u32) -> (u32, u32) { // (3)!
(words_per_row(width) as u32, height)
}
- The grid size, read back out of the parameters. The engine clamps both dimensions to at least 1.
- One length per entry of
BUFFERS, inu32elements. The default is one element per cell, and our packed layout measures in words instead. - The step's dispatch domain, in invocations. The default is one per cell, and we override it to one per word, for the ownership reason above. Display and reduce are unaffected, because they only ever read.
words_per_row is the one helper the layout needs:
/// Returns the number of `u32` words in a row of `width` cells, at 32 cells to a word, rounded up.
pub fn words_per_row(width: u32) -> usize {
(width as usize).div_ceil(32)
}
Seeding¶
On the CPU the engine passed init a grid and a generator.
Here we build the initial buffer contents ourselves, on the CPU, and the engine uploads them once at construction.
For now the density stays hard-coded, as it did on the CPU page:
fn seed_buffers(width: u32, height: u32, _params: &[ParamValue], seed: Option<u64>) -> Vec<Vec<u32>> {
let rng = grid_init_rng(seed); // (1)!
vec![seed_random(width, height, 0.3, rng)] // (2)!
}
seedisSomewhen a caller requests a particular run, andNonewhile the app's Seed field shows Default.grid_init_rngmixes a given seed and falls back toGRID_INIT_SEEDwhen no seed is given, exactly as the CPU engine seeded ourinit, so both backends open on the same grid for the same seed.- One vector per entry of
BUFFERS, each exactly as long asbuffer_lenssaid.
The fill itself is the CPU init again, storing bits instead of bytes:
fn seed_random(width: u32, height: u32, density: f32, mut rng: u64) -> Vec<u32> {
let threshold = (density * u32::MAX as f32) as u32; // (1)!
let stride = words_per_row(width);
let mut words = vec![0u32; stride * (height as usize)];
for y in 0..height as usize {
for x in 0..width as usize {
if below(next_bits(&mut rng), threshold) { // (2)!
words[y * stride + (x / 32)] |= 1u32 << (x % 32); // (3)!
}
}
}
words
}
- The same threshold trick as the CPU page, for the same reason.
- The same generator, drawn once per cell in the same row-major order, and the same Bernoulli trial.
- Only the storage differs. A live cell sets its bit in the containing word, and the padding bits at the end of a ragged row stay zero.
Given identical parameters, the two backends therefore start from a bit-identical grid. We'll get real value out of that under Testing.
Finishing up¶
Three items are left: the stat series, the step's uniform and stats itself.
const STATS: &'static [StatDescriptor] = &[StatDescriptor::new("Alive", PALETTE[1])]; // (1)!
fn step_params_bytes(width: u32, height: u32, _params: &[ParamValue]) -> Vec<u8> { // (2)!
bytemuck::cast_slice(&[width, height]).to_vec()
}
fn stats(counts: &[u32]) -> Vec<StatValue> { // (3)!
vec![StatValue::Scalar(f64::from(counts[0]))]
}
- One series, coloured like a live cell. Its length has to match the number of counters
reduce.wgslaccumulates, and nothing checks that at compile time. - The step's uniform block as raw bytes.
paramsinstep.wgslis avec2<u32>, and twou32s laid end to end are exactly that. countsholds one entry per series, read back from the reduce pass. It arrives through an asynchronous readback rather than a stall, so a reported stat is a few milliseconds stale, and reads zero until the first readback completes.
The prelude holds every name these use, grid_init_rng included, so the file compiles.
Running it¶
The model compiles, but the app can only pick models from the set it receives, so we have to register it.
We register it in models() in src/lib.rs, as on the CPU page.
That page shows the template's src/lib.rs whole.
The mod gpu_life; line is already there, so first we import the function that registers a GPU grid model, if src/lib.rs does not import it already,
and add a line to models(), next to the insert lines already there:
The entry needs no device, and builds on whichever device the host passes it. On a machine with no usable adapter the app and the CLI leave the GPU models out of their lists, and a GPU model requested by ID is rejected with a message saying it needs a GPU. A GPU entry also carries a capacity check, so a grid too large for this device disables Build with a readable reason instead of crashing the process.
With the entry in place, we can finally run the model.
Make sure that --release is present to reach full performance.
Our model shows up as Game of Life (GPU) in the picker. Press Build, then play, and once it runs try a 16384×16384 grid, which is 268 million cells. See App tour for a quick overview of the UI.
Add --set grid_width=8192 --set grid_height=8192 to see the model at scale, and --global-warmup 1000 in front of --steps.
The warm-up runs untimed steps first, letting the GPU clocks ramp up and the first-use shader compilation get paid before anything is measured.
Then open http://127.0.0.1:8081.
The GPU models appear in the browser too, as long as it exposes WebGPU with compute support.
See App tour for a quick overview of the UI.
You should see the same gliders, now stepped on the GPU. We are now good to implement the missing features: a way to change the starting density without editing the source, and a test.
Parameters¶
Let's deal with the density first.
As on the CPU page, we hoist the hard-coded 0.3 out of the seeding and declare it as a parameter:
henad::params! {
const GRID_WIDTH = u32_param("grid_width", "Grid Width", 1024, 1, 16_384);
const GRID_HEIGHT = u32_param("grid_height", "Grid Height", 1024, 1, 16_384);
const DENSITY = f32_param("density", "Initial Density", 0.3, 0.0, 1.0, Some(0.01));
}
No .on_reload() this time.
Every parameter of a GPU model applies on reload whatever we declare, because the GPU state rejects a live edit, and the engine marks the whole list accordingly.
Then we read the value where the grid is seeded:
fn seed_buffers(width: u32, height: u32, params: &[ParamValue], seed: Option<u64>) -> Vec<Vec<u32>> {
let density = extract_f32(params, DENSITY, 0.3);
let rng = grid_init_rng(seed);
vec![seed_random(width, height, density, rng)]
}
f32_param and extract_f32 both come from the prelude, next to u32_param and extract_u32.
An operator sees exactly the three parameters we declared, in the order we declared them. Here it is for our model:
parameters for gpu_life (Game of Life (GPU)):
index=0 id=grid_width kind=u32 default=1024 min=1 max=16384 apply=reload label="Grid Width"
index=1 id=grid_height kind=u32 default=1024 min=1 max=16384 apply=reload label="Grid Height"
index=2 id=density kind=f32 default=0.3 min=0 max=1 apply=reload label="Initial Density"
Indexes looking normal!
On the CPU page density sat at index 2 in this list while DENSITY read 0, because the engine prepended width and height.
Nothing is prepended here, so DENSITY reads 2 and the list says 2.
Spelling the list out ourselves also lets it match the CPU model's composed list exactly, so a slider in the app means the same thing on either backend.
Testing¶
To convince ourselves the shaders are right, let's write a test. Life draws no random numbers during a step, and we seeded the grid from the same generator as the CPU model, so the two backends should produce the same results. That makes the CPU model a correctness oracle for this one, and the test is just a comparison:
#[cfg(test)]
mod tests {
use super::*;
use henad::engine::{GpuGridState, GridModelState};
use henad::gpu::{GpuContext, wgpu};
use henad::runner::{GpuSimState as _, SimState as _};
use henad::stats::StatEntry;
use henad::testing::{TestDeviceRequest, headless_test_device};
use crate::life::LifeModel;
fn alive(stats: &[StatEntry]) -> u64 {
match stats.first().map(|s| s.value.clone()) {
Some(StatValue::Scalar(v)) => v as u64,
other => panic!("expected a scalar Alive stat, got {other:?}"),
}
}
/// Runs display and reduce, then waits for the count to arrive, as a one-shot snapshot does.
fn refresh_stats(ctx: &GpuContext, state: &mut GpuGridState<GpuLifeModel>) { // (1)!
let mut encoder = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
state.encode_snapshot_passes(&mut encoder);
ctx.queue.submit(Some(encoder.finish()));
state.begin_stats_readback();
state.poll_stats_readback(&ctx.device, true);
}
#[test]
fn the_alive_count_matches_the_cpu_model() {
let Some(ctx) = headless_test_device(&TestDeviceRequest::baseline()) else { // (2)!
log::warn!("skipping the_alive_count_matches_the_cpu_model: no wgpu adapter available");
return;
};
// 50 is neither a multiple of 32 nor a power of two, so the ragged last word is covered.
let params = vec![ParamValue::U32(50), ParamValue::U32(30), ParamValue::F32(0.3)]; // (3)!
let mut gpu = GpuGridState::<GpuLifeModel>::new(&ctx, ¶ms);
let mut cpu = GridModelState::<LifeModel>::from_params(¶ms);
for tick in 0..10 {
refresh_stats(&ctx, &mut gpu);
assert_eq!(
alive(&gpu.stats()),
alive(&cpu.stats()),
"the GPU alive count must match the CPU model's at tick {tick}"
);
let mut encoder = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
gpu.encode_steps(&mut encoder, 1, None); // (4)!
ctx.queue.submit(Some(encoder.finish()));
cpu.step();
}
assert!(
alive(&cpu.stats()) > 0,
"the grid died out, so the comparison proves nothing"
);
}
}
- A GPU state reports whatever its last readback delivered, so a test has to drive the snapshot passes itself and wait, exactly as a one-shot snapshot in the app does.
headless_test_deviceopens a device with no window behind it, at the WebGPU baseline, and returnsNoneon a machine with no adapter, where the test skips. It comes fromhenad::testing, behind the facade'stestingfeature. SetHENAD_REQUIRE_GPU=1to turn that skip into a failure, so a green run actually means the GPU tests ran.- Both models take the same parameter vector, thanks to the list we spelled out.
- One step per submission here, for simplicity. The real runner encodes many steps per submission, capped at 64, because one oversized submission can trip the OS GPU watchdog and silently zero every later readback.
Registering the model also opted us into the template's test.
The test runs the testing kit over every model in models().
Where a device exists, it builds every GPU model on a stock baseline device and checks that the capacity check agrees with what actually builds.
It also checks that one submission of 64 steps reads back what 64 single steps do.
Run both tests with cargo test.
The finished files¶
Here is everything we wrote on this page, gathered into four files.
gpu_life/mod.rs completed
//! GPU Game of Life as `docs/guide/first-model/gpu-game-of-life.md` builds it.
//!
//! The id is `gpu_life`. The example model uses `gpu_game_of_life`, and a set holds each id once.
//!
//! The three shaders beside this file are copies of the example model's shaders, line for line.
//! No shader refers to a model id.
use henad::authoring::prelude::*;
use crate::life::PALETTE;
henad::params! {
const GRID_WIDTH = u32_param("grid_width", "Grid Width", 1024, 1, 16_384);
const GRID_HEIGHT = u32_param("grid_height", "Grid Height", 1024, 1, 16_384);
const DENSITY = f32_param("density", "Initial Density", 0.3, 0.0, 1.0, Some(0.01));
}
pub struct GpuLifeModel;
impl GpuGridModel for GpuLifeModel {
const NAME: &'static str = "Game of Life (GPU)";
const ID: &'static str = "gpu_life";
const DESCRIPTION: &'static str = "Conway's Game of Life on a toroidal grid, stepped entirely on the GPU";
const PALETTE: &'static [[u8; 4]] = &PALETTE;
const STATS: &'static [StatDescriptor] = &[StatDescriptor::new("Alive", PALETTE[1])];
const BUFFERS: &'static [&'static str] = &["state"];
const STEP_SHADER: &'static str = crate::shader_bindings::gpu_life::step::SHADER_STRING;
const DISPLAY_SHADER: &'static str = crate::shader_bindings::gpu_life::display::SHADER_STRING;
const REDUCE_SHADER: &'static str = crate::shader_bindings::gpu_life::reduce::SHADER_STRING;
const STEP_BINDINGS: &'static [BindingDecl] = crate::binding_decls::bindings::GPU_LIFE_STEP;
const DISPLAY_BINDINGS: &'static [BindingDecl] = crate::binding_decls::bindings::GPU_LIFE_DISPLAY;
const REDUCE_BINDINGS: &'static [BindingDecl] = crate::binding_decls::bindings::GPU_LIFE_REDUCE;
fn param_descriptors() -> Vec<ParamDescriptor> {
descriptors()
}
fn dims(params: &[ParamValue]) -> (u32, u32) {
(
extract_u32(params, GRID_WIDTH, 1024),
extract_u32(params, GRID_HEIGHT, 1024),
)
}
fn buffer_lens(width: u32, height: u32) -> Vec<usize> {
vec![words_per_row(width) * (height as usize)]
}
fn step_dims(width: u32, height: u32) -> (u32, u32) {
(words_per_row(width) as u32, height)
}
fn seed_buffers(width: u32, height: u32, params: &[ParamValue], seed: Option<u64>) -> Vec<Vec<u32>> {
let density = extract_f32(params, DENSITY, 0.3);
let rng = grid_init_rng(seed);
vec![seed_random(width, height, density, rng)]
}
fn step_params_bytes(width: u32, height: u32, _params: &[ParamValue]) -> Vec<u8> {
bytemuck::cast_slice(&[width, height]).to_vec()
}
fn stats(counts: &[u32]) -> Vec<StatValue> {
vec![StatValue::Scalar(f64::from(counts[0]))]
}
}
/// Returns the number of `u32` words in a row of `width` cells, at 32 cells to a word, rounded up.
pub fn words_per_row(width: u32) -> usize {
(width as usize).div_ceil(32)
}
fn seed_random(width: u32, height: u32, density: f32, mut rng: u64) -> Vec<u32> {
let threshold = (density * u32::MAX as f32) as u32;
let stride = words_per_row(width);
let mut words = vec![0u32; stride * (height as usize)];
for y in 0..height as usize {
for x in 0..width as usize {
if below(next_bits(&mut rng), threshold) {
words[y * stride + (x / 32)] |= 1u32 << (x % 32);
}
}
}
words
}
#[cfg(test)]
mod tests {
use super::*;
use henad::engine::{GpuGridState, GridModelState};
use henad::gpu::{GpuContext, wgpu};
use henad::runner::{GpuSimState as _, SimState as _};
use henad::stats::StatEntry;
use henad::testing::{TestDeviceRequest, headless_test_device};
use crate::life::LifeModel;
fn alive(stats: &[StatEntry]) -> u64 {
match stats.first().map(|s| s.value.clone()) {
Some(StatValue::Scalar(v)) => v as u64,
other => panic!("expected a scalar Alive stat, got {other:?}"),
}
}
/// Runs display and reduce, then waits for the count to arrive, as a one-shot snapshot does.
fn refresh_stats(ctx: &GpuContext, state: &mut GpuGridState<GpuLifeModel>) {
let mut encoder = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
state.encode_snapshot_passes(&mut encoder);
ctx.queue.submit(Some(encoder.finish()));
state.begin_stats_readback();
state.poll_stats_readback(&ctx.device, true);
}
#[test]
fn the_alive_count_matches_the_cpu_model() {
let Some(ctx) = headless_test_device(&TestDeviceRequest::baseline()) else {
log::warn!("skipping the_alive_count_matches_the_cpu_model: no wgpu adapter available");
return;
};
// 50 is neither a multiple of 32 nor a power of two, so the ragged last word is covered.
let params = vec![ParamValue::U32(50), ParamValue::U32(30), ParamValue::F32(0.3)];
let mut gpu = GpuGridState::<GpuLifeModel>::new(&ctx, ¶ms);
let mut cpu = GridModelState::<LifeModel>::from_params(¶ms);
for tick in 0..10 {
refresh_stats(&ctx, &mut gpu);
assert_eq!(
alive(&gpu.stats()),
alive(&cpu.stats()),
"the GPU alive count must match the CPU model's at tick {tick}"
);
let mut encoder = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
gpu.encode_steps(&mut encoder, 1, None);
ctx.queue.submit(Some(encoder.finish()));
cpu.step();
}
assert!(
alive(&cpu.stats()) > 0,
"the grid died out, so the comparison proves nothing"
);
}
}
gpu_life/step.wgsl completed
// Bit-packed Game of Life step, 32 cells per u32 and one invocation per word.
//
// The rule is evaluated SWAR-style. A u32 is 32 independent 1-bit lanes, and the neighbour count
// is kept bit-sliced, so sb0/sb1/sb2 each hold one bit position of all 32 counts rather than one
// 4-bit count per lane. Summing is then a carry-save adder made of plain XOR/AND, and all 32 cells
// resolve at once with no loop.
@group(0) @binding(0) var<storage, read> state_in: array<u32>;
@group(0) @binding(1) var<storage, read_write> state_out: array<u32>;
@group(0) @binding(2) var<uniform> params: vec2<u32>;
// Preloaded row window, with west and east being the cells shifted by 1 bit left and right, respectively.
struct Row {
cells: u32, // bit j = cell (word*32 + j)
west: u32, // bit j = its west neighbour
east: u32, // bit j = its east neighbour
}
// One column of the adder tree. `sum` is the weight-w result, and `carry` feeds weight 2w.
struct Adder {
sum: u32,
carry: u32,
}
fn full_add(a: u32, b: u32, c: u32) -> Adder {
let t = a ^ b;
return Adder(t ^ c, (a & b) | (c & t));
}
fn half_add(a: u32, b: u32) -> Adder {
return Adder(a ^ b, a & b);
}
fn load_row(row: u32, word: u32, stride: u32, width: u32) -> Row {
let base = row * stride;
let mid = state_in[base + word];
let left = state_in[base + (word + stride - 1u) % stride];
let right = state_in[base + (word + 1u) % stride];
var r: Row;
r.cells = mid;
r.west = (mid << 1u) | (left >> 31u); // bit 0 comes from the previous word's bit 31
r.east = (mid >> 1u) | (right << 31u); // bit 31 comes from the next word's bit 0
// Those two shifts assume the grid's x-wrap lands on a word edge, which holds only when
// width % 32 == 0. When the last word is ragged, exactly two bits are wrong, and need to be fixed.
// When it isn't ragged, both patches rewrite the value that's already there.
let last = width - 1u;
if word == 0u {
r.west = (r.west & ~1u) | ((left >> (last % 32u)) & 1u);
}
if word == last / 32u {
let b = last % 32u;
r.east = (r.east & ~(1u << b)) | ((right & 1u) << b);
}
return r;
}
@compute
@workgroup_size(16, 16)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
let width = params.x;
let height = params.y;
let stride = (width + 31u) / 32u;
let word = global_id.x;
let y = global_id.y;
if word >= stride || y >= height {
return;
}
let up = (y + height - 1u) % height;
let down = (y + 1u) % height;
let r_up = load_row(up, word, stride, width);
let r_mid = load_row(y, word, stride, width);
let r_down = load_row(down, word, stride, width);
// Compress the 8 neighbours into weight-1 sums and weight-2 carries.
let a = full_add(r_up.west, r_up.cells, r_up.east);
let b = full_add(r_down.west, r_down.cells, r_down.east);
let c = half_add(r_mid.west, r_mid.east);
// Weight 1 adds the three sums into one bit, with a carry into weight 2.
let d = full_add(a.sum, b.sum, c.sum);
let sb0 = d.sum;
// Weight 2 adds four terms, the three stage-1 carries and d.carry.
let e = full_add(a.carry, b.carry, c.carry);
let f = half_add(e.sum, d.carry);
let sb1 = f.sum;
// Weight 4 adds two terms, and the weight-8 carry is dropped. Only n == 8 sets it, and n == 8
// has sb1 == 0, so the rule below already excludes it.
let sb2 = e.carry ^ f.carry;
// A cell lives on a count of 3, and on 2 if it is alive already. Bit-sliced, 3 is 011 and 2 is
// 010. Both counts need sb2 == 0 and sb1 == 1, and the term (sb0 | cells) covers the bit that differs.
let alive = ~sb2 & sb1 & (sb0 | r_mid.cells);
// Trailing bits of a ragged last word hold no cell, and nothing reads them. load_row's patches
// keep real cells off them, and display and reduce stop at the width. The layout still requires
// them to be zero, and the mask below clears them.
let cells_here = min(width - word * 32u, 32u);
var mask = 0xFFFFFFFFu;
if cells_here < 32u {
mask = (1u << cells_here) - 1u;
}
state_out[y * stride + word] = alive & mask;
}
gpu_life/display.wgsl completed
#import henad::dims::{Dims, cell_at}
@group(0) @binding(0) var<storage, read> state: array<u32>;
@group(0) @binding(1) var output: texture_storage_2d<rgba8unorm, write>;
@group(0) @binding(2) var<uniform> dims: Dims;
// Palette matches the CPU model's `PALETTE`: dead = 0x15/0x15/0x15, alive = 0x00/0xE6/0x76.
const DEAD_COLOR: vec4<f32> = vec4<f32>(21.0 / 255.0, 21.0 / 255.0, 21.0 / 255.0, 1.0);
const ALIVE_COLOR: vec4<f32> = vec4<f32>(0.0 / 255.0, 230.0 / 255.0, 118.0 / 255.0, 1.0);
@compute
@workgroup_size(16, 16)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
if (global_id.x >= dims.tex.x || global_id.y >= dims.tex.y) {
return;
}
let width = dims.grid.x;
let cell_xy = cell_at(global_id.xy, dims);
let x = cell_xy.x;
let y = cell_xy.y;
// Read the containing word and extract this cell's bit, unlike the per-word step pass.
let words_per_row = (width + 31u) / 32u;
let word = state[y * words_per_row + (x / 32u)];
let cell = (word >> (x % 32u)) & 1u;
let color = select(DEAD_COLOR, ALIVE_COLOR, cell == 1u);
textureStore(output, vec2<i32>(global_id.xy), color);
}
gpu_life/reduce.wgsl completed
// Counts alive cells entirely on the GPU, so `SimState::stats()` never has to read the grid back
// to the CPU. The pass runs only when the stats are sampled.
//
// The reduction has two levels. Every invocation adds its cell into a workgroup-local atomic, then
// one invocation per workgroup adds that total into the global counter, once per 256 cells.
#import henad::dims::Dims
@group(0) @binding(0) var<storage, read> state: array<u32>;
@group(0) @binding(1) var<storage, read_write> counters: atomic<u32>;
@group(0) @binding(2) var<uniform> dims: Dims;
var<workgroup> partial: atomic<u32>;
@compute
@workgroup_size(16, 16)
fn main(
@builtin(global_invocation_id) global_id: vec3<u32>,
@builtin(local_invocation_index) local_index: u32,
) {
if (local_index == 0u) {
atomicStore(&partial, 0u);
}
workgroupBarrier();
// The bounds check wraps the work in an `if`. Every invocation in the workgroup has to reach
// both barriers, and an early `return` in a partial grid tile would skip them.
let width = dims.grid.x;
let height = dims.grid.y;
if (global_id.x < width && global_id.y < height) {
// Each invocation reads the word holding its cell and extracts the cell's bit. A per-word
// countOneBits would have to dispatch over words and mask off the padding bits of each
// row's last word, and this pass runs only when the stats are sampled.
let words_per_row = (width + 31u) / 32u;
let word = state[global_id.y * words_per_row + (global_id.x / 32u)];
if (((word >> (global_id.x % 32u)) & 1u) == 1u) {
atomicAdd(&partial, 1u);
}
}
workgroupBarrier();
if (local_index == 0u) {
atomicAdd(&counters, atomicLoad(&partial));
}
}
The listings above are stored in the repository at examples/tutorial/src/gpu_life/.
The three shaders there are copies of the example port's shaders, at crates/henad-models/src/gpu_game_of_life/, since a shader carries no model ID and what we wrote is the same file line for line.
The example model is at crates/henad-models/src/gpu_game_of_life/mod.rs.
It runs under its own ID, and its tests pin the adder tree and the ragged wrap.
Next¶
The GPU Agent model tutorial takes the ant colony to the GPU, where a step becomes a list of passes and determinism needs more careful handling.
For the trait from the reference side, read GPU grid models, and for the WGSL tooling, shaders and bindings.
-
A word is 32 bits, i.e. a
u32. ↩