All articles GPU Systems & Kernels

FlashAttention-3 & GPU Kernel Optimization: Accelerating Transformer Attention

An in-depth technical analysis of FlashAttention-1, 2, and 3: IO-awareness, FP8 Tensor Cores, async memory pipelines, and custom Triton kernel design.

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^T computation for block i+1 with the softmax reduction for block i.
  • 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 quadratically O(N^2).
← Back to all articles