Split-K, streams and overlap
Double buffering hides copies behind compute; split-K makes enough blocks to fill the GPU when the output is small.
Loading the animation…
Concept
A GPU has separate engines for moving data and for computing on it, and a fast kernel keeps them all busy at once. The same pattern appears at two scales:
- Inside a kernel, a block loads the next tile into shared memory while it computes on the current one (double buffering, or a deeper pipeline). On the A100 the loads are
cp.asynccopies that bypass the registers; on the H100 the Tensor Memory Accelerator (TMA) moves whole tiles. - Across kernels, CUDA streams let a host-to-device copy, a kernel and a device-to-host copy of different chunks run at the same time, because the copy engines and the SMs are separate hardware.
The animation is the model's schedule for six tiles on three engines: copy in, compute, copy out. With one buffer, a tile cannot load until the previous tile's compute has freed the buffer, so loading and computing alternate and the compute engine waits (the hatching). With two, tile loads while tile computes: after the first load, the kernel runs at the pace of its slowest engine. A third buffer only helps when the engines' times vary; with the fixed times here it changes nothing. Try the copy-heavy workload: overlap still helps, but now the copy engine sets the pace and the compute engine waits instead.
Loading the animation…
Concept
Overlap needs enough independent work. A GEMM with a small output and a long (a 512 × 512 result from 16,384-long dot products, as in some projection layers at small batch) has only 16 output tiles of 128 × 128: 16 blocks for 108 SMs on an A100, so 85% of the GPU sits idle. Split-K cuts the dimension into slices: each block computes a partial sum of one tile over one slice, and a second pass adds the partial sums. More blocks fill more SMs, at the price of writing and re-reading the partial sums.
Step through the split counts: the time falls almost as while the blocks still fit in one wave, then jumps when they spill into a second wave of mostly idle SMs. On the A100 the best split is 6 (96 blocks, one wave): 39.1 µs instead of 186 µs. On the H100, with 132 SMs, 8 splits (128 blocks) still fit one wave. Libraries choose the split by such a model; Stream-K (Osama et al., 2023) goes further and gives every SM an equal share of the total multiply-add work, so there is no partial wave at all.
Maths
Overlap. Each engine handles one tile at a time, and tile 's load needs a free buffer, that is, tile 's compute finished ( buffers). With , loads and computes alternate, so
With and fixed times, the slowest engine runs without a break after the pipeline fills, so (exactly so when compute is the slowest; the model computes the schedule event by event). For the compute-heavy case, and : and , a speed-up of 1.48.
Split-K. With slices there are blocks and waves; each block does flops. The compute time falls as only while the wave count stays the same; the reduction adds bytes. The best is the largest that fills one wave, unless the reduction outgrows the saving.
Code
The schedule rule, cut from src/lib/gpu/model.ts: tile 's load waits for the previous load and for a free buffer.
let ls = i > 0 ? (loadEnd[i - 1] as number) : 0.0;
if (i >= buffers && (compEnd[i - buffers] as number) > ls)
ls = compEnd[i - buffers] as number;
Double buffering in CUDA, illustrative and not compiled by this site's CI (Ampere's asynchronous copies through the pipeline API):
// a two-stage cuda::pipeline: start tile 0, then each iteration starts
// tile t + 1 before waiting for tile t
pipe.producer_acquire();
cuda::memcpy_async(buf[0], src(0), bytes, pipe);
pipe.producer_commit();
for (int t = 0; t < ntiles; ++t) {
if (t + 1 < ntiles) {
pipe.producer_acquire(); // load tile t + 1 ...
cuda::memcpy_async(buf[(t + 1) % 2], src(t + 1), bytes, pipe);
pipe.producer_commit();
}
pipe.consumer_wait(); // tile t has arrived
compute(buf[t % 2]); // ... while computing tile t
pipe.consumer_release();
}
Split-K on the A100 (model, 512 × 512 × 16384, BF16)
| Splits | Blocks | Waves | Time |
|---|---|---|---|
| 1 | 16 | 1 | 186 µs |
| 4 | 64 | 1 | 51.9 µs |
| 6 | 96 | 1 | 39.1 µs |
| 7 | 112 | 2 | 62.5 µs |
| 12 | 192 | 2 | 47.1 µs |