Introduction
In our last lesson, we profiled the standard attention mechanism and proved, with data from ncu, that its performance is crippled by a memory-access bottleneck. The root cause is the need to materialize a massive N×N intermediate matrix, which is too large for fast on-chip SRAM and must be constantly read from and written to the much slower High-Bandwidth Memory (HBM).
Today, we will dissect the solution to this problem: FlashAttention. This lesson directly addresses your goal of understanding what actually happens inside key optimizations. We'll explore the two core algorithmic innovations that allow FlashAttention to compute the exact same attention result without the crippling memory overhead.
Your learning outcome for this lesson is to explain the tiling and online softmax algorithm used by FlashAttention to optimize memory access. By the end, you will understand how these two techniques work together to transform attention from a memory-bound operation into a compute-bound one, unlocking significant performance gains.
The Core Insight: IO-Aware Algorithm Design
The fundamental idea behind FlashAttention is to design an algorithm that is explicitly aware of the GPU's memory hierarchy. As we've established, moving data between slow HBM and fast SRAM is the bottleneck. Therefore, an optimal algorithm must minimize this data transfer.
The author of FlashAttention, Tri Dao, frames this as creating a "Hardware-Aware Algorithm." Let's hear him introduce the concept.
Hardware-aware Algorithms for Sequence Modeling - Tri Dao | Stanford MLSys #87
In this video from the Stanford MLSys Seminar Series, Tri Dao introduces FlashAttention and frames it as an 'IO-aware' algorithm designed to reduce memory reads and writes.
Watch from 10:07 to 12:35. Focus on the core motivation: reducing the amount of reads and writes to GPU memory (HBM) is the key to significant speedups.
The simple but profound goal is to avoid materializing the large intermediate matrices in HBM. To achieve this, FlashAttention uses a classic computer science technique called tiling.
Tiling: Computing Attention in SRAM-Sized Blocks
Tiling, or block processing, involves breaking a large computation into smaller chunks that can fit into a processor's fast local memory (in this case, SRAM).
Instead of computing the full N×N score matrix S = QKᵀ in one go, FlashAttention performs this computation in blocks. It loads a block of queries, and then iterates through blocks of keys and values. For each pair of query and key blocks, it computes a small score sub-matrix entirely within SRAM.
Let's explore this strategy in detail.
Artifact 2.2a: FlashAttention — The Tiling Strategy
This Hugging Face blog post provides an excellent, clear explanation of the tiling strategy with helpful visualizations.
Read the sections 'Tiling: The High-Level Idea', 'Visualizing Standard vs Tiled Computation', and 'The Block Size: Fitting in SRAM'. Pay close attention to: The visualization comparing the full matrix computation of standard attention to the block-by-block processing in FlashAttention. The memory calculation that demonstrates how a typical block of Q, K, V, and the intermediate S block can comfortably fit within the ~192KB of SRAM available on a modern GPU SM.
As the article makes clear, the critical difference is that the intermediate score blocks (S_ij) are computed, used, and then discarded, all without ever leaving the high-speed SRAM. The only writes to HBM are for the final output matrix O.
However, this strategy introduces a significant algorithmic challenge. The softmax function is non-local; to compute the softmax for any single element in a row, you need the statistics (the max and the sum of exponentials) of the entire row. If we only have a small block of scores S_ij in SRAM, how can we possibly compute the correct softmax value?
This is where the second key innovation, the online softmax algorithm, comes into play.
The Online Softmax Algorithm
The "online" or "one-pass" softmax algorithm is a method for computing softmax correctly while only looking at a portion of the input data at a time. It's the mathematical trick that makes tiling feasible for attention.
The algorithm works by maintaining running statistics: a running maximum m and a running normalization sum l. As it iterates through blocks of the input vector (our S_ij score blocks), it updates these statistics.
The most crucial part is the rescaling step. When a new block is processed, if a new, larger maximum value is found, the algorithm cleverly rescales the previous running sum to be correct with respect to this new maximum.
Let's walk through this process by hand to build a solid intuition.
This article provides a step-by-step numerical example of the online softmax algorithm. It first shows the traditional 'safe softmax' (which requires multiple passes) for contrast, and then demonstrates how online softmax achieves the same result in a single pass.
First, quickly review the section 'Conceptual overview: tiled "safe" softmax' to understand the multi-pass problem. Then, read 'Conceptual Overview: Online Softmax' and 'Mathematical conversion from element-wise to tile-wise formula' carefully. Follow the numerical example for processing Tile 1 and Tile 2. The key formula to understand is the update rule for the denominator d_new: d_new ← d_old × exp(m_old - m_new) + sum_of_exponentials_in_current_tile Focus on the exp(m_old - m_new) term—this is the rescaling factor that makes the whole algorithm work.
By applying this update rule iteratively across the blocks of keys for a given block of queries, FlashAttention can compute the correctly normalized attention values without ever needing the full score row at once.
Putting It All Together: The FlashAttention Algorithm
Now let's see how tiling and online softmax are combined in the full FlashAttention forward pass. The algorithm uses a nested loop structure:
- Outer Loop: Iterates over blocks of the Query matrix
Q. - Inner Loop: For each query block, iterates over blocks of the Key
Kand ValueVmatrices.
All computations inside the inner loop happen in fast SRAM.

