Parallel Reduction

Parallel Reduction for Sums, Minimums and Maximums

A reduction combines n values with an associative operator (+, min, max). As a tree it takes log2(n) steps: invocation l combines element l with element l + stride, and the stride halves each step, with a barrier between steps. One tree can carry several reductions at once:

A tree reduction computing the sum, minimum and maximum of 256 pricesJavaScript
// out: 3 x u32
@group(0) @binding(0) var<storage, read_write> out: array<u32>;
var<workgroup> total: array<u32, 256>;
var<workgroup> low: array<u32, 256>;
var<workgroup> high: array<u32, 256>;
@compute @workgroup_size(256) fn main(@builtin(local_invocation_index) l: u32) {
  let cents = 500 + (l * 2654435761u) % 4500;      // 256 hashed prices, $5.00 to $49.99
  total[l] = cents; low[l] = cents; high[l] = cents;
  workgroupBarrier();
  for (var stride = 128u; stride > 0; stride >>= 1) {   // 128, 64, ..., 1: 8 steps
    if (l < stride) {
      total[l] += total[l + stride];                   // the same tree for all three
      low[l] = min(low[l], low[l + stride]);
      high[l] = max(high[l], high[l + stride]);
    }
    workgroupBarrier();                                // uniform: every invocation reaches it
  }
  if (l == 0) { out[0] = total[0]; out[1] = low[0]; out[2] = high[0]; }
}
Output
out: 693720, 500, 4999
gpu: 0.00 ms (median of 5)

A JavaScript loop over the same 256 prices agreed. Half the active invocations idle at each level, so fast reductions finish the last five levels with subgroupAdd() (Subgroups). Integer cents keep the sum exact.

A tree reduction of 256 prices computing sum, minimum and maximum at once, with every level of the tree 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="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);
}

// The book's tree, plus a copy of each level's partial sums (256 + 128 + ... + 1 = 511 values).
const code = /* wgsl */ `
@group(0) @binding(0) var<storage, read_write> out: array<u32>;     // sum, min, max, then the levels
var<workgroup> total: array<u32, 256>;
var<workgroup> low: array<u32, 256>;
var<workgroup> high: array<u32, 256>;
@compute @workgroup_size(256) fn main(@builtin(local_invocation_index) l: u32) {
  let cents = 500 + (l * 2654435761u) % 4500;      // 256 hashed prices, $5.00 to $49.99
  total[l] = cents; low[l] = cents; high[l] = cents;
  out[3 + l] = cents;                              // level 0
  workgroupBarrier();
  var offset = 3u + 256u;
  for (var stride = 128u; stride > 0; stride >>= 1) {   // 128, 64, ..., 1: 8 steps
    if (l < stride) {
      total[l] += total[l + stride];               // the same tree for all three
      low[l] = min(low[l], low[l + stride]);
      high[l] = max(high[l], high[l + stride]);
      out[offset + l] = total[l];                  // record this level for the picture
    }
    offset += stride;
    workgroupBarrier();                            // uniform: every invocation reaches it
  }
  if (l == 0) { out[0] = total[0]; out[1] = low[0]; out[2] = high[0]; }
}`;
// Row k shows the 256 >> k partial sums of level k, each scaled by its level's expected size.
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));
  var level = 0u;  var start = 0u;  var count = 256u;
  while (i >= start + count) { start += count; count >>= 1; level++; }
  let k = i - start;
  let value = f32(out[3 + i]) / f32(1u << level) / 5000;          // an average price, 0..1
  let w = 560 / f32(count);
  let px = vec2f(20 + f32(k) * w + q.x * max(w - 1, 0.6), 44 + f32(level) * 32 + (1 - q.y) * 26 * (1 - value));
  let isMin = level == 0u && out[3 + i] == out[1];
  let isMax = level == 0u && out[3 + i] == out[2];
  var color = mix(vec3f(0.55, 0.62, 0.72), vec3f(0.08, 0.40, 0.75), f32(level) / 8);
  if (isMin) { color = vec3f(0.16, 0.56, 0.30); }
  if (isMax) { color = vec3f(0.85, 0.35, 0.25); }
  return Out(vec4f(px.x / 300 - 1, 1 - px.y / 180, 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, SIZE = (3 + 511) * 4;
  const out = device.createBuffer({ size: SIZE, usage: B.STORAGE | B.COPY_SRC });
  const read = device.createBuffer({ size: 12, 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, 511);
  pass.end();
  encoder.copyBufferToBuffer(out, 0, read, 0, 12);
  device.queue.submit([encoder.finish()]);
  await read.mapAsync(GPUMapMode.READ);
  const [sum, low, high] = new Uint32Array(read.getMappedRange());
  read.unmap();
  let jsSum = 0, jsLow = Infinity, jsHigh = 0;              // the same prices in JavaScript
  for (let l = 0; l < 256; l++) { const c = 500 + Number((BigInt(l) * 2654435761n) % 2n ** 32n) % 4500; jsSum += c; jsLow = Math.min(jsLow, c); jsHigh = Math.max(jsHigh, c); }

  ink.font = 'bold 13px system-ui, sans-serif'; ink.fillStyle = '#222';
  ink.fillText('256 prices, then 8 halving steps (each level: pairwise sums, shown as averages)', 20, 26);
  ink.font = '11px system-ui, sans-serif'; ink.fillStyle = '#666';
  for (let level = 0; level <= 8; level++) ink.fillText(`${256 >> level}`, 584, 64 + level * 32);
  ink.font = '12.5px system-ui, sans-serif'; ink.fillStyle = '#222';
  ink.fillText(`sum $${(sum / 100).toFixed(2)}   min $${(low / 100).toFixed(2)} (green)   max $${(high / 100).toFixed(2)} (red)`, 20, 344);
  ink.fillStyle = jsSum === sum && jsLow === low && jsHigh === high ? '#1e6b3a' : '#8a2b2b';
  ink.fillText(jsSum === sum ? 'matches JavaScript' : `JavaScript: ${jsSum}`, 470, 344);
}
main();
</script>