Bitonic Sort on the GPU

Quicksort branches on data, which suits a CPU. A sorting network fixes in advance which pairs to compare, so every invocation does the same work. In Batcher's bitonic sort (1968), stage k merges sorted runs of k/2 into runs of k, in alternating directions, through compare-exchanges at distances j = k/2, ..., 1. Each arrow is one invocation's compare-exchange, pointing to where the larger value goes:

A bitonic network sorting eight values in six stages of four compare-exchanges
A bitonic network sorting eight values in six stages of four compare-exchanges

For n elements that is log2(n)(log2(n) + 1)/2 stages. Inside a workgroup a stage is a loop iteration and a barrier (GPU Price Sort); beyond one, each stage is a dispatch whose invocations compare i with i ^ j and swap when (keys[i] > keys[partner]) == ((i & k) == 0), reading k and j from a uniform at a dynamic offset (Dynamic Offsets).

In one pass, 136 dispatches sorted 65,536 random u32 keys in 0.52 ms and 210 sorted 1,048,576 in 11.08 ms, matching Uint32Array.prototype.sort() (2.5 and 46-50 ms). With upload and read-back the million-key sort took 16-19 ms, about three times faster than JavaScript.

Bitonic sort of 512 prices on the GPU, one compare-exchange stage per dispatch, animated stage by stageHTMLLive
<!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 = 512;
// One stage of the network: invocation i compares keys[i] with keys[i ^ j].
const code = /* wgsl */ `
struct Stage { k: u32, j: u32 }
@group(0) @binding(0) var<storage, read_write> keys: array<f32>;
@group(0) @binding(1) var<uniform> stage: Stage;
@compute @workgroup_size(64) fn step(@builtin(global_invocation_id) g: vec3u) {
  let i = g.x;
  let partner = i ^ stage.j;
  if (partner <= i) { return; }                        // each pair handled once
  let ascending = (i & stage.k) == 0;                  // runs alternate direction
  let a = keys[i];  let b = keys[partner];
  if ((a > b) == ascending) { keys[i] = b; keys[partner] = a; }   // compare-exchange
}
@group(0) @binding(0) var<storage> shown: array<f32>;
@vertex fn vs(@builtin(vertex_index) v: u32, @builtin(instance_index) i: u32) -> @builtin(position) vec4f {
  let q = vec2f(f32(v & 1), f32(v >> 1));
  return vec4f(-0.95 + f32(i) / ${N} * 1.9 + q.x * 0.003, -0.8 + q.y * shown[i] / 50 * 1.5, 0, 1);
}
@fragment fn fs(@builtin(position) p: vec4f) -> @location(0) vec4f {
  let t = clamp((280 - p.y) / 220, 0, 1);
  return vec4f(mix(vec3f(0.08, 0.40, 0.75), vec3f(0.85, 0.40, 0.18), t), 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 keys = device.createBuffer({ size: N * 4, usage: B.STORAGE | B.COPY_DST });
  const module = device.createShaderModule({ code });
  const sortPipeline = device.createComputePipeline({ layout: 'auto', compute: { module, entryPoint: 'step' } });
  const render = device.createRenderPipeline({ layout: 'auto', primitive: { topology: 'triangle-strip' },
    vertex: { module, entryPoint: 'vs' }, fragment: { module, entryPoint: 'fs', targets: [{ format }] } });
  const drawGroup = device.createBindGroup({ layout: render.getBindGroupLayout(0), entries: [{ binding: 0, resource: { buffer: keys } }] });

  // All stages in order: k = 2, 4, ..., N; for each, j = k/2, ..., 1. One uniform (and group) per stage.
  const stages = [];
  for (let k = 2; k <= N; k <<= 1) for (let j = k >> 1; j > 0; j >>= 1) stages.push([k, j]);
  const stageGroups = stages.map(([k, j]) => {
    const u = device.createBuffer({ size: 8, usage: B.UNIFORM | B.COPY_DST });
    device.queue.writeBuffer(u, 0, new Uint32Array([k, j]));
    return device.createBindGroup({ layout: sortPipeline.getBindGroupLayout(0), entries: [
      { binding: 0, resource: { buffer: keys } }, { binding: 1, resource: { buffer: u } }] });
  });

  let next = 0, wait = 0;
  const shuffle = () => device.queue.writeBuffer(keys, 0, new Float32Array(N).map(() => 5 + Math.random() * 45));
  shuffle();
  function frame() {
    const encoder = device.createCommandEncoder();
    if (wait > 0) { if (--wait === 0) { shuffle(); next = 0; } }
    else if (next < stages.length) {
      const cp = encoder.beginComputePass();              // one stage: a dispatch of its own
      cp.setPipeline(sortPipeline); cp.setBindGroup(0, stageGroups[next]); cp.dispatchWorkgroups(N / 64); cp.end();
      if (++next === stages.length) wait = 90;            // show the sorted result for a while
    }
    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, drawGroup); pass.draw(4, N);
    pass.end();
    device.queue.submit([encoder.finish()]);
    ink.clearRect(0, 0, 600, 320);
    ink.font = '12.5px ui-monospace, monospace'; ink.fillStyle = '#222';
    const [k, j] = stages[Math.max(next - 1, 0)];
    ink.fillText(`stage ${next} of ${stages.length}: k = ${k}, j = ${j}  (log2(n)(log2(n)+1)/2 = ${stages.length} stages)`, 12, 22);
    ink.font = '11.5px system-ui, sans-serif'; ink.fillStyle = '#555';
    ink.fillText('Every invocation does the same work: compare i with i ^ j, swap when (a > b) == ((i & k) == 0).', 12, 310);
    requestAnimationFrame(frame);
  }
  requestAnimationFrame(frame);
}
main();
</script>