Activation Checkpointing & FlashAttention Mechanics
Storing all layer activations during the forward pass consumes tens of GBs of VRAM. Activation Checkpointing drops intermediate layer activations and recomputes them on-the-fly during the backward pass.
Trading Compute for Memory
By storing activations only at segment boundaries (e.g. every sqrt(L) layers), activation memory drops from O(L) to O(sqrt(L)). The cost is ~33% additional forward pass compute time.
FlashAttention-2 / 3 (Tiled Compute)
Standard attention writes N x N score matrices to slow GPU HBM. FlashAttention tiles Query, Key, and Value blocks inside fast GPU SRAM (20 TB/s bandwidth), computing online Softmax without ever materializing the quadratic matrix in HBM.
Selective Activation Recompute
Recomputing only memory-heavy, compute-cheap ops (like RMSNorm, SwiGLU, or Attention Dropout) saves ~70% memory with less than 10% compute penalty.
Online Softmax tracks running row maximums and rescale factors across SRAM tiles without full matrix materialization.
- FlashAttention-3 achieves over 800 TFLOPS (75%+ MFU) on H100 GPUs by leveraging FP8 Tensor Cores and asynchronous CUDA WGMMA instructions.