refac: split code into classes & add reset button

This commit is contained in:
Florian Sylvain
2025-04-06 05:23:20 +02:00
parent 9d1327e9c5
commit 296ea76d6b
2 changed files with 270 additions and 275 deletions
+1
View File
@@ -24,6 +24,7 @@
/>
</div>
</div>
<button id="reset-button">Reset Grid</button>
<script type="module" src="/src/main.ts"></script>
</body>
</html>
+247 -253
View File
@@ -3,19 +3,25 @@ import "./style.css"
const GRID_SIZE = 64
let COMPUTE_MS_INTERVAL = 100
function setupSliders() {
const computeIntervalSlider = document.getElementById(
class UIController {
static onReset(callback: () => void) {
document.getElementById("reset-button")?.addEventListener("click", callback)
}
static setup() {
const slider = document.getElementById(
"compute-interval-slider",
) as HTMLInputElement
const computeIntervalValue = document.getElementById(
const valueLabel = document.getElementById(
"compute-interval-value",
) as HTMLElement
computeIntervalSlider.addEventListener("input", () => {
COMPUTE_MS_INTERVAL = parseInt(computeIntervalSlider.value)
computeIntervalValue.textContent = computeIntervalSlider.value
slider.addEventListener("input", () => {
COMPUTE_MS_INTERVAL = parseInt(slider.value)
valueLabel.textContent = slider.value
})
}
}
async function getCanvas(): Promise<HTMLCanvasElement> {
const canvas = document.querySelector<HTMLCanvasElement>("#canvas")
@@ -23,102 +29,32 @@ async function getCanvas(): Promise<HTMLCanvasElement> {
return canvas
}
async function getAdapter(): Promise<GPUAdapter> {
if (!navigator.gpu) throw new Error("WebGPU not supported on this browser.")
async function getDevice(): Promise<GPUDevice> {
if (!navigator.gpu) throw new Error("WebGPU not supported")
const adapter = await navigator.gpu.requestAdapter()
if (!adapter) throw new Error("No appropriate GPUAdapter found.")
return adapter
if (!adapter) throw new Error("No GPUAdapter found")
return adapter.requestDevice()
}
async function getDevice(adapter: GPUAdapter): Promise<GPUDevice> {
return await adapter.requestDevice()
}
class Grid {
readonly uniformBuffer: GPUBuffer
function configureContext(
canvas: HTMLCanvasElement,
constructor(
readonly size: number,
device: GPUDevice,
): GPUCanvasContext {
const context = canvas.getContext("webgpu")
if (!context) throw new Error("Failed to get WebGPU context.")
const canvasFormat = navigator.gpu.getPreferredCanvasFormat()
context.configure({ device, format: canvasFormat })
return context
}
function createVertexBufferLayout(): GPUVertexBufferLayout {
return {
arrayStride: 8,
attributes: [
{ format: "float32x2" as GPUVertexFormat, offset: 0, shaderLocation: 0 },
],
}
}
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",
code: (await import("./shaders/cell.wgsl?raw")).default,
})
return device.createRenderPipeline({
label: "Cell pipeline",
layout: pipelineLayout,
vertex: {
module: cellShaderModule,
entryPoint: "vertexMain",
buffers: [createVertexBufferLayout()],
},
fragment: {
module: cellShaderModule,
entryPoint: "fragmentMain",
targets: [{ format: canvasFormat }],
},
})
}
function createGridUniformBuffer(device: GPUDevice): GPUBuffer {
const gridUniformArray = new Float32Array([GRID_SIZE, GRID_SIZE])
const gridUniformBuffer = device.createBuffer({
label: "Grid Uniforms",
size: gridUniformArray.byteLength,
) {
const data = new Float32Array([size, size])
this.uniformBuffer = device.createBuffer({
label: "Grid Uniform",
size: data.byteLength,
usage: GPUBufferUsage.UNIFORM | GPUBufferUsage.COPY_DST,
})
device.queue.writeBuffer(gridUniformBuffer, 0, gridUniformArray)
return gridUniformBuffer
device.queue.writeBuffer(this.uniformBuffer, 0, data)
}
}
function createTimeUniformBuffer(device: GPUDevice): GPUBuffer {
class Buffers {
static createTimeBuffer(device: GPUDevice): GPUBuffer {
return device.createBuffer({
label: "Time Uniform",
size: 4,
@@ -126,216 +62,274 @@ function createTimeUniformBuffer(device: GPUDevice): GPUBuffer {
})
}
function createStateStorageBuffers(device: GPUDevice): GPUBuffer[] {
const cellStateArray = new Uint32Array(GRID_SIZE * GRID_SIZE)
const size = cellStateArray.byteLength
const usage = GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_DST
const cellStateStorage = [
device.createBuffer({ label: "Cell State A", size, usage }),
device.createBuffer({ label: "Cell State B", size, usage }),
]
for (let i = 0; i < cellStateArray.length; ++i) {
cellStateArray[i] = Math.random() > 0.6 ? 1 : 0
}
device.queue.writeBuffer(cellStateStorage[0], 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,
gridUniformBuffer: GPUBuffer,
timeUniformBuffer: GPUBuffer,
cellStateStorage: GPUBuffer[],
bindGroupLayout: GPUBindGroupLayout,
): GPUBindGroup[] {
return [
device.createBindGroup({
label: "Cell renderer bind group A",
layout: bindGroupLayout,
entries: [
{ binding: 0, resource: { buffer: gridUniformBuffer } },
{ 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: bindGroupLayout,
entries: [
{ binding: 0, resource: { buffer: gridUniformBuffer } },
{ binding: 1, resource: { buffer: timeUniformBuffer } },
{ binding: 2, resource: { buffer: cellStateStorage[1] } },
{ binding: 3, resource: { buffer: cellStateStorage[0] } },
],
}),
]
}
function createVertexBuffer(device: GPUDevice): GPUBuffer {
const vertices = new Float32Array([
static createVertexBuffer(device: GPUDevice): [GPUBuffer, Float32Array] {
const data = 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,
])
const vertexBuffer = device.createBuffer({
const buffer = device.createBuffer({
label: "Cell vertices",
size: vertices.byteLength,
size: data.byteLength,
usage: GPUBufferUsage.VERTEX | GPUBufferUsage.COPY_DST,
})
device.queue.writeBuffer(vertexBuffer, 0, vertices)
return vertexBuffer
device.queue.writeBuffer(buffer, 0, data)
return [buffer, data]
}
static createStateBuffers(device: GPUDevice, size: number): GPUBuffer[] {
const data = new Uint32Array(size * size).map(() =>
Math.random() > 0.6 ? 1 : 0,
)
const usage = GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_DST
const buffers = [
device.createBuffer({ label: "State A", size: data.byteLength, usage }),
device.createBuffer({ label: "State B", size: data.byteLength, usage }),
]
device.queue.writeBuffer(buffers[0], 0, data)
device.queue.writeBuffer(buffers[1], 0, data)
return buffers
}
}
class PipelineFactory {
static createBindGroupLayout(device: GPUDevice): GPUBindGroupLayout {
return device.createBindGroupLayout({
label: "BindGroupLayout",
entries: [
{ binding: 0, visibility: 7, buffer: { type: "uniform" } },
{ binding: 1, visibility: 2, buffer: { type: "uniform" } },
{ binding: 2, visibility: 5, buffer: { type: "read-only-storage" } },
{ binding: 3, visibility: 4, buffer: { type: "storage" } },
],
})
}
static createPipelineLayout(device: GPUDevice, layout: GPUBindGroupLayout) {
return device.createPipelineLayout({ bindGroupLayouts: [layout] })
}
static async createCompute(
device: GPUDevice,
layout: GPUPipelineLayout,
): Promise<GPUComputePipeline> {
const module = device.createShaderModule({
label: "Compute shader",
code: (await import("./shaders/simulation.wgsl?raw")).default,
})
return device.createComputePipeline({
label: "Simulation pipeline",
layout,
compute: { module, entryPoint: "computeMain" },
})
}
static async createRender(
device: GPUDevice,
format: GPUTextureFormat,
layout: GPUPipelineLayout,
): Promise<GPURenderPipeline> {
const module = device.createShaderModule({
label: "Cell shader",
code: (await import("./shaders/cell.wgsl?raw")).default,
})
return device.createRenderPipeline({
label: "Cell pipeline",
layout,
vertex: {
module,
entryPoint: "vertexMain",
buffers: [
{
arrayStride: 8,
attributes: [{ format: "float32x2", offset: 0, shaderLocation: 0 }],
},
],
},
fragment: {
module,
entryPoint: "fragmentMain",
targets: [{ format }],
},
})
}
}
class Simulation {
public bindGroups: GPUBindGroup[]
private stateBuffers: GPUBuffer[]
constructor(
private device: GPUDevice,
grid: Grid,
private time: GPUBuffer,
state: GPUBuffer[],
layout: GPUBindGroupLayout,
) {
this.stateBuffers = state
this.bindGroups = this.createBindGroups(grid, layout)
}
reset(size: number) {
const data = new Uint32Array(size * size).map(() =>
Math.random() > 0.6 ? 1 : 0,
)
for (const buffer of this.stateBuffers) {
this.device.queue.writeBuffer(buffer, 0, data)
}
}
private createBindGroups(
grid: Grid,
layout: GPUBindGroupLayout,
): GPUBindGroup[] {
return [
this.device.createBindGroup({
layout,
entries: [
{ binding: 0, resource: { buffer: grid.uniformBuffer } },
{ binding: 1, resource: { buffer: this.time } },
{ binding: 2, resource: { buffer: this.stateBuffers[0] } },
{ binding: 3, resource: { buffer: this.stateBuffers[1] } },
],
}),
this.device.createBindGroup({
layout,
entries: [
{ binding: 0, resource: { buffer: grid.uniformBuffer } },
{ binding: 1, resource: { buffer: this.time } },
{ binding: 2, resource: { buffer: this.stateBuffers[1] } },
{ binding: 3, resource: { buffer: this.stateBuffers[0] } },
],
}),
]
}
}
class Renderer {
private lastFrameTime: number = 0
private frameCount: number = 0
private fps: number = 0
private bindGroupIndex: number = 0
private lastBindGroupSwitchTime: number = 0
private lastFrame = 0
private frameCount = 0
private fps = 0
private bindGroupIndex = 0
private lastSwitch = 0
constructor(
private device: GPUDevice,
private context: GPUCanvasContext,
private cellPipeline: GPURenderPipeline,
private simulationPipeline: GPUComputePipeline,
private renderPipeline: GPURenderPipeline,
private computePipeline: GPUComputePipeline,
private vertexBuffer: GPUBuffer,
private bindGroups: GPUBindGroup[],
private timeBuffer: GPUBuffer,
private vertices: Float32Array,
private simulation: Simulation,
private timeBuffer: GPUBuffer,
) {}
private updateFpsCount(timeMs: number): void {
if (this.lastFrameTime === 0) this.lastFrameTime = timeMs
render(timeMs: number) {
this.updateFPS(timeMs)
const encoder = this.device.createCommandEncoder()
this.runComputePass(encoder, timeMs)
this.runRenderPass(encoder)
this.updateTimeBuffer(timeMs)
this.device.queue.submit([encoder.finish()])
requestAnimationFrame(this.render.bind(this))
}
private updateFPS(timeMs: number) {
if (!this.lastFrame) this.lastFrame = timeMs
this.frameCount++
if (timeMs - this.lastFrameTime >= 1000) {
if (timeMs - this.lastFrame >= 1000) {
this.fps = this.frameCount
this.frameCount = 0
this.lastFrameTime = timeMs
this.lastFrame = timeMs
document.querySelector("#fps")!.textContent = `FPS: ${this.fps}`
}
}
private updateBindGroupIndex(timeMs: number): void {
if (timeMs - this.lastBindGroupSwitchTime >= COMPUTE_MS_INTERVAL) {
this.bindGroupIndex = (this.bindGroupIndex + 1) % this.bindGroups.length
this.lastBindGroupSwitchTime = timeMs
}
private updateTimeBuffer(timeMs: number) {
const seconds = timeMs / 1000
this.device.queue.writeBuffer(
this.timeBuffer,
0,
new Float32Array([seconds]),
)
}
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.device.queue.writeBuffer(this.timeBuffer, 0, new Float32Array([time]))
this.updateBindGroupIndex(timeMs)
private runComputePass(encoder: GPUCommandEncoder, timeMs: number) {
if (timeMs - this.lastSwitch >= COMPUTE_MS_INTERVAL) {
this.bindGroupIndex = (this.bindGroupIndex + 1) % 2
this.lastSwitch = timeMs
}
const pass = encoder.beginComputePass()
pass.setPipeline(this.computePipeline)
pass.setBindGroup(0, this.simulation.bindGroups[this.bindGroupIndex])
pass.dispatchWorkgroups(Math.ceil(GRID_SIZE / 8), Math.ceil(GRID_SIZE / 8))
pass.end()
}
private runRenderPass(encoder: GPUCommandEncoder) {
const pass = encoder.beginRenderPass({
colorAttachments: [
{
view: this.context.getCurrentTexture().createView(),
loadOp: "clear",
clearValue: [0.0, 0.0, 0.4, 1.0],
clearValue: [0.0, 0.0, 0.4, 1],
storeOp: "store",
},
],
})
pass.setPipeline(this.cellPipeline)
pass.setPipeline(this.renderPipeline)
pass.setVertexBuffer(0, this.vertexBuffer)
pass.setBindGroup(0, this.bindGroups[this.bindGroupIndex])
pass.setBindGroup(0, this.simulation.bindGroups[this.bindGroupIndex])
pass.draw(this.vertices.length / 2, GRID_SIZE * GRID_SIZE)
pass.end()
this.device.queue.submit([encoder.finish()])
requestAnimationFrame((time) => this.render(time))
}
}
async function init() {
setupSliders()
class Engine {
async start() {
UIController.setup()
const device = await getDevice(await getAdapter())
const context = configureContext(await getCanvas(), device)
const canvasFormat = navigator.gpu.getPreferredCanvasFormat()
const device = await getDevice()
const canvas = await getCanvas()
const context = canvas.getContext("webgpu")!
const format = navigator.gpu.getPreferredCanvasFormat()
context.configure({ device, format })
const gridUniformBuffer = createGridUniformBuffer(device)
const timeBuffer = createTimeUniformBuffer(device)
const cellStateStorage = createStateStorageBuffers(device)
const bindGroupLayout = createBindGroupLayout(device)
const bindGroups = createBindGroups(
device,
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,
const grid = new Grid(GRID_SIZE, device)
const timeBuffer = Buffers.createTimeBuffer(device)
const [vertexBuffer, vertices] = Buffers.createVertexBuffer(device)
const stateBuffers = Buffers.createStateBuffers(device, GRID_SIZE)
const layout = PipelineFactory.createBindGroupLayout(device)
const pipelineLayout = PipelineFactory.createPipelineLayout(device, layout)
const [renderPipeline, computePipeline] = await Promise.all([
PipelineFactory.createRender(device, format, pipelineLayout),
PipelineFactory.createCompute(device, pipelineLayout),
])
const simulation = new Simulation(
device,
grid,
timeBuffer,
stateBuffers,
layout,
)
const renderer = new Renderer(
device,
context,
cellPipeline,
simulationPipeline,
renderPipeline,
computePipeline,
vertexBuffer,
bindGroups,
timeBuffer,
vertices,
simulation,
timeBuffer,
)
requestAnimationFrame((time) => renderer.render(time))
UIController.onReset(() => simulation.reset(GRID_SIZE))
requestAnimationFrame(renderer.render.bind(renderer))
}
}
init().catch((error) => {
console.error("Error initializing WebGPU:", error)
})
new Engine().start().catch(console.error)