{"componentChunkName":"component---src-templates-blog-post-js","path":"/blog/2026-08-11-how-flashattention-speeds-up-attention/","result":{"data":{"site":{"siteMetadata":{"title":"M.Hassan Ahmed","author":"Hassan11196"}},"markdownRemark":{"id":"71eff3b6-d7de-5e08-8de6-c2538e5512e7","excerpt":"When I wrote about why GPU memory runs out during long LLM runs, the culprit was the KV cache: memory that grows with every token you keep around. This post…","html":"<p>When I wrote about <a href=\"/blog/2026-07-12-llm-kv-cache-gpu-memory/\">why GPU memory runs out during long LLM runs</a>, the culprit was the KV cache: memory that grows with every token you keep around. This post covers a different memory problem in the same layer. It shows up during the attention computation itself, before anything is cached.</p>\n<p>Here is the puzzle FlashAttention set out to solve. The attention math in a transformer is not that many floating-point operations (FLOPs), and a modern GPU does those operations blindingly fast. Yet for a long time the attention layer was one of the slowest parts of training and inference, and it got quadratically slower as the sequence grew. If the GPU has compute to spare, where does the time go?</p>\n<p>The answer is that attention is <em>memory-bound</em>, not compute-bound: the GPU spends most of its time moving data, not doing arithmetic. This post is for engineers who serve or train transformer models and want to understand why. I walk through:</p>\n<ul>\n<li>the GPU memory hierarchy that makes attention memory-bound,</li>\n<li>the specific thing standard attention does wrong,</li>\n<li>the trick FlashAttention uses to fix it while computing exactly the same result,</li>\n<li>what the later versions changed, and where the abstraction leaks.</li>\n</ul>\n<p>The reference throughout is the original paper, <a href=\"https://arxiv.org/abs/2205.14135\">FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness</a> (Dao et al., 2022).</p>\n<h2>Two kinds of memory on a GPU: HBM and SRAM</h2>\n<p>“Memory-bound” only means something once you know that a GPU has more than one kind of memory, and that they differ in speed by more than an order of magnitude.</p>\n<p><strong>HBM</strong> (high-bandwidth memory) is the large pool you think of as “the GPU’s memory.” On an A100 that is 40 to 80 GB, running at roughly 1.5 to 2.0 TB/s. When you read that a model “doesn’t fit in 40 GB,” this is the 40 GB in question.</p>\n<p><strong>SRAM</strong> is the on-chip memory right next to the compute units. It is tiny: about 192 KB per streaming multiprocessor (one of the GPU’s many independent compute blocks) and around 20 MB across the whole A100. But it runs at roughly 19 TB/s, an order of magnitude faster than HBM. (These figures come from the <a href=\"https://arxiv.org/abs/2205.14135\">FlashAttention paper</a>, Section 2.) SRAM is where computation actually happens fast; HBM is the warehouse you fetch from.</p>\n<p>The gap between those two bandwidths is the whole story. If an algorithm keeps shuttling data to and from HBM, the compute units sit idle waiting for it, no matter how many teraFLOPs the chip is rated for. This is a memory hierarchy in exactly the same sense as the <a href=\"https://en.wikipedia.org/wiki/Memory_hierarchy\">CPU cache hierarchy</a>: small and fast at the top, large and slow below. The same performance lesson applies. You win by keeping working data in the fast tier and touching the slow tier as little as possible.</p>\n<p><img src=\"/b0f33048b0fa307d4b162edbb8b81774/memory-hierarchy.svg\" alt=\"A side-by-side comparison of GPU data movement. On the left, standard attention writes the full N by N attention matrix out to slow HBM and reads it back multiple times for the softmax. On the right, FlashAttention streams tiles of Q, K and V into fast SRAM, computes the whole attention block there in one pass, and writes only the small output back to HBM.\"></p>\n<h2>What standard attention does wrong</h2>\n<p>Attention takes three matrices, the queries <code class=\"language-text\">Q</code>, keys <code class=\"language-text\">K</code>, and values <code class=\"language-text\">V</code>, and computes <code class=\"language-text\">softmax(Q Kᵀ) V</code>. Written out the obvious way, following the <a href=\"https://arxiv.org/abs/1706.03762\">original transformer formulation</a>, it takes three steps:</p>\n<ol>\n<li>Compute the scores <code class=\"language-text\">S = Q Kᵀ</code>. For a sequence of length <code class=\"language-text\">N</code>, this is an <code class=\"language-text\">N × N</code> matrix.</li>\n<li>Apply the softmax row by row to get the attention weights <code class=\"language-text\">P = softmax(S)</code>, still <code class=\"language-text\">N × N</code>.</li>\n<li>Multiply <code class=\"language-text\">O = P V</code> to get the output.</li>\n</ol>\n<p>The problem is steps 1 and 2 together. The <code class=\"language-text\">N × N</code> matrix is written out to HBM, read back to apply the softmax, written out again, and read back once more for the final multiply. At a sequence length of 4,000, that matrix has 16 million entries per attention head, and it sits in HBM for the whole round trip.</p>\n<p>Here is the part that surprises people: the matrix multiplies are not the bottleneck; the softmax is. Softmax, dropout, and masking are elementwise operations. Each does cheap arithmetic per element, so its speed is limited almost entirely by how fast it can read and write its operands. The paper measures this directly: the elementwise steps dominate wall-clock time even though they do a tiny fraction of the FLOPs. The GPU is starved for data, not for compute.</p>\n<p>The memory cost is also quadratic. Materializing (fully storing) the <code class=\"language-text\">N × N</code> matrix means attention uses O(N²) memory, which is why long-context models were historically so expensive. Double the sequence and you quadruple the matrix.</p>\n<h2>The FlashAttention idea: never write the big matrix</h2>\n<p>Once you see it, the insight is almost annoyingly simple: if the <code class=\"language-text\">N × N</code> matrix is what’s killing you, don’t create it. FlashAttention computes exactly the same output without ever materializing the full attention matrix in HBM.</p>\n<p>It does this with <strong>tiling</strong>, which means working on the matrices in small blocks:</p>\n<ol>\n<li>Split <code class=\"language-text\">Q</code>, <code class=\"language-text\">K</code>, and <code class=\"language-text\">V</code> into blocks small enough to fit in SRAM.</li>\n<li>For each block of query rows, walk across the blocks of keys and values.</li>\n<li>For each key/value block, compute the partial attention right there in SRAM and fold it into a running output.</li>\n<li>After the last key/value block, the running output for that query block is finished and correct. Write only that output back to HBM.</li>\n</ol>\n<p>The scores never leave the chip.</p>\n<p><img src=\"/6f0c14b2658e46554a164ba44f21e0e0/tiling-loop.svg\" alt=\"The FlashAttention tiling loop. Q, K and V are split into blocks. For each block of query rows, the algorithm walks the key/value blocks left to right, computes a partial score for each, and folds it into three running numbers: a running max, a running sum, and a running weighted output. After the last block, the running output is the exact softmax result and is written to HBM.\"></p>\n<p>The catch, and the clever part, is the softmax. Softmax normalizes each row by a sum over the <em>entire</em> row, and it first subtracts the row maximum for <a href=\"https://en.wikipedia.org/wiki/Softmax_function#Numerical_stability\">numerical stability</a>. If SRAM only ever holds one block of the row at a time, you don’t know the true maximum or the true sum until you’ve seen every block. So how can a single pass produce an exact result?</p>\n<h2>Online softmax makes the tiling exact</h2>\n<p>The answer is the <strong>online softmax</strong>, a running-statistics trick that predates FlashAttention. It comes from Milakov and Gimelshein’s <a href=\"https://arxiv.org/abs/1805.02867\">Online normalizer calculation for softmax</a> (2018). Instead of waiting for the whole row, you keep three running values for each query row:</p>\n<ul>\n<li><code class=\"language-text\">m</code>, the maximum score seen so far,</li>\n<li><code class=\"language-text\">ℓ</code>, the running sum of <code class=\"language-text\">exp(score − m)</code>, which is the softmax denominator,</li>\n<li><code class=\"language-text\">O</code>, the running weighted sum of the value vectors.</li>\n</ul>\n<p>When a new block contains a score larger than the current <code class=\"language-text\">m</code>, you rescale the values you already have by the correct factor and carry on. Because every partial result is rescaled to a consistent reference whenever the max moves, the value you end up with after the last block is <em>identical</em> to a full-row softmax.</p>\n<p>That property is why you can drop FlashAttention into an existing model with no retraining. It is <strong>exact attention</strong>, not an approximation: same numbers, different memory schedule.</p>\n<h3>Recomputing instead of storing in the backward pass</h3>\n<p>Training needs one more piece. The backward pass normally uses the <code class=\"language-text\">N × N</code> attention matrix that the forward pass produced, and FlashAttention never stored it. So it <em>recomputes</em> the relevant blocks on the fly during the backward pass, from <code class=\"language-text\">Q</code>, <code class=\"language-text\">K</code>, and <code class=\"language-text\">V</code>, which are cheap to reload.</p>\n<p>Spending a bit of extra compute to save a lot of memory traffic is a good deal precisely because the layer is memory-bound. The recomputed matmuls are close to free compared with the HBM reads and writes they save.</p>\n<h2>What this buys you: less memory and less time</h2>\n<p>The gains come in two forms, and they compound.</p>\n<p><strong>Memory drops from quadratic to linear in sequence length.</strong> Because the big matrix is never stored, FlashAttention’s memory use is O(N) instead of O(N²). The paper reports roughly 10× less memory at a sequence length of 2K and 20× at 4K (<a href=\"https://arxiv.org/abs/2205.14135\">FlashAttention paper</a>, Section 4). That is the difference between fitting a long context and hitting a CUDA out-of-memory error. It’s the same failure I traced to the KV cache in the <a href=\"/blog/2026-07-12-llm-kv-cache-gpu-memory/\">earlier post</a>, showing up here for a different reason.</p>\n<p><strong>Wall-clock time drops too</strong>, even though recomputation adds FLOPs, because FLOPs were never the bottleneck. The paper reports:</p>\n<ul>\n<li>training BERT-large (sequence length 512) 15% faster than the MLPerf 1.1 record,</li>\n<li>training GPT-2 (sequence length 1K) 3× faster than the standard HuggingFace and Megatron-LM implementations,</li>\n<li>running Long-Range Arena 2.4× faster.</li>\n</ul>\n<p>It also made sequence lengths practical that weren’t before, reaching better-than-chance results on the 16K-token Path-X task.</p>\n<h2>FlashAttention-2 and -3: same idea, better fit to the hardware</h2>\n<p>The follow-ups keep the core trick and get more out of the GPU.</p>\n<p><a href=\"https://arxiv.org/abs/2307.08691\">FlashAttention-2</a> (2023) reworks how the computation is split across the GPU’s parallel units:</p>\n<ul>\n<li>it parallelizes over the sequence-length dimension,</li>\n<li>it partitions work between warps (groups of threads the GPU schedules together) to cut shared-memory traffic,</li>\n<li>it reduces the non-matmul operations that Tensor Cores (the GPU’s dedicated matrix-multiply units) handle poorly.</li>\n</ul>\n<p>The result is up to about 230 TFLOPs/s on an A100 forward pass, roughly 73% of the theoretical peak, up from the 25% to 40% the first version reached.</p>\n<p><a href=\"https://pytorch.org/blog/flashattention-3/\">FlashAttention-3</a> (2024) targets the Hopper architecture (H100). It overlaps the matmuls with the softmax using the chip’s asynchronous instructions and warp specialization, and it adds FP8 support. Together these give roughly 1.5 to 2 times the throughput of FlashAttention-2 in FP16, and push FP8 attention toward 1.2 PFLOPs/s.</p>\n<p>Each version justified a rewrite because “memory-bound” is a moving target. As you claw back HBM traffic, the next bottleneck appears somewhere else: first warp scheduling, then instruction-level overlap. Each release chases the newly exposed one.</p>\n<h2>Where the abstraction leaks</h2>\n<p>FlashAttention is not a free win everywhere. Three things are worth knowing before you assume it is.</p>\n<p><strong>It’s a kernel, and kernels are hardware-specific.</strong> The speedups above depend on particular GPUs and particular data types. The optimized FP8 path lives on Hopper; older cards fall back to slower paths. If your model runs on mixed hardware, which is normal in a shared HPC cluster, you get a range of speedups, not one number. Check what your actual GPUs support instead of quoting the headline figure.</p>\n<p><strong>It speeds up attention, not the KV cache.</strong> These are two separate memory stories in the same layer. FlashAttention avoids materializing the attention matrix during the computation. It does nothing about the KV cache that grows as you generate tokens. On long autoregressive runs, the KV cache is still what eats your HBM, and the <a href=\"/blog/2026-07-12-llm-kv-cache-gpu-memory/\">fixes for that</a> (shorter context, cache eviction, a quantized cache) are a different toolbox.</p>\n<p><strong>You usually consume it rather than write it.</strong> Unless you do kernel work, FlashAttention reaches you through a framework: PyTorch’s <code class=\"language-text\">scaled_dot_product_attention</code>, an inference server like vLLM, or the <a href=\"https://github.com/Dao-AILab/flash-attention\">flash-attention library</a> directly. The useful skill is knowing whether it’s actually being used. A silent fallback to the naive path (wrong dtype, an unsupported mask, a head dimension the kernel doesn’t handle) quietly costs you the speedup, and profiling is the only way to find out.</p>\n<h2>Why I care about this</h2>\n<p>Most of my work sits a layer or two above the attention kernel. <a href=\"/project/archi/\">Archi</a>, the RAG copilot I built for CMS computing operations at CERN, and tools like <a href=\"/project/llm-dev-mate/\">LLM DevMate</a> live well above it. But the abstraction leaks downward the moment latency or GPU memory becomes the constraint, and on shared HPC hardware it always eventually does. Knowing that attention is memory-bound lets you read a profiler trace and tell “the model is too big” apart from “the kernel isn’t the one you think it is.”</p>\n<p>The lesson generalizes past attention. When something is slower than the arithmetic says it should be, the answer is usually in the memory hierarchy, not the FLOP count. That reframing, from counting operations to counting data movement, is the most useful thing FlashAttention taught me, and it applies far outside the transformer.</p>\n<hr>\n<p><em>Diagrams by M. Hassan Ahmed, released under CC0. No external image was used for this post; the figures are original work by the author.</em></p>","frontmatter":{"title":"How FlashAttention Speeds Up the Attention Layer","date":"2026-08-11T00:00:00.000Z","description":"Attention is memory-bound, not compute-bound. Here's how FlashAttention uses tiling and online softmax to skip the N×N matrix and run exact attention faster.","thumbnail":{"childImageSharp":{"fluid":{"base64":"data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAABQAAAALCAIAAADwazoUAAAACXBIWXMAAAsSAAALEgHS3X78AAAB0klEQVQoz0XQa6+aQBAGYNK0gsIuiwJ7X64iyE2PtrX1cjym3/p/+vs7FpMmTybDsu/sgrU1HeiS4bq7fRwfj6+P++EO9f3tDs79GRbPw+W2v13G623/PplS1szXtq+hsnzc7G7VeCm6n3n3oxqv0Gfb02Z3zdpTvX8v+jNsWPcXsznOsIKU5QVmgsOMxIUf5UFcLmkJjy7RbmBsJGwkHSznvpo4WCx8uSDaeu4gGi0TU+51Psi0U9lgijHk1QzxMCoS06fJIGUbszrmNWU1FQ2XLaWVtXgOk3OiHF/ZWH5B4pPLfGLCVQoy3dWb07g95/kx5i0THRW91KPJDkr1EJYeUctVItla8srw6lf9XYktWmUk0IY3KqqyeMNWZUI3QETrkjcla+AK1hwLQpSKiyEb+rT/Vuz+nH9rPbhhYSPKWZWbXZ7stGhTPQDFt6kZ82QPB0CYO4g7UDH8FWFj8RnxBVFwI3iFiYLzg2UCojgXqqG8AqtV6gfamiMGFli4RE284Fkh7HgxDnORv6Fo7YUlYTXhjU+flYgWRQWE6T9s4YsnLF6NLxw3CmhV9neWHUCc7CEW6gFAH8A3w/gXRCfTOMej07rtxnPMXhD9v9+L/wIza0l/29XgfQAAAABJRU5ErkJggg==","aspectRatio":1.899441340782123,"src":"/static/2d6b9a4a8b053eec65ba5b226b112ebe/40a76/hero.png","srcSet":"/static/2d6b9a4a8b053eec65ba5b226b112ebe/c972b/hero.png 340w,\n/static/2d6b9a4a8b053eec65ba5b226b112ebe/27625/hero.png 680w,\n/static/2d6b9a4a8b053eec65ba5b226b112ebe/40a76/hero.png 1360w,\n/static/2d6b9a4a8b053eec65ba5b226b112ebe/ed396/hero.png 2000w","sizes":"(max-width: 1360px) 100vw, 1360px"}}}}}},"pageContext":{"slug":"/2026-08-11-how-flashattention-speeds-up-attention/","previous":"blog/2026-08-12-matryoshka-embeddings-rag/","next":"blog/2026-08-14-fastapi-token-bucket-rate-limiting/"}},"staticQueryHashes":["32046230"]}