Chained Scans

Chaining Scans Across Workgroups

A workgroup scans only its own block, so a long array takes three dispatches, the scan-then-propagate scheme: (1) each workgroup scans its 256 elements in place, and its last invocation writes the block total to blockSums[w.x]; (2) one workgroup scans blockSums with the same entry point; (3) a third entry point adds the scanned total of all earlier blocks to every element: if (w.x > 0) { data[g.x] += blockSums[w.x - 1]; }.

Built from Workgroup Scan's scan, the three dispatches scanned 65,536 values (256 blocks), matching a JavaScript running sum everywhere; longer arrays recurse. CUDA's single-pass decoupled look-back scan needs a forward-progress guarantee WebGPU does not give, so portable code keeps the multi-dispatch version.

Scanning 65,536 values in three dispatches (block scans, a scan of block sums, add back), with the before and after curves drawnHTMLLive
<!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 = 65536, BLOCKS = N / 256;
const code = /* wgsl */ `
@group(0) @binding(0) var<storage, read_write> data: array<u32>;
@group(0) @binding(1) var<storage, read_write> blockSums: array<u32>;
var<workgroup> s: array<u32, 256>;
// (1) and (2): each workgroup scans its 256 elements in place; its last invocation writes the block total.
@compute @workgroup_size(256) fn scanBlocks(@builtin(global_invocation_id) g: vec3u,
    @builtin(local_invocation_index) l: u32, @builtin(workgroup_id) w: vec3u) {
  s[l] = data[g.x];
  workgroupBarrier();
  for (var d = 1u; d < 256; d <<= 1) {
    let left = select(0u, s[l - d], l >= d);
    workgroupBarrier();
    s[l] += left;
    workgroupBarrier();
  }
  data[g.x] = s[l];
  if (l == 255) { blockSums[w.x] = s[l]; }
}
// (3): add the scanned total of all earlier blocks to every element.
@compute @workgroup_size(256) fn addBack(@builtin(global_invocation_id) g: vec3u, @builtin(workgroup_id) w: vec3u) {
  if (w.x > 0) { data[g.x] += blockSums[w.x - 1]; }
}`;
// Draw 512 samples of the scanned array as a filled curve, and the 256 block totals as bars.
const view = /* wgsl */ `
@group(0) @binding(0) var<storage> data: array<u32>;
@group(0) @binding(1) var<storage> original: 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));
  if (i < 512) {                                              // the running total, sampled
    let value = f32(data[i * 128 + 127]) / f32(data[${N - 1}]);
    return Out(vec4f(-0.95 + f32(i) / 512 * 1.9 + q.x * 0.004, -0.25 + q.y * value * 1.05, 0, 1), vec3f(0.08, 0.40, 0.75));
  }
  let k = i - 512;                                            // the raw input: 256 block averages
  var sum = 0u;
  for (var j = 0u; j < 256; j += 16) { sum += original[k * 256 + j]; }
  return Out(vec4f(-0.95 + f32(k) / 256 * 1.9 + q.x * 0.006, -0.84 + q.y * f32(sum) / 16 / 9 * 0.45, 0, 1), vec3f(0.85, 0.55, 0.15));
}
@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 input = new Uint32Array(N).map((_, i) => Math.round(4 + 4 * Math.sin(i / 3000) + Math.random()));   // daily orders
  const data = device.createBuffer({ size: N * 4, usage: B.STORAGE | B.COPY_DST | B.COPY_SRC });
  const original = device.createBuffer({ size: N * 4, usage: B.STORAGE | B.COPY_DST });
  device.queue.writeBuffer(data, 0, input); device.queue.writeBuffer(original, 0, input);
  const blockSums = device.createBuffer({ size: BLOCKS * 4, usage: B.STORAGE });
  const scratch = device.createBuffer({ size: 4, usage: B.STORAGE });      // the block sums' own block sum
  const read = device.createBuffer({ size: N * 4, usage: B.COPY_DST | B.MAP_READ });
  const module = device.createShaderModule({ code });
  const scan = device.createComputePipeline({ layout: 'auto', compute: { module, entryPoint: 'scanBlocks' } });
  const add = device.createComputePipeline({ layout: 'auto', compute: { module, entryPoint: 'addBack' } });
  const group = (p, a, b) => device.createBindGroup({ layout: p.getBindGroupLayout(0), entries: [
    { binding: 0, resource: { buffer: a } }, { binding: 1, resource: { buffer: b } }] });
  const viewModule = device.createShaderModule({ code: view });
  const render = device.createRenderPipeline({ layout: 'auto', primitive: { topology: 'triangle-strip' },
    vertex: { module: viewModule }, fragment: { module: viewModule, targets: [{ format }] } });

  const encoder = device.createCommandEncoder();
  const pass = encoder.beginComputePass();
  pass.setPipeline(scan); pass.setBindGroup(0, group(scan, data, blockSums)); pass.dispatchWorkgroups(BLOCKS);   // (1)
  pass.setBindGroup(0, group(scan, blockSums, scratch)); pass.dispatchWorkgroups(1);                            // (2) same entry point
  pass.setPipeline(add); pass.setBindGroup(0, group(add, data, blockSums)); pass.dispatchWorkgroups(BLOCKS);     // (3)
  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, group(render, data, original));
  draw.draw(4, 512 + 256);
  draw.end();
  encoder.copyBufferToBuffer(data, 0, read, 0, N * 4);
  device.queue.submit([encoder.finish()]);
  await read.mapAsync(GPUMapMode.READ);
  const out = new Uint32Array(read.getMappedRange());
  let running = 0, mismatches = 0;
  for (let i = 0; i < N; i++) { running += input[i]; if (out[i] !== running) mismatches++; }
  read.unmap();

  ink.font = 'bold 13px system-ui, sans-serif'; ink.fillStyle = '#222';
  ink.fillText(`65,536 values = 256 blocks: scan blocks, scan the block sums, add them back`, 12, 22);
  ink.font = '12px system-ui, sans-serif'; ink.fillStyle = '#1f4f8a';
  ink.fillText(`scanned (blue), final total ${out[N - 1] ?? running}`, 12, 44);
  ink.fillStyle = '#b07a10'; ink.fillText('input, averaged per block (orange)', 12, 230);
  ink.fillStyle = mismatches ? '#8a2b2b' : '#1e6b3a';
  ink.fillText(mismatches ? `${mismatches} values differ from JavaScript` : 'every one of the 65,536 values matches a JavaScript running sum', 12, 332);
}
main();
</script>