Mapping Problems

Mapping Problems onto Workgroups and Dispatches

Ask what each invocation produces and what it reads. If each output depends on one input, the job is a map: one invocation per element, no communication. Otherwise invocations cooperate, cheaply within a workgroup and only through a new dispatch across workgroups, so the dispatch count follows how often data crosses workgroups: one for a map, a histogram or a blur, two for a large reduction, three for a scan, one per stage for a large sort. Give each invocation enough work to hide memory latency, keep intermediate buffers on the GPU, and read back only the answer.

How often data crosses workgroups decides the dispatch count: map, histogram, blur, reduction, scan and sort side by sideHTMLLive
<!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="392"></canvas>
  <canvas id="labels" width="600" height="392"></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);
}
function label(text, x, y, size = 12, color = '#2b2b2b', align = 'left', weight = '') {
  ink.font = `${weight} ${size}px system-ui, sans-serif`; ink.fillStyle = color; ink.textAlign = align; ink.fillText(text, x, y);
}

// pattern, dispatches, what each invocation reads, how outputs map from 8 inputs
const patterns = [
  ['Map', '1', 'one input per output', (i) => [i]],
  ['Histogram', '1', 'atomics into bins', (i) => [i % 3]],
  ['Blur (stencil)', '1', 'a neighbourhood', (i) => [Math.max(i - 1, 0), i, Math.min(i + 1, 7)]],
  ['Reduction (large)', '2', 'partials, then partials of partials', () => [0]],
  ['Scan (large)', '3', 'block scan, block sums, add back', (i) => [i]],
  ['Sort (large)', 'per stage', 'compare-exchange pairs', (i) => [i ^ 1]]];
const BLUE = [0.08, 0.40, 0.75], PALE = [0.80, 0.86, 0.94], LINE = [0.55, 0.55, 0.55];

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 boxes = [];
  label('pattern', 12, 22, 12, '#666'); label('inputs -> outputs', 190, 22, 12, '#666'); label('dispatches', 520, 22, 12, '#666');
  patterns.forEach(([name, dispatches, reads, map], row) => {
    const y = 36 + row * 56;
    boxes.push([8, y, 584, 48, 1, 1, 1, 8]);
    label(name, 18, y + 21, 13, '#222', 'left', 'bold'); label(reads, 18, y + 38, 10.5, '#666');
    const outputs = name.startsWith('Histogram') ? 3 : name.startsWith('Reduction') ? 1 : 8;
    for (let i = 0; i < 8; i++) boxes.push([190 + i * 18, y + 6, 14, 12, ...PALE, 2]);          // inputs
    for (let o = 0; o < outputs; o++) boxes.push([190 + o * 18, y + 30, 14, 12, ...BLUE, 2]);  // outputs
    for (let i = 0; i < 8; i++) for (const o of map(i)) {
      if (o >= outputs) continue;
      const x0 = 197 + i * 18, x1 = 197 + o * 18;            // a slanted link drawn as a dotted run of tiny boxes
      for (let t = 0; t <= 1; t += 0.2) boxes.push([x0 + (x1 - x0) * t - 1, y + 18 + 12 * t - 1, 2, 2, ...LINE, 1]);
    }
    boxes.push([520, y + 12, dispatches.length > 2 ? 64 : 28, 24, ...(dispatches === '1' ? [0.16, 0.56, 0.30] : [0.85, 0.55, 0.15]), 12]);
    label(dispatches, 534 + (dispatches.length > 2 ? 18 : 0), y + 29, 12.5, '#fff', 'center', 'bold');
  });
  label('Cooperate cheaply inside a workgroup; cross workgroups only with a new dispatch. Keep data on the GPU.', 12, 386, 11.5, '#444');

  const module = device.createShaderModule({ code: `
    struct Box { rect: vec4f, style: vec4f }          // style: r, g, b, corner radius
    @group(0) @binding(0) var<storage> boxes: array<Box>;
    struct Out { @builtin(position) pos: vec4f, @location(0) local: vec2f,
                 @location(1) @interpolate(flat) i: u32 }
    @vertex fn vs(@builtin(vertex_index) v: u32, @builtin(instance_index) i: u32) -> Out {
      let corner = vec2f(f32(v & 1), f32(v >> 1));   // triangle-strip corners
      let r = boxes[i].rect;
      let px = r.xy + corner * r.zw;
      return Out(vec4f(px.x / 300 - 1, 1 - px.y / 196, 0, 1), corner * r.zw, i);
    }
    @fragment fn fs(in: Out) -> @location(0) vec4f {
      let b = boxes[in.i];
      let half = b.rect.zw / 2;
      let q = abs(in.local - half) - half + b.style.w; // signed distance to a rounded box
      let d = length(max(q, vec2f(0))) + min(max(q.x, q.y), 0) - b.style.w;
      let a = clamp(0.5 - d, 0, 1);
      return vec4f(b.style.rgb * a, a);               // premultiplied alpha
    }` });
  const blend = { srcFactor: 'one', dstFactor: 'one-minus-src-alpha' };
  const pipeline = device.createRenderPipeline({ layout: 'auto',
    vertex: { module }, primitive: { topology: 'triangle-strip' },
    fragment: { module, targets: [{ format, blend: { color: blend, alpha: blend } }] } });
  const data = new Float32Array(boxes.flat());
  const buffer = device.createBuffer({ size: data.byteLength, usage: GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_DST });
  device.queue.writeBuffer(buffer, 0, data);
  const encoder = device.createCommandEncoder();
  const pass = encoder.beginRenderPass({ colorAttachments: [{ view: context.getCurrentTexture().createView(),
    clearValue: [0.93, 0.91, 0.87, 1], loadOp: 'clear', storeOp: 'store' }] });
  pass.setPipeline(pipeline);
  pass.setBindGroup(0, device.createBindGroup({ layout: pipeline.getBindGroupLayout(0), entries: [{ binding: 0, resource: { buffer } }] }));
  pass.draw(4, boxes.length);
  pass.end();
  device.queue.submit([encoder.finish()]);
}
main();
</script>