Two-Pass Reduction

A Two-Pass Reduction Across Many Workgroups

Workgroups cannot share memory, so a large reduction runs twice. The first dispatch launches a fixed 256 workgroups, and each invocation first sums a strided slice of the input in a register, a grid-stride loop (for (var i = g.x; i < arrayLength(&input); i += n.x * 256) { sum += input[i]; }, with n the num_workgroups built-in); the workgroup reduces its 256 sums as in Parallel Reduction, and invocation 0 writes a partial. The second dispatch runs the same shader over the 256 partials with one workgroup, in the same pass with another bind group.

On 4,194,304 prices (16 MiB) the pair took 0.13 ms of GPU time against 9-10 ms for a JavaScript loop, and returned 20,942,532 against the exact 20,942,531.76. A loop rounding to f32 after each add, as a naive shader would, drifted to 20,938,192: small values added to a huge total lose their low bits, while the tree adds numbers of similar size. The 16 MiB upload took 9-14 ms, so reduce data that already lives on the GPU.

Summing 1,048,576 prices in two dispatches: 256 workgroup partials (drawn as bars), then one workgroup over the partialsHTMLLive
<!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="340"></canvas>
  <canvas id="labels" width="600" height="340"></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 = 1 << 20;
// One entry point for both passes: a grid-stride loop, then a tree over the workgroup's 256 sums.
const code = /* wgsl */ `
@group(0) @binding(0) var<storage> input: array<f32>;
@group(0) @binding(1) var<storage, read_write> partials: array<f32>;
var<workgroup> sums: array<f32, 256>;
@compute @workgroup_size(256) fn reduce(@builtin(global_invocation_id) g: vec3u,
    @builtin(local_invocation_index) l: u32, @builtin(workgroup_id) w: vec3u, @builtin(num_workgroups) n: vec3u) {
  var sum = 0.0;
  for (var i = g.x; i < arrayLength(&input); i += n.x * 256) { sum += input[i]; }   // grid stride
  sums[l] = sum;
  workgroupBarrier();
  for (var stride = 128u; stride > 0; stride >>= 1) {
    if (l < stride) { sums[l] += sums[l + stride]; }
    workgroupBarrier();
  }
  if (l == 0) { partials[w.x] = sums[0]; }         // one partial per workgroup
}`;
const view = /* wgsl */ `
@group(0) @binding(0) var<storage> partials: array<f32>;
@vertex fn vs(@builtin(vertex_index) v: u32, @builtin(instance_index) i: u32) -> @builtin(position) vec4f {
  let q = vec2f(f32(v & 1), f32(v >> 1));
  let h = (partials[i] / 110000 - 0.8) * 3;          // each partial is about 4,096 prices x $27
  return vec4f(-0.95 + f32(i) * (1.9 / 256) + q.x * 0.006, -0.35 + q.y * clamp(h, 0.02, 1.0), 0, 1);
}
@fragment fn fs() -> @location(0) vec4f { return vec4f(0.08, 0.40, 0.75, 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 prices = new Float32Array(N);
  let exact = 0, naive = 0;                          // double precision, and a float32 running total
  for (let i = 0; i < N; i++) {
    prices[i] = 5 + Math.round(Math.random() * 4500) / 100;
    exact += prices[i]; naive = Math.fround(naive + prices[i]);
  }
  const input = device.createBuffer({ size: N * 4, usage: B.STORAGE | B.COPY_DST });
  device.queue.writeBuffer(input, 0, prices);
  const partials = device.createBuffer({ size: 256 * 4, usage: B.STORAGE });
  const result = device.createBuffer({ size: 4, usage: B.STORAGE | B.COPY_SRC });
  const read = device.createBuffer({ size: 4, usage: B.COPY_DST | B.MAP_READ });
  const pipeline = device.createComputePipeline({ layout: 'auto', compute: { module: device.createShaderModule({ code }) } });
  const group = (a, b) => device.createBindGroup({ layout: pipeline.getBindGroupLayout(0), entries: [
    { binding: 0, resource: { buffer: a } }, { binding: 1, resource: { buffer: b } }] });
  const module = device.createShaderModule({ code: view });
  const render = device.createRenderPipeline({ layout: 'auto', primitive: { topology: 'triangle-strip' },
    vertex: { module }, fragment: { module, targets: [{ format }] } });

  const encoder = device.createCommandEncoder();
  const pass = encoder.beginComputePass();
  pass.setPipeline(pipeline);
  pass.setBindGroup(0, group(input, partials)); pass.dispatchWorkgroups(256);   // pass 1: 256 partials
  pass.setBindGroup(0, group(partials, result)); pass.dispatchWorkgroups(1);    // pass 2: same shader, one workgroup
  pass.end();
  const draw = encoder.beginRenderPass({ colorAttachments: [{ view: context.getCurrentTexture().createView(),
    clearValue: [0.97, 0.96, 0.93, 1], loadOp: 'clear', storeOp: 'store' }] });
  draw.setPipeline(render);
  draw.setBindGroup(0, device.createBindGroup({ layout: render.getBindGroupLayout(0), entries: [{ binding: 0, resource: { buffer: partials } }] }));
  draw.draw(4, 256);
  draw.end();
  encoder.copyBufferToBuffer(result, 0, read, 0, 4);
  device.queue.submit([encoder.finish()]);
  await read.mapAsync(GPUMapMode.READ);
  const gpu = new Float32Array(read.getMappedRange())[0];
  read.unmap();

  const fmt = (x) => x.toLocaleString('en-US', { maximumFractionDigits: 2 });
  ink.font = 'bold 13px system-ui, sans-serif'; ink.fillStyle = '#222';
  ink.fillText('1,048,576 prices: pass 1 writes one partial per workgroup (256 bars)', 12, 24);
  ink.font = '12.5px ui-monospace, monospace';
  ink.fillText(`exact (JavaScript doubles)   ${fmt(exact)}`, 12, 262);
  ink.fillStyle = '#1e6b3a'; ink.fillText(`two-pass tree on the GPU      ${fmt(gpu)}`, 12, 282);
  ink.fillStyle = '#8a2b2b'; ink.fillText(`naive float32 running total  ${fmt(naive)}`, 12, 302);
  ink.font = '11.5px system-ui, sans-serif'; ink.fillStyle = '#444';
  ink.fillText('The tree adds numbers of similar size; the running total loses low bits once it grows large.', 12, 328);
}
main();
</script>