← All Posts
From Scratch · C++ · Transformers · Attention Variants· Part 8 of 13

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:

For each Q block $Q_i$ (rows from $i \cdot B_r$ to $(i+1) \cdot B_r - 1$):

  1. Initialize: $m_i \leftarrow -\infty$ (per-row max of attention scores), $\ell_i \leftarrow 0$ (exponential sum), $O_i \leftarrow 0$ (output accumulator).
  2. 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}}$.
  3. 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:

torch::Tensor flash_attention_cpu( torch::Tensor Q, torch::Tensor K, torch::Tensor V, int Br, int Bc) { // Q, K, V: (seq_len, d_k) [simplified: batch=1] // Br, Bc: block row, block col sizes int T = Q.size(0), d_k = Q.size(1); int N = (T + Br - 1) / Br; // # Q blocks int M = (T + Bc - 1) / Bc; // # KV blocks torch::Tensor O = torch::zeros({T, d_k}, Q.options()); torch::Tensor m = torch::full({T}, -1e9, Q.options()); // row-wise max torch::Tensor ℓ = torch::zeros({T}, Q.options()); // exp sum per row // Outer loop: Q blocks for (int i = 0; i < N; i++) { int q_start = i * Br, q_end = min(q_start + Br, T); int q_len = q_end - q_start; torch::Tensor Q_i = Q.slice(0, q_start, q_end); // (q_len, d_k) torch::Tensor m_i = torch::full({q_len}, -1e9, Q.options()); torch::Tensor ℓ_i = torch::zeros({q_len}, Q.options()); torch::Tensor O_i = torch::zeros({q_len, d_k}, Q.options()); // Inner loop: KV blocks for (int j = 0; j < M; j++) { int kv_start = j * Bc, kv_end = min(kv_start + Bc, T); int kv_len = kv_end - kv_start; torch::Tensor K_j = K.slice(0, kv_start, kv_end); // (kv_len, d_k) torch::Tensor V_j = V.slice(0, kv_start, kv_end); // (kv_len, d_k) // Compute S_ij = Q_i @ K_j^T [shape: (q_len, kv_len)] torch::Tensor S_ij = torch::matmul(Q_i, K_j.t()); // Online softmax: compute m_ij_new, exp-weights torch::Tensor m_ij = S_ij.max(1).values; // (q_len,) row-wise max torch::Tensor m_ij_new = torch::max(m_i, m_ij); torch::Tensor alpha = torch::exp(m_i.unsqueeze(1) - m_ij_new.unsqueeze(1)); torch::Tensor beta = torch::exp(S_ij - m_ij_new.unsqueeze(1)); // Rescale accumulators O_i = alpha * O_i; ℓ_i = alpha.squeeze(1) * ℓ_i; // Accumulate O_i = O_i + torch::matmul(beta, V_j); ℓ_i = ℓ_i + beta.sum(1); // Update max m_i = m_ij_new; } // Normalize: O_i / ℓ_i (per-row division) O_i = O_i / ℓ_i.unsqueeze(1); O.slice(0, q_start, q_end).copy_(O_i); } return O; }

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.

Production note: Real-world implementations use CUDA kernels for tiling with fused operations. This C++ skeleton is for understanding the algorithm. Integrate with at::parallel_for for CPU parallelism, or call optimized CUDA routines from PyTorch/Flash-Attention libraries for production use.