feat: game of life compute shader

This commit is contained in:
Florian Sylvain
2025-04-06 04:15:05 +02:00
parent 63420f39a1
commit 69f1e7cdd9
3 changed files with 139 additions and 22 deletions
+1 -1
View File
@@ -7,7 +7,7 @@
<title>Whatever WebGPU</title>
</head>
<body>
<canvas id="canvas" width="512" height="512"></canvas>
<canvas id="canvas" width="1024" height="1024"></canvas>
<p id="fps">FPS:</p>
<script type="module" src="/src/main.ts"></script>
</body>
+98 -21
View File
@@ -1,6 +1,7 @@
import "./style.css"
const GRID_SIZE = 32
const GRID_SIZE = 64
const COMPUTE_MS_INTERVAL = 100
async function getCanvas(): Promise<HTMLCanvasElement> {
const canvas = document.querySelector<HTMLCanvasElement>("#canvas")
@@ -39,9 +40,37 @@ function createVertexBufferLayout(): GPUVertexBufferLayout {
}
}
function createPipelineLayout(
device: GPUDevice,
bindGroupLayout: GPUBindGroupLayout,
): GPUPipelineLayout {
return device.createPipelineLayout({
label: "Cell Pipeline Layout",
bindGroupLayouts: [bindGroupLayout],
})
}
async function createComputePipeline(
device: GPUDevice,
pipelineLayout: GPUPipelineLayout,
): Promise<GPUComputePipeline> {
return device.createComputePipeline({
label: "Simulation pipeline",
layout: pipelineLayout,
compute: {
module: device.createShaderModule({
label: "Game of Life simulation shader",
code: (await import("./shaders/simulation.wgsl?raw")).default,
}),
entryPoint: "computeMain",
},
})
}
async function createCellPipeline(
device: GPUDevice,
canvasFormat: GPUTextureFormat,
pipelineLayout: GPUPipelineLayout,
): Promise<GPURenderPipeline> {
const cellShaderModule = device.createShaderModule({
label: "Cell shader",
@@ -50,7 +79,7 @@ async function createCellPipeline(
return device.createRenderPipeline({
label: "Cell pipeline",
layout: "auto",
layout: pipelineLayout,
vertex: {
module: cellShaderModule,
entryPoint: "vertexMain",
@@ -75,7 +104,7 @@ function createGridUniformBuffer(device: GPUDevice): GPUBuffer {
return gridUniformBuffer
}
function createTimeBuffer(device: GPUDevice): GPUBuffer {
function createTimeUniformBuffer(device: GPUDevice): GPUBuffer {
return device.createBuffer({
label: "Time Uniform",
size: 4,
@@ -91,41 +120,70 @@ function createStateStorageBuffers(device: GPUDevice): GPUBuffer[] {
device.createBuffer({ label: "Cell State A", size, usage }),
device.createBuffer({ label: "Cell State B", size, usage }),
]
for (let i = 0; i < cellStateArray.length; i += 3) {
cellStateArray[i] = 1
for (let i = 0; i < cellStateArray.length; ++i) {
cellStateArray[i] = Math.random() > 0.6 ? 1 : 0
}
device.queue.writeBuffer(cellStateStorage[0], 0, cellStateArray)
for (let i = 0; i < cellStateArray.length; i++) {
cellStateArray[i] = i % 2
}
device.queue.writeBuffer(cellStateStorage[1], 0, cellStateArray)
return cellStateStorage
}
function createBindGroupLayout(device: GPUDevice): GPUBindGroupLayout {
return device.createBindGroupLayout({
label: "Cell Bind Group Layout",
entries: [
{
binding: 0,
visibility:
GPUShaderStage.FRAGMENT |
GPUShaderStage.VERTEX |
GPUShaderStage.COMPUTE,
buffer: { type: "uniform" },
},
{
binding: 1,
visibility: GPUShaderStage.FRAGMENT,
buffer: { type: "uniform" },
},
{
binding: 2,
visibility: GPUShaderStage.VERTEX | GPUShaderStage.COMPUTE,
buffer: { type: "read-only-storage" },
},
{
binding: 3,
visibility: GPUShaderStage.COMPUTE,
buffer: { type: "storage" },
},
],
})
}
function createBindGroups(
device: GPUDevice,
cellPipeline: GPURenderPipeline,
gridUniformBuffer: GPUBuffer,
timeBuffer: GPUBuffer,
timeUniformBuffer: GPUBuffer,
cellStateStorage: GPUBuffer[],
bindGroupLayout: GPUBindGroupLayout,
): GPUBindGroup[] {
return [
device.createBindGroup({
label: "Cell renderer bind group A",
layout: cellPipeline.getBindGroupLayout(0),
layout: bindGroupLayout,
entries: [
{ binding: 0, resource: { buffer: gridUniformBuffer } },
{ binding: 1, resource: { buffer: timeBuffer } },
{ binding: 1, resource: { buffer: timeUniformBuffer } },
{ binding: 2, resource: { buffer: cellStateStorage[0] } },
{ binding: 3, resource: { buffer: cellStateStorage[1] } },
],
}),
device.createBindGroup({
label: "Cell updater bind group B",
layout: cellPipeline.getBindGroupLayout(0),
layout: bindGroupLayout,
entries: [
{ binding: 0, resource: { buffer: gridUniformBuffer } },
{ binding: 1, resource: { buffer: timeBuffer } },
{ binding: 1, resource: { buffer: timeUniformBuffer } },
{ binding: 2, resource: { buffer: cellStateStorage[1] } },
{ binding: 3, resource: { buffer: cellStateStorage[0] } },
],
}),
]
@@ -155,6 +213,7 @@ class Renderer {
private device: GPUDevice,
private context: GPUCanvasContext,
private cellPipeline: GPURenderPipeline,
private simulationPipeline: GPUComputePipeline,
private vertexBuffer: GPUBuffer,
private bindGroups: GPUBindGroup[],
private timeBuffer: GPUBuffer,
@@ -173,19 +232,29 @@ class Renderer {
}
private updateBindGroupIndex(timeMs: number): void {
if (timeMs - this.lastBindGroupSwitchTime >= 500) {
if (timeMs - this.lastBindGroupSwitchTime >= COMPUTE_MS_INTERVAL) {
this.bindGroupIndex = (this.bindGroupIndex + 1) % this.bindGroups.length
this.lastBindGroupSwitchTime = timeMs
}
}
public async render(timeMs: number): Promise<void> {
const encoder = this.device.createCommandEncoder()
const computePass = encoder.beginComputePass()
computePass.setPipeline(this.simulationPipeline)
computePass.setBindGroup(0, this.bindGroups[this.bindGroupIndex])
const workgroupCount = Math.ceil(GRID_SIZE / 8)
computePass.dispatchWorkgroups(workgroupCount, workgroupCount)
computePass.end()
const time = timeMs / 1000
this.updateFpsCount(timeMs)
this.updateBindGroupIndex(timeMs)
this.device.queue.writeBuffer(this.timeBuffer, 0, new Float32Array([time]))
this.updateBindGroupIndex(timeMs)
const encoder = this.device.createCommandEncoder()
const pass = encoder.beginRenderPass({
colorAttachments: [
{
@@ -215,17 +284,24 @@ async function init() {
const context = configureContext(await getCanvas(), device)
const canvasFormat = navigator.gpu.getPreferredCanvasFormat()
const cellPipeline = await createCellPipeline(device, canvasFormat)
const gridUniformBuffer = createGridUniformBuffer(device)
const timeBuffer = createTimeBuffer(device)
const timeBuffer = createTimeUniformBuffer(device)
const cellStateStorage = createStateStorageBuffers(device)
const bindGroupLayout = createBindGroupLayout(device)
const bindGroups = createBindGroups(
device,
cellPipeline,
gridUniformBuffer,
timeBuffer,
cellStateStorage,
bindGroupLayout,
)
const pipelineLayout = createPipelineLayout(device, bindGroupLayout)
const cellPipeline = await createCellPipeline(
device,
canvasFormat,
pipelineLayout,
)
const simulationPipeline = await createComputePipeline(device, pipelineLayout)
const vertexBuffer = createVertexBuffer(device)
const vertices = new Float32Array([
-0.8, -0.8, 0.8, -0.8, 0.8, 0.8, -0.8, -0.8, 0.8, 0.8, -0.8, 0.8,
@@ -235,6 +311,7 @@ async function init() {
device,
context,
cellPipeline,
simulationPipeline,
vertexBuffer,
bindGroups,
timeBuffer,
+40
View File
@@ -0,0 +1,40 @@
@group(0) @binding(0) var<uniform> grid: vec2f;
@group(0) @binding(2) var<storage> cellStateIn: array<u32>;
@group(0) @binding(3) var<storage, read_write> cellStateOut: array<u32>;
fn cellIndex(cell: vec2u) -> u32 {
return (cell.y % u32(grid.y)) * u32(grid.x) + (cell.x % u32(grid.x));
}
fn cellActive(x: u32, y: u32) -> u32 {
return cellStateIn[cellIndex(vec2(x, y))];
}
@compute
@workgroup_size(8, 8)
fn computeMain(@builtin(global_invocation_id) cell: vec3u) {
let activeNeighbors =
cellActive(cell.x+1, cell.y+1) +
cellActive(cell.x+1, cell.y) +
cellActive(cell.x+1, cell.y-1) +
cellActive(cell.x, cell.y-1) +
cellActive(cell.x-1, cell.y-1) +
cellActive(cell.x-1, cell.y) +
cellActive(cell.x-1, cell.y+1) +
cellActive(cell.x, cell.y+1);
let i = cellIndex(cell.xy);
switch activeNeighbors {
case 2: {
cellStateOut[i] = cellStateIn[i];
}
case 3: {
cellStateOut[i] = 1;
}
default: {
cellStateOut[i] = 0;
}
}
}