Player not loading? Watch on YouTube
This tutorial derives FlashAttention by examining the memory traffic in standard multi-head attention. The speaker assumes familiarity with standard attention and focuses on the forward pass, stopping before the output projection.
The explanation starts with the score and probability matrices that a conventional three-stage implementation writes to GPU memory and reads back. In the speaker's FP16 example, a 32K-token score matrix occupies roughly 2GB per attention head. A simplified account of registers, SRAM and VRAM explains why moving these matrices can be expensive.
A small matrix multiplication example introduces tiling: threads reuse blocks loaded into SRAM instead of repeatedly fetching the same inputs from VRAM. Softmax needs an additional step because each probability depends on the whole row. The speaker derives running maximum and denominator updates, then extends them to maintain a weighted average of value vectors.
The FlashAttention walkthrough combines these updates with tiled computation and follows the paper's pseudocode. The closing analysis gives O(N²D) compute complexity and O(ND) storage, attributing the speed benefit to reduced memory traffic rather than fewer query-key comparisons. The speaker describes the result as the same attention output. GPU hardware details are deliberately simplified; backward-pass derivations and sparse attention extensions remain outside the tutorial's scope.