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:
@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.
<!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>