Spatial Partitioning

Spatial Partitioning to Speed Up N-Body Simulation

Quadratic cost limits brute force to tens of thousands of bodies per frame. Two structures cut it:

Rebuild the grid every frame and size cells to the interaction radius.

A uniform grid built on the GPU every frame (atomic histogram, scan, scatter) so a probe body examines only its 3 x 3 neighbouring cellsHTMLLive
<!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>