← Back to all stories

FlashAttention and the Art of GPU Mechanics: How Hardware-Aware Algorithms Conquered the Attention Matrix

Picture an architect who needs to examine a colossal blueprint measuring 100 feet wide. The architect's desk, however, is only three feet across. If the architect runs down two flights of stairs to the basement archive to fetch a single three-foot square of the drawing, sketches on it, runs back down to the archive to store it, and repeats this for thousands of squares, the project will take months. But if the architect brings a small box of tiles to the desk and computes running totals on the fly without ever assembling the full 100-foot sheet, the work finishes in an hour. This is the essence of FlashAttention.

The Memory Hierarchy of Modern GPUs

To understand why standard attention was so slow, one must examine the physical anatomy of an AI accelerator. A modern GPU is not a uniform block of memory; it is a steep memory hierarchy:

  • SRAM (Static RAM on chip): Located mere nanometers from the compute cores. It delivers a staggering 19 terabytes per second of bandwidth, but is minuscule in capacity—often only 20 to 50 megabytes per GPU.
  • HBM (High Bandwidth Memory off chip): The main storage pool (80GB to 140GB). It holds model weights and long context buffers, but runs nearly 6x to 10x slower than on-chip SRAM.
[GPU Memory Hierarchy]
┌─────────────────────────────────────────────────────────────┐
│  GPU Compute Units (Tensor Cores)                           │
│  └── SRAM (On-chip Cache: ~20-50MB, ~19 TB/s Bandwidth)     │
└──────────────────────────────▲──────────────────────────────┘
                               │ (Slow Interconnect Bottleneck)
┌──────────────────────────────┴──────────────────────────────┐
│  HBM (High Bandwidth Memory: ~80-140GB, ~3.35 TB/s)         │
└─────────────────────────────────────────────────────────────┘

The Quadratic Memory Catastrophe

In standard self-attention ($\text{Softmax}(QK^T / \sqrt{d})V$), computing how much attention each word pays to every other word requires calculating an $N \times N$ matrix. For an 8,000-token prompt, this intermediate matrix contains 64 million numbers; for a 32,000-token prompt, it explodes to over 1 billion values.

In original PyTorch implementations, the GPU computed these attention scores in SRAM, wrote the full $N \times N$ matrix out to slow global HBM, read it back into SRAM to compute the Softmax normalization, wrote the normalized matrix back out to HBM, and then read it a third time to multiply by the Value vectors. The GPU spent over 80% of its execution cycles simply waiting for data transfers over memory buses.

[Standard Attention: 3 Expensive Roundtrips to Slow HBM]
SRAM ──► Write N×N Matrix ──► [HBM] ──► Read Back ──► Compute Softmax ──► Write Back ──► [HBM] (Slow!)

[FlashAttention: Tiled Fusion with Online Softmax]
Tiled Q, K, V Blocks ──► Loaded ONCE into SRAM ──► [Fused Online Softmax + Output Accumulation]
(The massive N×N intermediate matrix is NEVER written to slow HBM!)

The FlashAttention Breakthrough: Tiling & Online Softmax

Tri Dao and his collaborators approached this problem with a first-principles question: Do we ever actually need to write down the full $N \times N$ attention matrix?

By dividing the Query, Key, and Value matrices into small blocks that fit entirely inside fast SRAM registers, FlashAttention computes attention block-by-block. The major mathematical obstacle was the Softmax function: Softmax requires dividing each value by the sum of exponentials across the entire sequence. How can you normalize a row if you only see one block at a time?

FlashAttention solved this using a technique called Online Softmax. As each new block of Keys and Values is processed in SRAM, the algorithm maintains running normalization factors and dynamically rescales the accumulated output vector. The intermediate attention scores are computed, multiplied, and discarded entirely within fast SRAM registers without a single byte ever touching global HBM.

The Systems Lesson

FlashAttention is an exact mathematical algorithm: it produces the identical numerical output down to floating-point precision as naive attention, yet executes 3x to 5x faster while consuming $O(N)$ memory instead of $O(N^2)$. It stands as proof that in modern AI systems, understanding memory hardware physics is far more powerful than theoretical operation counting.

Reference Paper / Context: FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning (Tri Dao) — Read source ↗
👨‍💻
About the Author

I am Vikram Samal, an AI systems architect exploring how intelligent systems reason, adapt, and act—and how to make them reliable at scale. I connect emerging AI capabilities with the architectural decisions that shape performance, trust, and practical value. Through this blog, I share insights into the ideas and engineering choices shaping AI’s next chapter. As a proud father of two, I believe curiosity, human judgment, and continuous learning are essential in a world being transformed by AI.

Read full bio & connect on LinkedIn →
Previous
← Sparse MoE and Multi-Head Latent Attention: The Story of How Architecture Beat the Memory Wall
Next
The Speculative Gamble: How Guessing the Future Made Autoregressive Models Twice as Fast →