FlashAttention-2 and -3: From Algebra to GPU Scheduling
Keep the unchanged mathematical target in view
The running maximum, denominator, and numerator still define the output. Correctness comes from their merge identity. Performance comes from the schedule: which queries a thread block owns, which data are reused, where partial results live, and when synchronization occurs.
A fast kernel may deliberately perform some redundant arithmetic to avoid expensive transfers or synchronization. Counting only matrix-multiply FLOPs misses costs from reductions, exponentials, divisions, memory operations, and thread coordination.
FlashAttention-2: expose enough parallel work
FlashAttention-2 improves work partitioning and reduces non-matmul overhead. One important source of parallelism is the query sequence: different query tiles can be processed independently while each iterates over the relevant key/value tiles. This matters when batch size and head count alone do not provide enough work to occupy the GPU.
Within a thread block, arranging work so warps own different query rows can reduce communication of partial outputs compared with splitting a query's reduction across multiple warps. The useful question is who owns the accumulated numerator and denominator, and how often those quantities must cross thread or memory boundaries.
These changes do not turn dense attention into a linear-time operation. They reduce constant factors and improve hardware utilization for a particular workload. A small context or unusual head dimension can produce a different bottleneck.
Why “more parallelism” has limits
Suppose a kernel has only eight independent thread blocks but the device can execute far more concurrently. Splitting work along the query dimension can expose enough blocks to use the device. However, making tiles too small reduces reuse and increases overhead. Making them too large can consume registers or shared memory and reduce the number of active blocks.
Occupancy is the number of active execution resources, not a direct measurement of useful throughput. A kernel with high occupancy may spend much of its time waiting on memory. A kernel with lower occupancy may run efficiently if its work is compute-dense and well pipelined.
FlashAttention-3: overlap work on Hopper
FlashAttention-3 targets Hopper capabilities including asynchronous data movement and matrix-multiply execution. Its techniques overlap data transfer with computation and interleave matrix operations with softmax work, rather than forcing every stage to finish before another begins.
A useful conceptual schedule is: while a tile's values participate in an output product, prepare another score tile and prefetch later data. Actual dependencies still constrain execution. Probabilities cannot be consumed before their scores and normalization are ready; buffers cannot be overwritten before their consumers finish.
TMA transfers and asynchronous matrix instructions require thread-issued commands and synchronization. They do not mean memory movement happens without coordination or that softmax becomes free. A pipeline's steady state, startup, and drain phases have different utilization.
A pipeline is constrained by its slowest stage
Imagine illustrative stage times of 3 units for loading, 5 for matrix work, and 2 for scalar reductions. Serial execution would take 10 units per tile. An ideal overlapped steady state cannot run faster than the limiting resource permits; five units is an optimistic bound only if those operations use independent resources and all dependencies can be scheduled accordingly.
This is an analytical illustration, not a FlashAttention benchmark. If two nominal stages compete for the same tensor cores or shared-memory bandwidth, simply taking their maximum understates the cost. Measure the actual kernel and include fill/drain overhead for short workloads.
Low precision changes numerical error
FP8 storage and multiplication can increase hardware throughput, but quantization introduces error. FP32 accumulation or softmax reductions do not restore information lost when inputs were quantized. Block scaling and the distribution of outliers affect how much error is introduced.
The paper's incoherent-processing approach uses a shared orthogonal transformation to spread outliers before quantization. For an orthogonal $R$, $(QR)(KR)^\top=QK^\top$ in exact arithmetic. A random sign flip alone preserves coordinate magnitudes and therefore does not by itself spread an outlier across coordinates; mixing transformations are the relevant distinction.
How to compare kernels honestly
Record GPU model, dtype, head dimensions, query/key lengths, batch and head counts, causal mode, dropout, and whether timing includes backward. Synchronize asynchronous work and exclude or separately report compilation and warm-up. Compare numerical errors against a reference and measure the same operation on both sides.
Peak hardware throughput, attention-kernel throughput, and end-to-end model throughput are different denominators. Faster attention may have a limited overall effect if MLPs, communication, or the vocabulary head dominate. Avoid transferring a reported speedup between GPUs or workloads without remeasurement.
Try it: If a kernel is twice as fast but attention was only one quarter of total runtime, what is the ideal whole-model speedup?
With all other work unchanged, runtime becomes $0.75+0.25/2=0.875$ of the original, a speedup of about 1.143 times. This is an idealized Amdahl-law calculation; actual integration can add other effects.
Another use of mergeable statistics
Ring Attention distributes key/value tiles across devices and combines their contributions using the same softmax mathematics.