FlashAttention: Part 2
This is Part 2 of a two-part FlashAttention series. If you haven’t read Part 1 yet, start there. Part 1 covers the fundamental memory problem, the mathematical trick, and the FA1 implementation.
From FA1 to FA4: The Evolution of FlashAttention
The reference handbook for this blog post was inspired by this tweet and is titled: Understanding FlashAttention: How IO-Aware Attention Makes Transformers Faster Without Approximating Attention.
In Part 1, we covered the fundamental problem FlashAttention-1 (FA1) addresses, the mathematical tricks it uses, the GPU implementation of those tricks, and the architectural compatibility (MHA/MQA/GQA) of FlashAttention.
In this blog post, we continue the story through the evolution of the FlashAttention family (FA2, FA3, FA4), its relationship to other efficient attention techniques (PagedAttention, sparse, linear), the training vs. inference distinction, current PyTorch integration, and the practical mental model to walk away with.
Table of Contents
- 1.1 FlashAttention-2: what changed
- 1.2 FA2 parallelism across sequence tiles
- 1.3 FA2 work partitioning and non-matmul FLOPs
- 1.4 FlashAttention-3: the Hopper generation
- 1.5 FA3 asynchrony: overlap data movement, GEMM, and softmax
- 1.6 FA3 FP8: performance without pretending precision is free
- 1.7 FlashAttention-4: the Blackwell generation
- 1.8 FA4 and asymmetric hardware scaling
- 1.9 FA4 implementation and current status
- 1.10 FlashAttention-1 through FA-4 compared
2. Using FlashAttention in Frameworks
3. FlashAttention vs. Other Techniques
5. Practical Engineering for FlashAttention
Appendix
1. FlashAttention Evolution
1.1 FlashAttention-2: What Changed
FA2 didn’t change the goal. It still computes exact attention without materializing quadratic intermediates. However, it made much better use of the GPU.1 The FA2 paper identifies three main changes:
- Reduce the number of non-matmul FLOPs.
- Parallelize attention across sequence tiles, even for a single head, to improve occupancy.
- Repartition work among warps within a thread block to reduce shared-memory communication.
The roughly 2x speedup of FA2 over FA1 on A100 in the original paper’s experiments is real,1 but it’s a benchmark- and GPU-specific number that depends on dtype, shapes, causal mode, and software versions. It is not a universal speed guarantee.
The question FA2 is really asking: “How can we utilize the GPU more effectively?”
- FA1’s breakthrough was IO-aware tiled exact attention.
- FA2’s breakthrough was better parallelism and work partitioning around the same high-level algorithmic idea.
Every subsequent FlashAttention evolution since has followed the same pattern: keep the semantics, and adapt the kernel pipeline to the resources and bottlenecks of newer GPU architectures.
1.2 FA2 Parallelism Across Sequence Tiles
The practical scenario where this matters:
- batch size is small,
- number of heads is small,
- sequence length is huge.
In that case, FA1 might not expose enough independent work. A GPU has many processing clusters known as streaming multiprocessors (SMs) that need enough independent work to be kept busy. If we only parallelize over batch times heads, some of those units sit idle. FA2 adds parallelism along the sequence dimension, so different query-row tiles can be processed independently. That gives the GPU more work to schedule, since more resident blocks means more available parallel work.
Let’s consider a simplified mental picture:
Q rows [ block 0 | block 1 | block 2 | block 3 | block 4 | block 5 ]
↓ ↓ ↓
CTA CTA CTA
- Query rows get split into blocks (block 0, block 1, …, block 5), and each block gets assigned to a CTA.
- A CTA, Cooperative Thread Array also known as a thread block, is a logical grouping of threads that execute together on a single SM.
More parallel blocks don’t automatically mean better performance. We still have constraints: occupancy, registers, shared memory, tile size, data reuse. Kernel design is an optimization problem over hardware resources, not just a request for maximum thread count.2
Long sequence length itself can become a source of parallelism.
1.3 FA2 Work Partitioning and Non-Matmul FLOPs
FA2 also changes how work is divided among warps in a thread block. In FA1, warp partitioning could require partial results to be written to shared memory and then combined.1 That creates communication overhead.
Picture it: a warp writes a partial result to shared memory, another warp reads it, combines it, and writes again. Every round trip through shared memory costs cycles.
[warp 0] --> shared memory --> [warp 1] --> shared memory --> ...
FA2 repartitions the work so warps cooperate in a way that cuts down those shared-memory reads and writes, thereby reducing communication. It also reduces non-matmul operations via algebraic manipulation around softmax rescaling and backward computations.1 This matters because:
- A FLOP isn’t a universal unit of performance across all operation types.
- Matrix multiply-accumulate (MMA) has significantly higher throughput than exponentials, scalar reductions, conversions, and shared-memory synchronization.
- A tensor-core matmul and an exponential aren’t equivalent from a hardware-throughput perspective.
We can track the bottlenecks addressed by each FlashAttention generation so far:
- FA1 attacked HBM traffic.
- FA2 exposed the following as the next important bottlenecks: occupancy, warp communication, non-matmul operations.
FA2 also simplifies what gets saved for backward. Instead of keeping probabilities around, it stores a single row-wise log-sum-exp value per query row and recomputes the tiles, keeping the same linear-memory footprint.
A useful optimization principle from the FlashAttention evolution so far: after fixing one bottleneck, another bottleneck emerges.
1.4 FlashAttention-3: The Hopper Generation
Why does FA3 exist at all? GPU hardware changed. Hopper introduced capabilities that changed the best way to schedule attention, and FA2 didn’t fully exploit what the H100 could do. The FA3 paper measured FA2 at only around 35% utilization on H100, which leaves a lot of headroom.
FA3 was designed around NVIDIA Hopper GPUs and attacks that utilization gap with four broad ideas:
- Asynchronous computation and data movement.
- Warp specialization.
- Interleaving GEMM (General Matrix Multiplication) and softmax work.
- FP8 support.
The important point: The mathematical algorithm didn’t suddenly change. The execution workflow did, to exploit Hopper’s hardware architecture.
On H100, the NeurIPS 2024 publication reports a 1.5 to 2.0x speedup over FA2 in its benchmark suite. BF16 throughput reaches up to 840 TFLOPs/s (85% utilization), and FP8 reaches 1.3 PFLOPs/s.3 (The earlier arXiv preprint reported 740 TFLOPs/s and close to 1.2 PFLOPs/s; the figures here are from the final NeurIPS version.) As before, these are hardware- and benchmark-specific empirical results, not universal across all models.
FA3 is not “a more approximate FlashAttention.” The FP16/BF16 path preserves the exact-attention algorithmic goal. The FP8 path intentionally introduces lower-precision arithmetic and therefore needs its own numerical-accuracy discussion.
1.5 FA3 Asynchrony: Overlap Data Movement, GEMM, and Softmax
Picture a naive pipeline:
load tile → GEMM → softmax → PV GEMM → load next tile
Each stage waits for the previous one. That leaves periods where some hardware resources sit idle while others work. Tensor cores go quiet during softmax. Memory buses go quiet during GEMM. FA3 builds a more overlapping pipeline:
Load K/V tile 1
Compute tile 1
Load K/V tile 2
Compute tile 2
Load K/V tile 3
Compute tile 3
More explicitly:
Load K/V tile 1
QK^T GEMM
softmax/update
PV GEMM
Load K/V tile 2
QK^T GEMM
softmax/update
PV GEMM
The exact implementation is more sophisticated than the toy depictions above, but the architectural intuition is what matters. Two Hopper capabilities make this overlapping possible:
- Tensor Memory Accelerator (TMA), which moves data asynchronously.
- Asynchronous warp-group matrix-multiply-accumulate (WGMMA), which lets tensor-core work proceed without blocking the issuing warp.
FA3 also uses warp specialization to allow data movement and computation to overlap.3 Different warps take responsibility for different stages of the pipeline: some handle data movement, some issue GEMMs, some handle softmax and update. That way a warp that’s waiting on a memory transfer doesn’t stall the whole block.
FA3 also uses a ping-pong style schedule to interleave block matrix multiplication and softmax. The goal is to avoid leaving tensor cores idle while non-matmul operations (scalar and special-function work) are performed, and vice versa.3
The online-softmax recurrence itself hasn’t changed. The normalization and value-accumulation stages still do what Part 1 describes in FlashAttention-1. What changed is the hardware pipeline used to execute them.
- FA1 asks “How do we tackle HBM traffic?”
- FA2 asks “How do we partition & parallelize the work better?”
- FA3 asks “How do we maximize the use of distinct execution resources on Hopper?”
1.6 FA3 FP8: Performance Without Pretending Precision Is Free
Now let's focus on precision. Hopper gives extremely high tensor-core throughput for FP8. That’s tempting, but it isn’t free. Simply converting every attention operand to FP8 produces unacceptable numerical error.2 Attention contains a chain of operations:
\[QK^T \rightarrow e^{x} \rightarrow \frac{e^{x}}{\sum e^{x}} \rightarrow PV.\]These operations have different numerical sensitivities. Dot products tolerate some loss. The exponential is much more sensitive. Normalization involves division. The value accumulation sums over long sequences. Each stage has its own error profile.
FA3 introduces an FP8 path with two techniques to control that error:
- Block quantization, which scales values relative to their local tile rather than globally.
- Incoherent processing, which multiplies Q and K by a random orthogonal matrix to spread outlier values across dimensions before quantization.
In the final NeurIPS 2024 publication, FA3 reaches 1.3 PFLOPs/s on H100 with this path, and reports 2.6x lower numerical error than the baseline FP8 attention method used for comparison.3
So low precision doesn’t automatically buy free performance. It shifts error into specific parts of the pipeline, therefore we have to design around that shift. This is exactly where “exact” starts to mean two different things:
- Algorithmic exactness. The algorithm still represents $\mathrm{softmax}(QK^T)V$. No pairs are skipped, no low-rank approximation is made.
- Numerical precision. The actual calculation may use FP8, and therefore experiences quantization and rounding error.
An algorithm can be exact in its mathematical formulation while a particular low-precision implementation of that algorithm is numerically less accurate. Both things are true at once, and neither cancels the other. A model’s acceptable precision depends on the full training and inference setup. FP8 support is a hardware-and-numerics feature layered onto the FlashAttention execution strategy, not proof that FP8 is universally interchangeable with BF16 or FP16.2
One more scope note: FA3’s FP8 path is Hopper-specific. The official repository still separates FA3’s implementation from other backends, so these aren’t universal features of “FlashAttention.”4
1.7 FlashAttention-4: The Blackwell Generation
In FA4, the core pattern of hardware-aware redesign continues, this time for NVIDIA Blackwell.
Blackwell exhibits asymmetric hardware scaling. Tensor-core throughput roughly doubled relative to Hopper, but several other resources didn’t scale at the same rate. Shared-memory bandwidth and exponential throughput, in particular, lag behind.5
This changes which stage is the bottleneck. See the illustration below:
Tensor cores GEMM [ | | | | | | | | | ]
Softmax/exponential [ | ]
The matrix multiplication has become extremely fast. Now the secondary softmax/exponential work becomes relatively expensive, and as such becomes the bottleneck. When one subsystem becomes much faster than the others, operations that used to be secondary can become the bottleneck.2
- On Hopper, tensor cores were fast, but the gap to softmax and shared-memory bandwidth was narrower.
- On Blackwell, the matrix-multiplication bar gets much longer, and the softmax/exponential bar stays roughly the same length, as shown in the illustration above.
FA4 redesigns both the forward and backward pipelines rather than just reusing the Hopper schedule. The major improvement techniques include:
- Fully asynchronous MMA pipelines and larger tiles.
- Software-emulated exponential plus conditional online-softmax rescaling in the forward pass.
- Tensor Memory (TMEM) and 2-CTA MMA techniques to reduce shared-memory traffic and atomic additions in the backward pass.
On B200 with BF16, the paper reports up to 1.3x speedup over cuDNN 9.13 and 2.7x over its Triton comparison, reaching 1613 TFLOPs/s (71% utilization) under the evaluated configurations.5
FA4 does not change the big-O arithmetic of dense attention. It attacks the new bottlenecks created by a GPU generation whose compute, memory, and function-unit balance differs from Hopper’s.
1.8 FA4 and Asymmetric Hardware Scaling
The FlashAttention lineage is a case study in why kernels can’t be optimized once and assumed optimal forever.2 This section contains a broader systems lesson. Let’s look at two hypothetical GPU generations (throughput in FLOP/s):
| FLOP/s | GPU A | GPU B |
|---|---|---|
| GEMM | 100 | 200 |
| Softmax | 50 | 50 |
| Memory | 50 | 50 |
GPU B makes GEMM 2x faster. But the overall system doesn’t become 2x faster automatically. Now that GEMM is less of a bottleneck, the relative cost of everything else goes up. Softmax and memory traffic, which were already on the critical path, are now even more clearly on it. This is asymmetric hardware scaling. The ratio between subsystems changes, and the optimal algorithm changes with it.
Asymmetric hardware scaling is one of the major systems principles behind the FA1 → FA4 evolution.
Blackwell is a case of this. Tensor-core throughput outran softmax and exponential throughput. So the forward pass becomes constrained by softmax and exponential work rather than GEMM. FA4 responds with a software-emulated exponential and by skipping online-softmax rescaling when the running maximum doesn’t need it. Both moves take pressure off the relatively slower non-matmul units.
The backward pass has a different bottleneck. FA4 uses Blackwell’s Tensor Memory (TMEM) and 2-CTA MMA mode to cut shared-memory traffic and atomic accumulation overhead.
This is why “FA4 is just a faster FA3” misses the point. The algorithm and kernel pipeline are co-designed around the new asymmetries, which is reflected in the paper title: Algorithm and Kernel Pipelining Co-Design for Asymmetric Hardware Scaling.5
That is also why performance numbers should always name the GPU generation. A speedup on B200 doesn’t reveal anything about H100, A100, or a non-NVIDIA accelerator. The ratios are the whole story.
A kernel is balanced against a particular hardware ratio. If tensor-core throughput improves faster than memory bandwidth or special-function throughput, the optimal algorithmic pipeline can change even when the mathematical function is identical.
1.9 FA4 Implementation and Current Status
FlashAttention-4 is a CuTe DSL implementation, optimized for Hopper and Blackwell GPUs such as the H100 and B200, and installed via the flash-attn-4 package.4 CuTe DSL is a Python-based domain-specific language from NVIDIA’s CUTLASS for writing and just-in-time compiling high-performance GPU kernels.
One engineering motivation behind the FA4 choice is compile-time productivity. The FA4 paper reports substantially faster compilation, roughly 20 to 30 times, compared with the traditional C++ template-based approach used in its comparison, while still retaining the required low-level expressivity.5 As of September 2026, PyPI classifies flash-attn-4 as Development Status :: 3 - Alpha, and its latest release (4.0.0b32, uploaded September 23, 2026) is a pre-release.6
That’s the caveat. Current implementation status is not the same thing as an evergreen property of FlashAttention. Packages, GPU support, CUDA compatibility, PyTorch, hardware, and framework integration all evolve quickly.
The durable lesson is the algorithmic progression. The current library support layer changes much faster than the underlying algorithmic ideas. Algorithmic knowledge evolves slowly. Library support can change overnight. When deploying FlashAttention, it’s worth separating the two. The algorithm tells us what should be possible. The current library tells us what’s actually available on your specific GPU, CUDA version, and framework.
Treat this as a current implementation snapshot, not an evergreen property. “FA4 exists in the official repository” and “every production environment should replace FA2 with FA4” are very different claims.
1.10 FlashAttention-1 Through FA-4 Compared
Here’s the whole evolution on one page:
| Version | Main Problem | Main Solution |
|---|---|---|
| FA1 | Excessive memory (IO/HBM) traffic | Tiling + online softmax + fused exact attention + IO awareness |
| FA2 | GPU underutilization | More sequence-level parallelism + better warp-level work partitioning + fewer non-matmul FLOPs |
| FA3 | Hopper execution imbalance | Asynchronous TMA/WGMMA pipeline + GEMM-softmax overlap + warp specialization + FP8 path |
| FA4 | Blackwell hardware asymmetry | Fully async MMA pipeline + larger tiles + optimized non-matmul work (software exp / conditional online-softmax rescale, TMEM, 2-CTA backward techniques) |
A few things to hold onto when reading those numbers:
- These are four successive generations of efficient implementations of attention, not four different attention mechanisms. The mathematical target remains fundamentally the same for dense attention.
- The exact feature matrix differs across implementations. FA3 is strongly associated with Hopper-specific optimization, while the current FA4 repository targets both Hopper and Blackwell through its CuTe DSL path.4
2. Using FlashAttention in Frameworks
2.1 PyTorch Scaled-Dot-Product Attention Today
Modern PyTorch provides a high-level API:7
torch.nn.functional.scaled_dot_product_attention(
query, key, value,
attn_mask=None,
dropout_p=0.0,
is_causal=False,
scale=None,
enable_gqa=False,
)
Conceptually, this function computes $\mathrm{softmax}(QK^T / \sqrt{d})V$. Internally, PyTorch picks an optimized implementation based on the inputs. The available backends include FlashAttention-2, a memory-efficient attention implementation, cuDNN attention, and the C++ math implementation. To restrict which backend runs, use torch.nn.attention.sdpa_kernel:8
from torch.nn.attention import SDPBackend, sdpa_kernel
import torch.nn.functional as F
with sdpa_kernel(SDPBackend.FLASH_ATTENTION):
y = F.scaled_dot_product_attention(q, k, v, is_causal=True)
If a requested fused kernel can’t run for the given inputs, PyTorch can warn with reasons when fallbacks are disabled. Backend selection and eligibility depend on device, dtype, shape, mask, and other arguments.8 The torch.nn.attention docs also expose registration and activation hooks for newer implementations, including FA3 and FA4. This surface moves fast, so treat the installed PyTorch version’s docs as the source of truth.9
- Calling the high-level API does not necessarily mean “I know exactly which kernel executed.”
- If we need to know which backend is running, restrict it explicitly with
sdpa_kernelor check the PyTorch version’s registration APIs. In production and benchmarking, verify rather than assume.
2.2 Exactness Is Not Bitwise Identity
Suppose two implementations, a math kernel and a FlashAttention kernel, compute the same dense SDPA definition. There will still be small numerical differences. The reason is straightforward. Floating-point addition is not associative:
\[(a + b) + c \neq a + (b + c)\]in finite precision, for some values. Tiling changes the reduction order. Fusion can change when values are rounded. Accumulation precision can differ.
Consider a simple example:
\[f(Q, K, V) = g(Q, K, V) = a + b + c + d\]But:
- Implementation A sums in order $1, 2, 3, 4$: $(((a + b) + c) + d)$
- Implementation B sums in order $3, 1, 4, 2$: $(((c + a) + d) + b)$
Same mathematical function, slightly different floating-point bits.
PyTorch’s reproducibility docs note this directly: SDPA backends can produce different results because they accumulate floating-point values in different orders.10 So the same mathematical function is not necessarily the same floating-point bits. Validation should use numerical tolerances, and those tolerances depend on:
- FP32 vs. BF16/FP16/FP8
- accumulation precision
- sequence length and score distribution
- hardware and backend
- dropout and randomness
Bitwise equality shouldn’t be used as the definition of exactness.
Determinism is a separate question from exactness. Libraries expose backend- and version-specific controls for it, and those controls can carry performance or memory costs.4 $^,$ 10
Exact algorithm means no deliberate mathematical approximation of dense attention. Bitwise identical implementation means every finite-precision operation lands on exactly the same bits. FlashAttention promises the first idea, not the second across arbitrary kernels.
3. FlashAttention vs. Other Techniques
3.1 FlashAttention vs. PagedAttention
The names sound related. They’re not. FlashAttention and PagedAttention answer two completely different questions:
- FlashAttention: how do I compute $QK^T$, softmax, and $PV$ efficiently?
- PagedAttention: how do I store and retrieve the growing KV cache for many serving requests?
| FlashAttention | PagedAttention | |
|---|---|---|
| Main problem | Efficient attention computation | KV-cache memory management |
| Main setting | Attention execution | LLM serving |
| Mechanism | Tiling, online softmax, fused/hardware-aware kernels | Paging KV cache |
| Main state | Attention tiles and row-wise normalization/output state | Persistent per-request KV cache |
| Changes dense attention formula? | No for dense FlashAttention | No, primarily changes cache memory management |
| Specific task to tackle | $N^2$ score/probability intermediates | fragmented per-request KV-cache allocation |
PagedAttention, introduced with vLLM, reduces KV-cache waste from fragmentation and duplication in high-throughput serving.11 FlashAttention doesn’t manage KV cache allocation. It handles how the attention operation consumes $Q/K/V$. How systems access the persistent KV state that lives across requests is a separate concern handled by PagedAttention. That separation is why they can coexist in the same serving stack without stepping on each other. Dao-AILab even ships a KV-cache-oriented FlashAttention interface with optional block tables. Kernel execution and cache paging are two modular concerns that play nicely together.4
In practice, the stack looks like this:
Paged KV cache → FlashAttention-style kernel → Attention output
- If the problem statement contains “$N^2$ score/probability intermediates,” think FlashAttention.
- If it contains “fragmented per-request KV-cache allocation,” think PagedAttention.
3.2 FlashAttention vs. Sparse and Linear Attention
Three very different strategies, grouped together because they all make attention cheaper in some regime. They are not the same kind of change.
| Dense FlashAttention | Sparse Attention | Linear Attention | |
|---|---|---|---|
| What it does | Same dense attention, different execution | Fewer interactions | Different mathematical formulation |
| Attention pattern | Dense (all pairs) | Sparse or local | Rearranged via feature maps or associativity |
| Scaling | Still $O(N^2 d)$ | Can fall below full $N^2$ work | Can fall below full $N^2$ work |
Dense FlashAttention
- computes $\mathrm{softmax}(QK^T)V$ for all required Q-K pairs.
- changes the execution of dense attention, not the mathematical function.
Sparse / local attention
- only computes a selected subset of $Q$-$K$ interactions.
- $i \rightarrow j$ only if $\lvert i-j \rvert \leq w$ where $w$ is the window size.
- fewer $Q$-$K$ interactions, and the model’s attention pattern changes.
Linear attention
- changes the mathematical formula itself.
- usually rearranges the attention computation with feature maps or associativity, avoiding an explicit dense softmax score matrix.
The phrase “memory-efficient attention” is too broad to identify a method by itself. Ask two questions instead:2
- Does it compute the same dense softmax attention?
- Does it reduce arithmetic pairs, or mostly memory traffic/intermediate storage?
The first question splits dense FlashAttention from sparse and linear approaches. The second separates FlashAttention from methods that change the number of interactions that get computed.
The original FlashAttention paper makes this line explicit. It contrasts its dense exact algorithm with approximate methods on one side, and explores block-sparse FlashAttention as an approximate/sparse extension on the other.12
FlashAttention reduces memory requirements and IO, but dense attention remains quadratic.
In summary:
- FlashAttention: same dense attention, different execution.
- Sparse attention: fewer interactions, different attention pattern.
- Linear attention: different mathematical formulation.
“FlashAttention makes attention linear” is wrong. If a system shows near-linear scaling because it uses a local window or another sparse pattern, the sparsity is what changed the number of interactions. FlashAttention may still be the kernel underneath.
4. Training vs. Inference
4.1 Training, Prefill, and Decode Are Different Regimes
FlashAttention shines when both query and key sequence dimensions are large. Plenty of query rows means avoiding the materialized attention matrix actually pays off, and tiled GEMMs have real work to do.
Training. Suppose $N_q = N_k = N$. There are many query and key rows. FlashAttention can exploit large GEMMs, tiling, and massive parallelism, and it avoids $N^2$ intermediates. This is a very favorable regime.
Prefill. Suppose the user sends a prompt of $8000$ tokens. The model processes those tokens together. Again $N_q \approx N_k$, so dense attention is substantial. FlashAttention can be very useful here. Prefill is typically compute-bound.
Decode. Now suppose the model has already generated $8000$ tokens and wants to generate token $8001$. The new query might have $N_q = 1$ while $N_k = 8000$. So the computation is now:
\[1 \times 8000 \text{, not } 8000 \times 8000\]The bottleneck changes. Decode is heavily interacting with the existing KV cache. The system’s behavior is now driven by how fast it can read that cache, not by how much arithmetic it can do. Decode is typically memory-bound. Serving and batching behavior matter here too.2 The official repository exposes flash_attn_with_kvcache, which updates and attends to the cache in one kernel and supports MQA/GQA and optional RoPE handling.4
The key distinction in one line:
- Prefill: many query rows → large attention computation.
- Decode: one query row → large KV-cache memory access.
FlashAttention also cannot eliminate other costs. MLP layers, communication in tensor parallelism, KV-cache capacity, model-weight bandwidth, sampling, and scheduler overhead remain separate concerns.2
So when someone says “FlashAttention makes LLM inference faster,” ask:
- Prefill or decode?
- Compute-bound or memory-bound?
- Sequence length?
- Batch size?
- KV cache?
- GPU?
- Which kernel?
This does not mean FlashAttention is irrelevant to inference. Prefill is an attention-heavy dense regime, and decode can still use specialized kernels from the same implementation family. It means “FlashAttention speeds up LLM inference” is incomplete unless you say which phase and bottleneck.
5. Practical Engineering for FlashAttention
5.1 Common Implementation Mistakes
A running list of things that bite people in practice:2
- Assuming the flash backend always ran. High-level frameworks may fall back when dtype, device, shape, or mask is unsupported. Use backend controls and profiling when it matters.
- Leaving dropout on during evaluation. PyTorch SDPA applies
dropout_pexactly as supplied, so pass0.0when evaluation requires no dropout. - Using the wrong mask convention. Boolean mask semantics differ across APIs.
- Expecting bitwise agreement. Fused kernels can reorder floating-point operations.
- Ignoring GQA head constraints. Current interfaces require compatible/divisible query and KV head counts.
- Treating current head-dimension/dtype support as timeless. These constraints change across releases and backends.
- Benchmarking without synchronization or warmup. Asynchronous GPU execution can make naive wall-clock measurements meaningless.
- Calling a sparse or local pattern “dense FlashAttention.” Kernel and attention pattern are separate choices.
- Confusing prefill with decode. The shapes and bottlenecks are different.
- Assuming the newest generation is always deployable. Current FA4 package metadata is alpha, and hardware/software compatibility must be checked.
5.2 Common Misconceptions
A list of things people say that aren’t quite right about FlashAttention, and what’s actually true.2
-
“FlashAttention approximates softmax.”
No. Dense FlashAttention evaluates the dense softmax-attention function using tiled online normalization. Nothing is thrown away, nothing is approximated. -
“FlashAttention makes attention $O(N)$.”
No. Dense arithmetic remains $O(N^2 d)$. What becomes linear is the large auxiliary attention-memory footprint with respect to sequence length. HBM traffic is reduced, but the underlying number of query-key interactions is unchanged. -
“It is just kernel fusion.”
Fusion is part of the implementation story. But online-softmax tiling and IO-aware scheduling are algorithmic, not merely concatenating existing kernels. Fusion alone wouldn’t produce the same memory profile or the same IO complexity bound. -
“It is the same as PagedAttention.”
No. PagedAttention manages KV-cache allocation for serving. FlashAttention manages how the attention operation itself is executed. Different problems, different solutions. -
“Exact means bitwise identical.”
No. Floating-point operation order can change rounding. Two exact implementations of the same function can produce slightly different bits, and that’s expected. -
“Long context becomes free.”
No. Dense pairwise arithmetic is still quadratic, and KV-cache and other model costs remain. What FlashAttention gives us is a smaller memory footprint and less HBM traffic, not a free lunch. -
“FA1, FA2, FA3, and FA4 are different attention architectures.”
No. They are successive generations of efficient kernels and algorithms, shaped by different hardware bottlenecks. The mathematical target is fundamentally the same for dense attention.
6. Conclusion
6.1 The Mental Model to Have
Think of FlashAttention as a streaming matrix computation with exact streaming softmax. Picture the full attention matrix, $Q$ on the $y$-axis and $K$ on the $x$-axis. It’s a big grid.
\[Q \quad \begin{array}{|c|c|c|c|} \hline & & & \\ \hline & & & \\ \hline & & & \\ \hline & & & \\ \hline \end{array} \\\] \[\qquad K\]- Naive attention:
- Calculate the entire matrix grid and store it.
- Flash attention:
- Bring one small region of the matrix grid close to the compute units. Calculate that region, normalize it using running statistics, immediately use the result to accumulate the output, then discard the region. Then move on to the next region.
- The mathematical interactions remain. The physical representation of the intermediate computation does not. That’s the central insight.
The evolution from FA1 to FA4, in one line each:
- FA1 → “Reduce expensive HBM movement.”
- FA2 → “Now use the GPU more effectively via better parallelism.”
- FA3 → “Now overlap communication and memory movement on Hopper.”
- FA4 → “Now redesign the pipeline for Blackwell’s new hardware balance.”
This evolution is why FlashAttention is much more interesting than just “a faster attention kernel.” It’s a case study in algorithm-hardware co-design. The mathematical function can stay identical while the optimal execution strategy changes dramatically as the memory hierarchy, compute throughput, and specialized hardware change.
The final practical mental model:
- keep score tiles near compute,
- carry only the row statistics needed to merge softmax blocks,
- immediately consume probabilities for each tile into the $V$ accumulation, and
- avoid sending the full $N^2$ attention matrix through HBM.
6.2 FlashAttention Summary
Let’s walk through the whole FlashAttention story in one derivation, and hit the key points to remember.
Ordinary attention: $QK^T \in \mathbb{R}^{N \times N}$. Materializing it creates an enormous intermediate matrix.13
\[O = \mathrm{softmax}\left(\frac{QK^T}{\sqrt{d}}\right) V.\]Naive execution: This creates huge memory traffic.
\[QK^T \to \text{store } N \times N \text{ scores} \to \text{softmax} \to \text{store } N \times N \text{ probabilities} \to \text{multiply by } V\]FlashAttention: Tiling is used. Split $Q \to Q_i$ and $K, V \to K_j, V_j$, then calculate the score tile:
\[S_{ij} = \frac{Q_i K_j^T}{\sqrt{d}}.\]-
The problem with tiling: Softmax needs the maximum and the denominator over the whole row.
-
The solution: online softmax
- Maintain $(m, \ell, a)$ where:
- $m$ = running maximum
- $\ell$ = running sum of exponentials (the softmax denominator)
- $a$ = running unnormalized value-weighted sum (accumulator)
- When a new block has a larger maximum $m’$, calculate a rescaling factor $\alpha$ and rescale the old state ($\ell$ and $a$) by it:
- Maintain $(m, \ell, a)$ where:
-
Therefore, we can process the full FlashAttention pipeline without ever materializing the $N \times N$ matrix:
\[\text{tile} \to \text{score} \to \text{online softmax} \to V \text{ accumulation} \to \text{discard tile}.\] -
Result: The same exact dense attention mathematics, but much better memory behavior: less HBM traffic, much smaller intermediates. Arithmetic remains $O(N^2 d)$.
FlashAttention in one table:
| Property | Complexity / Formulation |
|---|---|
| Memory complexity | $O(Nd)$ auxiliary attention storage |
| Compute complexity | $O(N^2 d)$ dense attention compute |
| Attention formulation | Exact dense softmax attention |
Same exact dense attention mathematics, but much better memory behavior. Less HBM traffic, much smaller intermediates, arithmetic unchanged at $O(N^2 d)$.
Appendix
Reference Map
The table below shows what paper to refer to for each question.2
| Question | Best starting source |
|---|---|
| What is the original IO-aware algorithm? | FlashAttention-1 paper 12 |
| Why does online normalization work? | Milakov & Gimelshein 14 |
| Was exact subquadratic-memory attention known earlier? | Rabe & Staats 15 |
| What specifically changed in FA2? | FlashAttention-2 paper 1 |
| What is Hopper-specific in FA3? | FlashAttention-3 paper 3 |
| What changes on Blackwell in FA4? | FlashAttention-4 paper 5 |
| What does the current official package support? | Dao-AILab repository 4 |
| How does current PyTorch select SDPA kernels? | PyTorch SDPA and sdpa_kernel docs 7 $^,$ 8 |
| Why can exact backends differ numerically? | PyTorch reproducibility notes 10 |
| How is PagedAttention different? | Kwon et al. 11 |
Implementation support changes faster than the mathematics. For deployed systems, treat the installed framework’s documentation and the exact kernel repository/release as authoritative for supported devices, dtypes, head dimensions, masks, dropout, and GQA behavior.
References
Citation
If you found this blog post helpful, please consider citing it:
@article{obasi2026FlashAttentionPt2,
title = "FlashAttention: Part 2",
author = "Obasi, Chizoba",
journal = "chizkidd.github.io",
year = "2026",
month = "Sep",
url = "https://chizkidd.github.io/2026/09/17/flashattention-2/"
}
-
Tri Dao. FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning. ICLR 2024. ↩ ↩2 ↩3 ↩4 ↩5
-
Understanding FlashAttention: How IO-Aware Attention Makes Transformers Faster Without Approximating Attention. Google Drive PDF. Accessed September 2026. ↩ ↩2 ↩3 ↩4 ↩5 ↩6 ↩7 ↩8 ↩9 ↩10 ↩11
-
Jay Shah et al. FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision. NeurIPS 2024. ↩ ↩2 ↩3 ↩4 ↩5
-
Dao-AILab. flash-attention official repository and current implementation documentation. GitHub. Accessed September 2026. ↩ ↩2 ↩3 ↩4 ↩5 ↩6 ↩7
-
Ted Zadouri et al. FlashAttention-4: Algorithm and Kernel Pipelining Co-Design for Asymmetric Hardware Scaling. arXiv, 2026. ↩ ↩2 ↩3 ↩4 ↩5
-
PyPI. flash-attn-4 package metadata and release history. PyPI. Accessed September 2026. ↩
-
PyTorch. torch.nn.functional.scaled_dot_product_attention documentation. PyTorch Documentation. Accessed September 2026. ↩ ↩2
-
PyTorch. torch.nn.attention.sdpa_kernel documentation. PyTorch Documentation. Accessed September 2026. ↩ ↩2 ↩3
-
PyTorch. torch.nn.attention module and FlashAttention implementation registration APIs. PyTorch Documentation. Accessed September 2026. ↩
-
PyTorch. Reproducibility documentation, including scaled-dot-product attention backend differences. PyTorch Documentation. Accessed September 2026. ↩ ↩2 ↩3
-
Woosuk Kwon et al. Efficient Memory Management for Large Language Model Serving with PagedAttention. SOSP 2023. ↩ ↩2
-
Tri Dao et al. FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness. NeurIPS 2022. ↩ ↩2
-
Ashish Vaswani et al. Attention Is All You Need. NeurIPS 2017. ↩
-
Maxim Milakov and Natalia Gimelshein. Online normalizer calculation for softmax. arXiv, 2018. ↩
-
Markus N. Rabe and Charles Staats. Self-attention Does Not Need O(n²) Memory. arXiv, 2021. ↩