FlashAttention in C++ (LibTorch)
FlashAttention (Dao et al., 2022) revolutionized transformer efficiency by reordering attention computation to minimize HBM (memory bandwidth) bottlenecks. This post walks through the core algorithm: tiling the Q, K, V matrices into blocks, performing online softmax updates, and rewriting attention as a sequence of sequential GEMM+softmax operations. We'll implement a CPU-friendly pseudocode version in C++ using LibTorch, then visualize the tiling and accumulator flow.
1 · The Key Insight: Block-Wise Softmax
Standard attention computes $S = QK^T$, $P = \text{softmax}(S)$, $O = PV$ all in one pass. For long sequences, this requires materializing $S$ and $P$ as (seq_len, seq_len) matrices — quadratic memory and register pressure.
FlashAttention observes that softmax can be computed online in a single forward pass over the blocks. Instead of computing $\text{softmax}(QK^T)$ over the entire K dimension at once, we split K and V into blocks $K_j$, $V_j$, process each pair $(K_j, V_j)$ in sequence, and maintain running statistics: a running max $m_i$ and exponential sum $\ell_i$ per Q block. After each KV block, we rescale the output accumulators.
This is the essence of online softmax and reduces intermediate memory from $O(N^2)$ to $O(N \cdot d_k)$ (only storing the current KV block and a small O(N) running state).
2 · Algorithm Outline
Let:
- $B_r$ = block size for Q (rows of Q)
- $B_c$ = block size for K, V (columns, or time dimension)
- $T$ = sequence length (multiples of $B_c$ for simplicity)
- $N$ = T / $B_r$ (number of Q blocks)
- $M$ = T / $B_c$ (number of KV blocks)
For each Q block $Q_i$ (rows from $i \cdot B_r$ to $(i+1) \cdot B_r - 1$):
- Initialize: $m_i \leftarrow -\infty$ (per-row max of attention scores), $\ell_i \leftarrow 0$ (exponential sum), $O_i \leftarrow 0$ (output accumulator).
- For each KV block $K_j$, $V_j$ (with $j = 0 \ldots M-1$):
- Compute $S_{ij} = Q_i K_j^T$ (shape: $B_r \times B_c$).
- Compute max-plus softmax: for each row of $S_{ij}$, find $m_{ij}^{\text{new}}$ (the new max), compute $e^{S_{ij} - m_{ij}^{\text{new}}}$, and track $\ell_{ij}^{\text{new}}$ (the exponential sum).
- Rescale previous accumulators: $O_i \leftarrow e^{m_i - m_{ij}^{\text{new}}} \cdot O_i$, $\ell_i \leftarrow e^{m_i - m_{ij}^{\text{new}}} \cdot \ell_i$.
- Accumulate: $O_i \leftarrow O_i + e^{S_{ij} - m_{ij}^{\text{new}}} \cdot V_j$, $\ell_i \leftarrow \ell_i + \text{sum}(e^{S_{ij} - m_{ij}^{\text{new}}}, \text{axis}=1)$.
- Update: $m_i \leftarrow m_{ij}^{\text{new}}$.
- Final normalize: $O_i \leftarrow O_i / \ell_i$ (broadcast divide per row).
3 · C++ Pseudocode with LibTorch
Below is a pseudo-implementation showing the two-level loop structure and the online softmax mechanics using LibTorch operations:
4 · Shape Flow Walkthrough
| Operation | Input Shapes | Output Shape | Notes |
|---|---|---|---|
Q_i = Q[q_start:q_end] |
Q: (T, d_k) | (q_len, d_k) | Extract Q block |
K_j, V_j = slice |
K, V: (T, d_k) | (kv_len, d_k) | Extract KV block |
S_ij = Q_i @ K_j^T |
Q_i: (q_len, d_k), K_j^T: (d_k, kv_len) | (q_len, kv_len) | Attention scores |
m_ij = max(S_ij, axis=1) |
S_ij: (q_len, kv_len) | (q_len,) | Row-wise max for online softmax |
beta = exp(S_ij - m_ij_new) |
S_ij: (q_len, kv_len), m_ij_new: (q_len,) | (q_len, kv_len) | Normalized attention weights |
O_i += beta @ V_j |
beta: (q_len, kv_len), V_j: (kv_len, d_k) | (q_len, d_k) | Accumulate weighted values |
O_i / ℓ_i |
O_i: (q_len, d_k), ℓ_i: (q_len,) | (q_len, d_k) | Final normalization per row |
5 · The Two-Level Tiling Animation
Block-Wise Attention Flow
Watch how each Q block processes all KV blocks sequentially, updating (m, ℓ, O) accumulators.
6 · Why This Matters
The tiling strategy dramatically reduces memory movement. Instead of materializing the full (T, T) attention matrix, we process (Br × Bc) tiles, and intermediate accumulators (m, ℓ, O) stay in fast memory across block iterations. In CUDA, this maps to SRAM reads and writes for S, P, and smaller tensors, while V is loaded sequentially from HBM. For long sequences, FlashAttention is 2–4× faster than standard attention on modern GPUs.
This implementation is CPU-side pseudocode (no CUDA kernel fusion). A true GPU FlashAttention kernel further exploits tiling to maximize compute intensity. But the algorithmic core — online softmax with block-wise accumulation — is exactly what we've walked through.
7 · Multi-Batch and Head-Wise Extension
For batched inputs with multiple attention heads, wrap the pseudocode in outer loops over batch and head. Each (batch, head) pair runs the block-wise algorithm independently. LibTorch's parallelization via torch::parallel_for or thread pool can accelerate these loops.
at::parallel_for for CPU parallelism, or call optimized CUDA routines from PyTorch/Flash-Attention libraries for production use.