Workgroup Memory

Shared Memory with var<workgroup>

A var<workgroup> variable lives in on-chip memory, one copy per workgroup (Local Address Spaces), 16,384 bytes by default. It costs a few cycles to read against hundreds for an uncached storage read, so the classic use is tiling: each invocation loads one element of a block, the workgroup waits, then all of them read the whole block from shared memory. Here each of 16,384 books counts the cheaper books, which is its price rank:

Ranking prices through 256-price tiles in workgroup memoryJavaScript
@group(0) @binding(0) var<storage> prices: array<f32>;
@group(0) @binding(1) var<storage, read_write> rank: array<u32>;
var<workgroup> tile: array<f32, 256>;                // 1 KiB of the 16 KiB allowed
@compute @workgroup_size(256) fn main(@builtin(global_invocation_id) g: vec3u,
                                      @builtin(local_invocation_index) l: u32) {
  let mine = prices[g.x];
  var cheaper = 0u;
  for (var start = 0u; start < arrayLength(&prices); start += 256) {   // 256 | length
    tile[l] = prices[start + l];                     // each invocation loads one price
    workgroupBarrier();                              // the tile is complete
    for (var k = 0u; k < 256; k++) {
      let other = tile[k];                           // 256 reads served on-chip
      cheaper += u32(other < mine || (other == mine && start + k < g.x));
    }
    workgroupBarrier();                              // all done before the next load
  }
  rank[g.x] = cheaper;                               // this book's position by price
}

Both this and a version reading prices[k] from storage matched a JavaScript sort for all 16,384 books; the tiled one took 1.31 ms against 2.36 ms (medians of five, twice). The gain is modest because all invocations read the same address, which the cache serves well; tiling pays more when neighbors read overlapping data, as in Blur and Edge Detection's blur. The second barrier stops a fast invocation from overwriting tile while a slow one still reads it.

Ranking 4,096 prices through 256-price tiles in var<workgroup> memory, drawn before and after sorting by rankHTMLLive
<!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="320"></canvas>
  <canvas id="labels" width="600" height="320"></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 = 4096;
const code = /* wgsl */ `
@group(0) @binding(0) var<storage> prices: array<f32>;
@group(0) @binding(1) var<storage, read_write> rank: array<u32>;
var<workgroup> tile: array<f32, 256>;                // 1 KiB of the 16 KiB allowed
@compute @workgroup_size(256) fn main(@builtin(global_invocation_id) g: vec3u,
                                      @builtin(local_invocation_index) l: u32) {
  let mine = prices[g.x];
  var cheaper = 0u;
  for (var start = 0u; start < arrayLength(&prices); start += 256) {   // 256 divides the length
    tile[l] = prices[start + l];                     // each invocation loads one price
    workgroupBarrier();                              // the tile is complete
    for (var k = 0u; k < 256; k++) {
      let other = tile[k];                           // 256 reads served on-chip
      cheaper += u32(other < mine || (other == mine && start + k < g.x));
    }
    workgroupBarrier();                              // all done before the next load
  }
  rank[g.x] = cheaper;                               // this book's position by price
}`;
// Each book is a dot: x from its index (left panel) or its rank (right panel), y from its price.
const view = /* wgsl */ `
@group(0) @binding(0) var<storage> prices: array<f32>;
@group(0) @binding(1) var<storage> rank: 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)) - 0.5;
  let book = i % ${N};  let ranked = i >= ${N};
  let x = select(f32(book), f32(rank[book]), ranked) / ${N};
  let y = (prices[book] - 5) / 45;
  let p = vec2f(select(-0.95, 0.05, ranked) + x * 0.9, -0.66 + y * 1.4) + q * vec2f(0.006, 0.012);
  return Out(vec4f(p, 0, 1), mix(vec3f(0.08, 0.40, 0.75), vec3f(0.85, 0.40, 0.18), y));
}
@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 priceList = new Float32Array(N).map(() => 5 + 45 * Math.random() ** 1.5);
  const prices = device.createBuffer({ size: N * 4, usage: B.STORAGE | B.COPY_DST });
  device.queue.writeBuffer(prices, 0, priceList);
  const rank = device.createBuffer({ size: N * 4, usage: B.STORAGE | B.COPY_SRC });
  const read = device.createBuffer({ size: N * 4, 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 both = (p) => device.createBindGroup({ layout: p.getBindGroupLayout(0), entries: [
    { binding: 0, resource: { buffer: prices } }, { binding: 1, resource: { buffer: rank } }] });

  const encoder = device.createCommandEncoder();
  const cp = encoder.beginComputePass();
  cp.setPipeline(compute); cp.setBindGroup(0, both(compute)); cp.dispatchWorkgroups(N / 256); 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, both(render)); pass.draw(4, 2 * N);
  pass.end();
  encoder.copyBufferToBuffer(rank, 0, read, 0, N * 4);
  device.queue.submit([encoder.finish()]);
  await read.mapAsync(GPUMapMode.READ);
  const ranks = new Uint32Array(read.getMappedRange());
  const sorted = [...priceList].map((p, i) => [p, i]).sort((a, b) => a[0] - b[0] || a[1] - b[1]);
  const agree = sorted.every(([, i], r) => ranks[i] === r);   // compare with a JavaScript sort
  read.unmap();

  ink.font = '12.5px system-ui, sans-serif'; ink.fillStyle = '#222';
  ink.fillText('prices in catalog order', 20, 24);
  ink.fillText('the same prices placed at rank[i]', 320, 24);
  ink.fillStyle = agree ? '#1e6b3a' : '#8a2b2b';
  ink.fillText(`${N.toLocaleString('en-US')} ranks ${agree ? 'match' : 'do not match'} a JavaScript sort. Each workgroup loads 256 prices at a time`, 20, 292);
  ink.fillText('into var<workgroup> tile, then every invocation reads the whole tile on-chip.', 20, 310);
}
main();
</script>