Workgroup Scan

A Workgroup-Level Inclusive Scan in Shared Memory

The simplest parallel scan (Hillis and Steele, 1986) runs log2(n) steps; at step d each element adds the value d places to its left. Everyone must read before anyone overwrites, so each step has two barriers:

An inclusive Hillis-Steele scan of 256 daily sales in workgroup memoryJavaScript
// out: 256 x u32
@group(0) @binding(0) var<storage, read_write> out: array<u32>;
var<workgroup> s: array<u32, 256>;
@compute @workgroup_size(256) fn main(@builtin(local_invocation_index) l: u32) {
  s[l] = (l * 7) % 5;                              // copies sold: 0, 2, 4, 1, 3, 0, 2, ...
  workgroupBarrier();
  for (var d = 1u; d < 256; d <<= 1) {             // add the value d places left: 8 steps
    let left = select(0u, s[l - d], l >= d);
    workgroupBarrier();                            // everyone has read before anyone writes
    s[l] += left;
    workgroupBarrier();                            // everyone has written before the next read
  }
  out[l] = s[l];                                   // inclusive: s[l] = sum of inputs 0..l
}
Output
out: 0, 2, 6, 7, 10, 10, 12, 16, 17, 20, 20, 22, ...
gpu: 0.00 ms (median of 5)

The totals match a JavaScript running sum, ending at 510. This scan does about n log2(n) additions where a loop does n - 1; Blelloch's work-efficient scan does about 2n in twice the steps, and with subgroups, subgroupInclusiveAdd() scans 32 values in one call.

A Hillis-Steele inclusive scan of 256 daily sales in workgroup memory, with the sales and their running total drawn from the buffersHTMLLive
<!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="330"></canvas>
  <canvas id="labels" width="600" height="330"></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 code = /* wgsl */ `
@group(0) @binding(0) var<storage, read_write> out: array<u32>;   // 256 sales, then 256 totals
var<workgroup> s: array<u32, 256>;
@compute @workgroup_size(256) fn main(@builtin(local_invocation_index) l: u32) {
  s[l] = (l * 7) % 5;                              // copies sold: 0, 2, 4, 1, 3, 0, 2, ...
  out[l] = s[l];
  workgroupBarrier();
  for (var d = 1u; d < 256; d <<= 1) {             // add the value d places left: 8 steps
    let left = select(0u, s[l - d], l >= d);
    workgroupBarrier();                            // everyone has read before anyone writes
    s[l] += left;
    workgroupBarrier();                            // everyone has written before the next read
  }
  out[256 + l] = s[l];                             // inclusive: s[l] = sum of inputs 0..l
}`;
const view = /* wgsl */ `
@group(0) @binding(0) var<storage> out: 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));
  let k = f32(i % 256);
  let x = -0.95 + k * (1.9 / 256) + q.x * 0.0068;
  if (i < 256) {                                   // daily sales, small bars at the bottom
    return Out(vec4f(x, -0.9 + q.y * f32(out[i]) * 0.05, 0, 1), vec3f(0.85, 0.55, 0.15));
  }
  return Out(vec4f(x, -0.6 + q.y * f32(out[i]) / 520 * 1.1, 0, 1), vec3f(0.08, 0.40, 0.75));   // running total
}
@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 out = device.createBuffer({ size: 2048, usage: B.STORAGE | B.COPY_SRC });
  const read = device.createBuffer({ size: 2048, usage: B.COPY_DST | B.MAP_READ });
  const compute = device.createComputePipeline({ layout: 'auto', compute: { module: device.createShaderModule({ code }) } });
  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 cp = encoder.beginComputePass();
  cp.setPipeline(compute);
  cp.setBindGroup(0, device.createBindGroup({ layout: compute.getBindGroupLayout(0), entries: [{ binding: 0, resource: { buffer: out } }] }));
  cp.dispatchWorkgroups(1);
  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, device.createBindGroup({ layout: render.getBindGroupLayout(0), entries: [{ binding: 0, resource: { buffer: out } }] }));
  pass.draw(4, 512);
  pass.end();
  encoder.copyBufferToBuffer(out, 0, read, 0, 2048);
  device.queue.submit([encoder.finish()]);
  await read.mapAsync(GPUMapMode.READ);
  const d = new Uint32Array(read.getMappedRange().slice(0));
  read.unmap();
  let running = 0, same = true;                    // a JavaScript running sum to compare
  for (let l = 0; l < 256; l++) { running += (l * 7) % 5; same &&= running === d[256 + l]; }

  ink.font = 'bold 13px system-ui, sans-serif'; ink.fillStyle = '#222';
  ink.fillText(`inclusive scan of 256 values in 8 steps, two barriers each: total ${d[511]}`, 12, 24);
  ink.font = '12px system-ui, sans-serif';
  ink.fillStyle = '#1f4f8a'; ink.fillText('blue: running total', 12, 44);
  ink.fillStyle = '#b07a10'; ink.fillText('orange: copies sold per day', 140, 44);
  ink.fillStyle = same ? '#1e6b3a' : '#8a2b2b';
  ink.fillText(same ? 'matches a JavaScript running sum at every index' : 'differs from JavaScript', 320, 44);
  ink.fillStyle = '#555';
  ink.fillText('About n log2(n) additions; Blelloch does ~2n; subgroupInclusiveAdd() scans a subgroup in one call.', 12, 62);
}
main();
</script>