Quadratic cost limits brute force to tens of thousands of bodies per frame. Two structures cut it:
A uniform grid suits short-range forces (collisions, fluids, flocking). It is built from GPGPU Patterns's patterns: an atomic histogram counts bodies per cell, a scan of the counts gives each cell's start, and a scatter writes cell-sorted indices. Each body then reads only its own and neighboring cells, about O(n) work.
A Barnes-Hut tree handles long-range gravity by treating distant groups as one mass, O(n log n) per step. Building it on the GPU takes a sort by Morton code (interleaved coordinate bits) and a bottom-up pass.
Rebuild the grid every frame and size cells to the interaction radius.
<!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 = 2048, G = 16; // bodies; a 16 x 16 grid of cells over the unit square
const code = /* wgsl */ `
const N = ${N}u; const G = ${G}u;
struct Grid { counts: array<atomic<u32>, ${G * G}>, starts: array<u32, ${G * G}>, fill: array<atomic<u32>, ${G * G}> }
@group(0) @binding(0) var<storage, read_write> pos: array<vec2f>;
@group(0) @binding(1) var<storage, read_write> grid: Grid;
@group(0) @binding(2) var<storage, read_write> sorted: array<u32>; // body indices, cell by cell
@group(0) @binding(3) var<storage, read_write> examined: array<u32>; // 1 if the probe looked at this body
@group(0) @binding(4) var<uniform> time: f32;
fn cellOf(p: vec2f) -> vec2u { return min(vec2u(p * f32(G)), vec2u(G - 1)); }
fn cellIndex(c: vec2u) -> u32 { return c.y * G + c.x; }
fn hash(i: u32) -> f32 { var h = i * 747796405u + 2891336453u; h = ((h >> ((h >> 28u) + 4u)) ^ h) * 277803737u; return f32((h >> 22u) ^ h) / 4294967296.0; }
@compute @workgroup_size(64) fn drift(@builtin(global_invocation_id) g: vec3u) { // drift, and clear the grid
if (g.x < G * G) { atomicStore(&grid.counts[g.x], 0); atomicStore(&grid.fill[g.x], 0); }
if (g.x >= N) { return; }
let a = time * (0.2 + hash(g.x)) + hash(g.x + 9) * 6.28;
pos[g.x] = fract(pos[g.x] + vec2f(cos(a), sin(a)) * 0.0015);
examined[g.x] = 0;
}
@compute @workgroup_size(64) fn count(@builtin(global_invocation_id) g: vec3u) { // 1. atomic histogram
if (g.x >= N) { return; }
atomicAdd(&grid.counts[cellIndex(cellOf(pos[g.x]))], 1);
}
var<workgroup> s: array<u32, 256>;
@compute @workgroup_size(256) fn scan(@builtin(local_invocation_index) l: u32) { // 2. exclusive scan of the counts
let own = atomicLoad(&grid.counts[l]);
s[l] = own;
workgroupBarrier();
for (var d = 1u; d < 256; d <<= 1) {
let left = select(0u, s[l - d], l >= d);
workgroupBarrier();
s[l] += left;
workgroupBarrier();
}
grid.starts[l] = s[l] - own; // where this cell's list begins
}
@compute @workgroup_size(64) fn scatter(@builtin(global_invocation_id) g: vec3u) { // 3. cell-sorted indices
if (g.x >= N) { return; }
let c = cellIndex(cellOf(pos[g.x]));
sorted[grid.starts[c] + atomicAdd(&grid.fill[c], 1)] = g.x;
}
@compute @workgroup_size(1) fn probe() { // body 0 visits only its neighbours
let c = vec2i(cellOf(pos[0]));
for (var dy = -1; dy <= 1; dy++) {
for (var dx = -1; dx <= 1; dx++) {
let n = c + vec2i(dx, dy);
if (any(n < vec2i(0)) || any(n >= vec2i(i32(G)))) { continue; }
let cell = cellIndex(vec2u(n));
for (var k = 0u; k < atomicLoad(&grid.counts[cell]); k++) { examined[sorted[grid.starts[cell] + k]] = 1; }
}
}
}
@group(0) @binding(0) var<storage> drawnPos: array<vec2f>;
@group(0) @binding(3) var<storage> drawnExamined: array<u32>;
struct Out { @builtin(position) pos: vec4f, @location(0) color: vec3f }
@vertex fn vs(@builtin(vertex_index) v: u32, @builtin(instance_index) i: u32) -> Out {
let q = (vec2f(f32(v & 1), f32(v >> 1)) - 0.5) * select(0.012, 0.03, i == 0u);
let p = drawnPos[i];
var color = vec3f(0.65, 0.68, 0.74);
if (drawnExamined[i] == 1u) { color = vec3f(0.16, 0.56, 0.30); }
if (i == 0u) { color = vec3f(0.85, 0.30, 0.20); }
return Out(vec4f((p.x * 2 - 1) * 0.6 - 0.35 + q.x * 0.6, (p.y * 2 - 1) * 0.9 + q.y, 0, 1), color);
}
@fragment fn fs(in: Out) -> @location(0) vec4f { return vec4f(in.color, 1); }`;
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;
const pos = device.createBuffer({ size: N * 8, usage: B.STORAGE | B.COPY_DST });
device.queue.writeBuffer(pos, 0, new Float32Array(N * 2).map(() => Math.random()));
const grid = device.createBuffer({ size: G * G * 12, usage: B.STORAGE });
const sorted = device.createBuffer({ size: N * 4, usage: B.STORAGE });
const examined = device.createBuffer({ size: N * 4, usage: B.STORAGE | B.COPY_SRC });
const time = device.createBuffer({ size: 4, usage: B.UNIFORM | B.COPY_DST });
const readback = [0, 1].map(() => device.createBuffer({ size: N * 4, usage: B.COPY_DST | B.MAP_READ }));
const module = device.createShaderModule({ code });
const layout = device.createBindGroupLayout({ entries: [0, 1, 2, 3].map((binding) =>
({ binding, visibility: GPUShaderStage.COMPUTE, buffer: { type: 'storage' } })).concat([{ binding: 4, visibility: GPUShaderStage.COMPUTE, buffer: { type: 'uniform' } }]) });
const pl = device.createPipelineLayout({ bindGroupLayouts: [layout] });
const passes = ['drift', 'count', 'scan', 'scatter', 'probe'].map((entryPoint) => device.createComputePipeline({ layout: pl, compute: { module, entryPoint } }));
const group = device.createBindGroup({ layout, entries: [pos, grid, sorted, examined, time].map((buffer, binding) => ({ binding, resource: { buffer } })) });
const render = device.createRenderPipeline({ layout: 'auto', primitive: { topology: 'triangle-strip' },
vertex: { module, entryPoint: 'vs' }, fragment: { module, entryPoint: 'fs', targets: [{ format }] } });
const renderGroup = device.createBindGroup({ layout: render.getBindGroupLayout(0), entries: [
{ binding: 0, resource: { buffer: pos } }, { binding: 3, resource: { buffer: examined } }] });
const dispatch = [Math.ceil(N / 64), Math.ceil(N / 64), 1, Math.ceil(N / 64), 1];
let checked = '?';
function frame(now) {
device.queue.writeBuffer(time, 0, new Float32Array([now / 1000]));
const encoder = device.createCommandEncoder();
const cp = encoder.beginComputePass();
cp.setBindGroup(0, group);
passes.forEach((p, k) => { cp.setPipeline(p); cp.dispatchWorkgroups(dispatch[k]); }); // five dispatches: each sees the last one's writes
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, renderGroup); pass.draw(4, N);
pass.end();
const staging = readback.find((b) => b.mapState === 'unmapped');
if (staging) encoder.copyBufferToBuffer(examined, 0, staging, 0, N * 4);
device.queue.submit([encoder.finish()]);
if (staging) staging.mapAsync(GPUMapMode.READ).then(() => {
checked = new Uint32Array(staging.getMappedRange()).reduce((a, b) => a + b, 0); staging.unmap();
});
ink.clearRect(0, 0, 600, 360);
ink.strokeStyle = 'rgba(0,0,0,0.12)';
for (let k = 0; k <= G; k++) { // the grid lines, over the square plot
const x = 15 + k * 360 / G, y = 18 + k * 324 / G;
ink.beginPath(); ink.moveTo(x, 18); ink.lineTo(x, 342); ink.stroke();
ink.beginPath(); ink.moveTo(15, y); ink.lineTo(375, y); ink.stroke();
}
ink.font = '12.5px system-ui, sans-serif'; ink.fillStyle = '#222';
['Each frame, on the GPU:', '1. atomic histogram of bodies per cell', '2. scan of the counts: cell starts', '3. scatter: indices sorted by cell', '',
'The red probe body reads only its', '3 x 3 cells (green):', `${checked} of ${N} bodies examined`, 'instead of all of them.', '',
'Cells sized to the interaction radius', 'turn O(n^2) into about O(n).'].forEach((l, i) => ink.fillText(l, 392, 40 + i * 21));
requestAnimationFrame(frame);
}
requestAnimationFrame(frame);
}
main();
</script>