Standard Transformer attention scales quadratically in memory IO operations with sequence length N. While exact attention computation requires O(N^2) floating-point operations, the primary bottleneck on modern GPUs is not compute capacity — it is memory bandwidth. FlashAttention-3 solves this bottleneck by reformulating scaled dot-product attention to exploit GPU SRAM tiling, asynchronous Tensor Core pipelines, and FP8 low-precision execution.
Standard attention is HBM memory-bandwidth bound; FlashAttention re-writes the algorithm to make attention SRAM compute-bound.
The memory hierarchy problem in GPU computing
Modern NVIDIA GPUs (such as H100 and A100) feature two distinct storage tiers: High Bandwidth Memory (HBM3) with ~3.3 TB/s bandwidth, and On-Chip SRAM (L1/Shared Memory) with ~33 TB/s bandwidth — an order of magnitude faster.
+-------------------------------------------------------+
| High Bandwidth Memory (HBM3): ~80GB @ 3.3 TB/sec |
+---------------------------+---------------------------+
| (Heavy IO Bottleneck)
+---------------------------v---------------------------+
| On-Chip SRAM (L1 Cache): ~50MB @ 33 TB/sec |
+---------------------------+---------------------------+
| (Fast Register Access)
+---------------------------v---------------------------+
| Tensor Cores / CUDA Vector Units (H100 FLOPS Engine) |
+-------------------------------------------------------+
Standard attention reads Q, K, V matrices from HBM to SRAM, writes the intermediate N x N attention matrix S back to HBM, reads S to compute P = softmax(S), writes P back to HBM, and finally reads P and V to compute output O. This results in O(N^2) HBM reads and writes.
Online softmax tiling mathematics
FlashAttention avoids writing the intermediate N x N matrix to HBM by computing attention in SRAM blocks using the Online Softmax trick:
# Standard Softmax:
m = max(x)
p = exp(x - m)
l = sum(p)
softmax(x) = p / l
# Online Softmax update when combining block 1 (m1, l1) and block 2 (m2, l2):
m_new = max(m1, m2)
l_new = exp(m1 - m_new) * l1 + exp(m2 - m_new) * l2
O_new = diag(exp(m1 - m_new)) * O1 + diag(exp(m2 - m_new)) * O2
FlashAttention-3: Async pipelines & Tensor Core overlapping
On NVIDIA Hopper (H100) architectures, FlashAttention-3 introduces three structural performance accelerations:
- Producer-Consumer Warp Separation: Dedicated warps issue asynchronous TMA (Tensor Memory Accelerator) loads while separate consumer warps perform GEMM computations.
- Interleaved GEMM and Softmax Pipelines: Overlaps the GEMM
Q * K^Tcomputation for blocki+1with the softmax reduction for blocki. - FP8 Low-Precision Support: Leverages FP8 (E4M3 format) Tensor Cores to double FLOPS density (1.98 PFLOPS on H100).
Triton Python kernel implementation
import torch
import triton
import triton.language as tl
@triton.jit
def _flash_attn_fwd_kernel(
Q, K, V, sm_scale,
L, Out,
stride_qz, stride_qh, stride_qm, stride_qk,
stride_kz, stride_kh, stride_kn, stride_kk,
stride_vz, stride_vh, stride_vn, stride_vk,
stride_oz, stride_oh, stride_om, stride_ok,
Z, H, N_CTX,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
):
start_m = tl.program_id(0)
off_hz = tl.program_id(1)
# Initialize pointers for Q, K, V blocks
offs_m = start_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = tl.arange(0, BLOCK_N)
offs_d = tl.arange(0, 64)
# Load Q block to SRAM registers
q_ptrs = Q + off_hz * stride_qh + offs_m[:, None] * stride_qm + offs_d[None, :] * stride_qk
q = tl.load(q_ptrs, mask=offs_m[:, None] < N_CTX, other=0.0)
# Initialize online softmax accumulators
m_i = tl.zeros([BLOCK_M], dtype=tl.float32) - float("inf")
l_i = tl.zeros([BLOCK_M], dtype=tl.float32)
acc = tl.zeros([BLOCK_M, 64], dtype=tl.float32)
# Loop over K, V blocks
for start_n in range(0, N_CTX, BLOCK_N):
k_ptrs = K + off_hz * stride_kh + (start_n + offs_n)[:, None] * stride_kn + offs_d[None, :] * stride_kk
k = tl.load(k_ptrs, mask=(start_n + offs_n)[:, None] < N_CTX, other=0.0)
# Compute Q * K^T
qk = tl.zeros([BLOCK_M, BLOCK_N], dtype=tl.float32)
qk += tl.dot(q, tl.trans(k)) * sm_scale
# Online Softmax update
m_ij = tl.max(qk, 1)
m_new = tl.maximum(m_i, m_ij)
alpha = tl.exp(m_i - m_new)
beta = tl.exp(qk - m_new[:, None])
l_i = l_i * alpha + tl.sum(beta, 1)
# Load V and accumulate output
v_ptrs = V + off_hz * stride_vh + (start_n + offs_n)[:, None] * stride_vn + offs_d[None, :] * stride_vk
v = tl.load(v_ptrs, mask=(start_n + offs_n)[:, None] < N_CTX, other=0.0)
acc = acc * alpha[:, None] + tl.dot(beta.to(tl.float16), v)
m_i = m_new
# Store final output block back to HBM
acc = acc / l_i[:, None]
out_ptrs = Out + off_hz * stride_oh + offs_m[:, None] * stride_om + offs_d[None, :] * stride_ok
tl.store(out_ptrs, acc.to(tl.float16), mask=offs_m[:, None] < N_CTX)
Performance benchmarks: H100 SXM5 throughput
On NVIDIA H100 SXM5 GPUs with 80GB HBM3 memory, FlashAttention-3 achieves:
- 1.2 PFLOPS FP16 Throughput: Over 75% of maximum theoretical hardware FLOPS limit.
- 1.9 PFLOPS FP8 Throughput: 2.2x speedup over FlashAttention-2.
- Zero Extra Memory Overhead: Memory requirement scales linearly
O(N)rather than quadraticallyO(N^2).