Block Pointers and Multi-Dim Support
PremiumScale from single sequence to Batch/Head parallelism and simplify pointer math with block pointers.
Get code accessIn the previous chapters we built a functional, autotuned Flash Attention kernel. Our parallelism was along the sequence dimension: each program handles one Q block, iterating over all K/V blocks.
But real Transformer inputs are (Batch, Head, SeqLen, Dim). Batch and Head are independent, so we should parallelize them too.
We solve two problems:
- Multi-dimensional parallelism: scale from SeqLen to Batch × Head × SeqLen
- Pointer management: use block pointers to simplify 4D addressing
From Single Sequence to Batch/Head Parallelism
4D Tensor Memory Layout
When input is (B, H, N, D), pointer math gets complex. The GPU memory is still 1D contiguous (see Tensor Layout).
Example: (2, 4, 8, 64) (2 batches, 4 heads, seq length 8, dim 64):
Logical view: Q[batch, head, seq, dim] → Q[2, 4, 8, 64]
Physical storage: 1D array, total 2 × 4 × 8 × 64 = 4096 elementsStrides tell how many elements to skip per dimension:
Log in to continue reading
This is premium content. Please log in to access the full article.
CookLLM Docs