FlashAttention tutorial: tiling and PyTorch attention

Learn how FlashAttention uses tiles and online softmax through a 28-line Python example, then compare its output with PyTorch.

Player not loading? Watch on YouTube

This tutorial explains FlashAttention through a simplified 28-line Python implementation. It assumes basic Python and high school algebra, then builds ordinary attention on one query, four keys and four values. The example produces 24.7115, the result the tiled implementation must reproduce.

The speaker explains the performance benefit in terms of GPU memory traffic: processing score tiles in fast memory avoids writing and rereading the full attention score table. Every key still contributes to the calculation. An illustrative sequence of 8,192 tokens produces about 67 million scores per head, showing why the intermediate table becomes costly.

The main walkthrough covers online softmax. A running maximum keeps exponentials from overflowing, while a running sum and weighted accumulator carry the result between tiles. When a later tile raises the maximum, the algorithm rescales both totals before adding the new contributions. A final division returns the attention output.

The teaching code uses one attention head and holds all queries in memory. The speaker also discusses causal masking and kernel fusion, then reports a largest difference of about one part in 10 million against PyTorch's reference implementation. The closing example uses scaled_dot_product_attention: PyTorch selects its backend according to hardware, data type and tensor shapes, so FlashAttention use depends on a supported GPU configuration.