Let's watch Tri Dao walk through the mathematics of this process, which directly corresponds to the diagram above.
Hardware-aware Algorithms for Sequence Modeling - Tri Dao | Stanford MLSys #87
Here, Tri Dao explains the mathematical details of the softmax rescaling trick. This provides an authoritative look at the core of the algorithm.
Watch from 40:15 to 50:39. This is a dense section, so focus on these key points: Tiling (41:25): The idea of breaking computation into blocks and the challenge posed by the softmax normalization constant. Softmax Rescaling (44:45): This is the core of the 'online softmax' we just studied. Follow his block diagram as he computes a partial output o1 with an incorrect denominator L1, and then later rescales it by L1/L2 to get the correct value before adding the contribution from the second block. Numerical Stability (49:45): He briefly mentions the max subtraction for numerical stability, which is incorporated into the running statistics (m in our notation) of the online softmax algorithm.
A Note on the Backward Pass: Recomputation
While our primary focus is on inference, it's insightful to know how FlashAttention optimizes the backward pass for training. A naive backward pass would require the N×N attention matrix P saved from the forward pass. Storing this would defeat the purpose of FlashAttention's memory savings.
The solution is again guided by the principle that compute is cheaper than memory I/O. Instead of storing P, FlashAttention simply recomputes the necessary blocks of P on-the-fly during the backward pass. This recomputation happens in SRAM and is faster than reading the massive matrix from HBM would have been.
Hardware-aware Algorithms for Sequence Modeling - Tri Dao | Stanford MLSys #87
To complete our understanding, let's briefly hear about the backward pass optimization.
Watch from 51:43 to 53:19. The key takeaway is the counter-intuitive idea that increasing FLOPs (by recomputing) can lead to a faster wall-clock time by drastically reducing memory I/O.
Conclusion
In this lesson, we deconstructed FlashAttention and revealed the two key algorithms that power its efficiency. By understanding these concepts, you've moved beyond simply using an optimized library to understanding the fundamental systems-level principles that make it work.
Key Takeaways:
- IO-Aware Design: FlashAttention's performance comes from minimizing data transfer between slow HBM and fast SRAM.
- Tiling: The computation is broken into small blocks that fit entirely within SRAM, avoiding the need to materialize the full
N×Nattention matrix in HBM. This changes the memory complexity fromO(N^2)toO(N). - Online Softmax: This one-pass algorithm makes tiling possible. By maintaining and rescaling running statistics (
maxandsum), it computes the correct softmax values without needing the entire score row at once. - Compute vs. Memory: FlashAttention is a prime example of trading increased computation (recomputing in the backward pass, more complex logic in the forward pass) for a massive reduction in memory I/O, resulting in a significant overall speedup.
Preview of the Next Lesson:
Now that we have the theory down, it's time for practical validation. In the next lesson, we will benchmark standard attention vs. a FlashAttention implementation. You will use profiling tools to measure the concrete speedup and memory savings across different sequence lengths, quantitatively proving the effectiveness of the algorithms we discussed today.