In an n-body simulation every body attracts every other: n x (n - 1) interactions per step. One invocation per body sums the pulls, tiled as in Workgroup Memory so each workgroup loads 256 bodies into shared memory at a time, and a softening term keeps close encounters finite. This lab shader computes accelerations for 4,096 bodies in a disc, scaled to integers for printing:
// out: 8192 x i32, dispatch 16
const N = 4096u; // bodies; the dispatch is N / 256
@group(0) @binding(0) var<storage, read_write> out: array<vec2i>; // accelerations
var<workgroup> tile: array<vec3f, 256>; // x, y, mass of 256 bodies
fn body(i: u32) -> vec3f { // a stand-in galaxy: hashed positions
let a = f32((i * 2654435761u) >> 8) / 16777216.0 * 6.2832;
let r = sqrt(f32(i) / f32(N));
return vec3f(r * cos(a), r * sin(a), 1.0 / f32(N));
}
@compute @workgroup_size(256) fn main(@builtin(global_invocation_id) g: vec3u,
@builtin(local_invocation_index) l: u32) {
let me = body(g.x);
var acc = vec2f(0);
for (var start = 0u; start < N; start += 256) {
tile[l] = body(start + l); // each invocation loads one body
workgroupBarrier();
for (var k = 0u; k < 256; k++) { // 256 interactions from shared memory
let d = tile[k].xy - me.xy;
let r2 = dot(d, d) + 0.0001; // softening: no infinite force at r = 0
acc += d * tile[k].z * inverseSqrt(r2 * r2 * r2);
}
workgroupBarrier();
}
out[g.x] = vec2i(round(acc * 1000)); // G = 1: a = sum of m d / |d|^3
}out: -352, -184, 299, 321, -59, -215, -155, 63, 167, -115, -104, 52, ... gpu: 0.26 ms (median of 5)
The six printed accelerations equal a double-precision JavaScript loop's, rounded the same way; that loop took 94-110 ms in Node.js 2,131 on this machine against 0.26-0.33 ms here. With 8,192 and 16,384 bodies the pass took 0.98 and 3.08 ms, quadratic growth at about 87 billion interactions per second. A full simulation adds an integration entry point and ping-pong position buffers.
<!doctype html>
<style>
body { margin: 0; background: #f7f4ee; font: 14px system-ui, sans-serif; }
.stage { position: relative; width: 100%; max-width: 600px; }
.stage canvas { display: block; width: 100%; }
.stage canvas + canvas { position: absolute; inset: 0; pointer-events: none; }
</style>
<div class="stage">
<canvas id="view" width="600" height="360"></canvas>
<canvas id="labels" width="600" height="360"></canvas>
</div>
<script>
const canvas = document.getElementById('view');
const ink = document.getElementById('labels').getContext('2d');
function showMessage(text) { // 2D fallback when WebGPU is missing
const ctx = canvas.getContext('2d');
ctx.fillStyle = '#fbeaea'; ctx.fillRect(0, 0, canvas.width, canvas.height);
ctx.fillStyle = '#8a2b2b'; ctx.font = '18px system-ui, sans-serif'; ctx.textAlign = 'center';
ctx.fillText(text, canvas.width / 2, canvas.height / 2);
}
const N = 1024;
const code = /* wgsl */ `
const N = ${N}u;
@group(0) @binding(0) var<storage> bodiesIn: array<vec4f>; // pos.xy, vel.xy
@group(0) @binding(1) var<storage, read_write> bodiesOut: array<vec4f>; // ping-pong: never read what we write
@group(0) @binding(2) var<uniform> dt: f32;
var<workgroup> tile: array<vec2f, 256>; // 256 positions at a time
@compute @workgroup_size(256) fn step(@builtin(global_invocation_id) g: vec3u,
@builtin(local_invocation_index) l: u32) {
let me = bodiesIn[g.x];
var acc = vec2f(0);
for (var start = 0u; start < N; start += 256) {
tile[l] = bodiesIn[start + l].xy; // each invocation loads one body
workgroupBarrier();
for (var k = 0u; k < 256; k++) { // 256 interactions from shared memory
let d = tile[k] - me.xy;
let r2 = dot(d, d) + 0.0004; // softening: no infinite force at r = 0
acc += d * (1.0 / f32(N)) * inverseSqrt(r2 * r2 * r2); // G = 1, equal masses
}
workgroupBarrier();
}
let vel = me.zw + acc * dt; // semi-implicit Euler
bodiesOut[g.x] = vec4f(me.xy + vel * dt, vel);
}
@group(0) @binding(0) var<storage> shown: array<vec4f>;
struct Out { @builtin(position) pos: vec4f, @location(0) color: vec4f }
@vertex fn vs(@builtin(vertex_index) v: u32, @builtin(instance_index) i: u32) -> Out {
let b = shown[i];
let q = (vec2f(f32(v & 1), f32(v >> 1)) - 0.5) * 0.012;
let speed = clamp(length(b.zw) / 1.4, 0, 1);
let c = mix(vec3f(0.85, 0.40, 0.15), vec3f(0.10, 0.35, 0.80), speed);
return Out(vec4f(b.x * 0.6 + q.x, b.y + q.y * 1.67, 0, 1), vec4f(c * 0.8, 0.8));
}
@fragment fn fs(in: Out) -> @location(0) vec4f { return in.color; }`;
async function main() {
const adapter = await navigator.gpu?.requestAdapter();
if (!adapter) return showMessage('WebGPU is not available in this browser');
const device = await adapter.requestDevice();
const context = canvas.getContext('webgpu');
const format = navigator.gpu.getPreferredCanvasFormat();
context.configure({ device, format });
const B = GPUBufferUsage;
// A disc with uniform density: enclosed mass grows as r^2, so a circular orbit needs v = sqrt(r).
const init = new Float32Array(N * 4);
for (let i = 0; i < N; i++) {
const r = 0.08 + 0.82 * Math.sqrt(Math.random()), a = Math.random() * 6.2832, v = Math.sqrt(r) * 0.95;
init.set([r * Math.cos(a), r * Math.sin(a), -Math.sin(a) * v, Math.cos(a) * v], i * 4);
}
const buffers = [0, 1].map(() => device.createBuffer({ size: init.byteLength, usage: B.STORAGE | B.COPY_DST }));
device.queue.writeBuffer(buffers[0], 0, init);
const dt = device.createBuffer({ size: 4, usage: B.UNIFORM | B.COPY_DST });
device.queue.writeBuffer(dt, 0, new Float32Array([0.004]));
const module = device.createShaderModule({ code });
const step = device.createComputePipeline({ layout: 'auto', compute: { module, entryPoint: 'step' } });
const blend = { srcFactor: 'one', dstFactor: 'one-minus-src-alpha' };
const render = device.createRenderPipeline({ layout: 'auto', primitive: { topology: 'triangle-strip' },
vertex: { module, entryPoint: 'vs' }, fragment: { module, entryPoint: 'fs', targets: [{ format, blend: { color: blend, alpha: blend } }] } });
// Two bind groups, swapped every step: A -> B, then B -> A.
const stepGroups = [0, 1].map((k) => device.createBindGroup({ layout: step.getBindGroupLayout(0), entries: [
{ binding: 0, resource: { buffer: buffers[k] } }, { binding: 1, resource: { buffer: buffers[1 - k] } }, { binding: 2, resource: { buffer: dt } }] }));
const drawGroups = [0, 1].map((k) => device.createBindGroup({ layout: render.getBindGroupLayout(0), entries: [{ binding: 0, resource: { buffer: buffers[k] } }] }));
ink.font = '12px ui-monospace, monospace'; ink.fillStyle = '#222';
ink.fillText(`${N} bodies: ${(N * N).toLocaleString('en-US')} interactions per step, 256 per tile in workgroup memory`, 10, 18);
ink.font = '11.5px system-ui, sans-serif'; ink.fillStyle = '#555';
ink.fillText('colour: speed (blue fast, orange slow); ping-pong buffers keep every invocation reading a consistent step', 10, 352);
let current = 0;
function frame() {
const encoder = device.createCommandEncoder();
const cp = encoder.beginComputePass();
cp.setPipeline(step);
for (let s = 0; s < 2; s++) { // two steps per frame
cp.setBindGroup(0, stepGroups[current]); cp.dispatchWorkgroups(N / 256);
current = 1 - current; // swap: this step's output is the next step's input
}
cp.end();
const pass = encoder.beginRenderPass({ colorAttachments: [{ view: context.getCurrentTexture().createView(),
clearValue: [0.97, 0.96, 0.93, 1], loadOp: 'clear', storeOp: 'store' }] });
pass.setPipeline(render); pass.setBindGroup(0, drawGroups[current]); pass.draw(4, N);
pass.end();
device.queue.submit([encoder.finish()]);
requestAnimationFrame(frame);
}
requestAnimationFrame(frame);
}
main();
</script>