Overview
Flash Attention provides 2-4x speedup and 10-20x memory reduction for transformer attention through IO-aware tiling and recomputation. It supports PyTorch native SDPA, the flash-attn library, H100 FP8, and sliding window attention.
When to Use
- Training transformers with sequences >512 tokens.
- Running inference with long context (>2K tokens).
- GPU memory constrained (OOM with standard attention).
- Need 2-4x speedup without accuracy loss.
- Using PyTorch 2.2+ or can install
flash-attn.
Prerequisites
- GPU: NVIDIA Ampere+ (A100, A10, A30) or AMD MI200+. Turing (T4) is supported. Volta (V100) is NOT supported.
- VRAM: Same as standard attention (Flash Attention doesn't increase memory).
- CUDA: 12.0+ (11.8 minimum).
- PyTorch: 2.2+ for native SDPA support.
Procedure
Workflow 1: Enable in existing PyTorch model (Native SDPA)
- Check PyTorch version (≥2.2):
python -c "import torch; print(torch.__version__)"
If <2.2, upgrade:
pip install --upgrade torch
- Enable Flash Attention backend: Replace standard attention:
# Before (standard attention)
attn_weights = torch.softmax(q @ k.transpose(-2, -1) / math.sqrt(d_k), dim=-1)
out = attn_weights @ v
# After (Flash Attention)
import torch.nn.functional as F
out = F.scaled_dot_product_attention(q, k, v, attn_mask=mask)
Force Flash Attention backend:
with torch.backends.cuda.sdp_kernel(
enable_flash=True,
enable_math=False,
enable_mem_efficient=False
):
out = F.scaled_dot_product_attention(q, k, v)
Workflow 2: Use flash-attn library for advanced features
- Install flash-attn library:
pip install flash-attn --no-build-isolation
- Modify attention code:
from flash_attn import flash_attn_func
# Input: [batch_size, seq_len, num_heads, head_dim]
# Transpose from [batch, heads, seq, dim] if needed
q = q.transpose(1, 2)
k = k.transpose(1, 2)
v = v.transpose(1, 2)
out = flash_attn_func(
q, k, v,
dropout_p=0.1,
causal=True,
window_size=(-1, -1), # No sliding window
softmax_scale=None # Auto-scale
)
out = out.transpose(1, 2) # Back to [batch, heads, seq, dim]
- Enable advanced features: Sliding window attention (local attention):
# Only attend to window of 256 tokens before/after
out = flash_attn_func(
q, k, v,
window_size=(256, 256), # (left, right) window
causal=True
)
Workflow 3: H100 FP8 optimization (FlashAttention-3)
- Verify H100 GPU:
nvidia-smi --query-gpu=name --format=csv
# Should show "H100" or "H800"
- Install flash-attn with FP8 support:
pip install flash-attn --no-build-isolation
- Convert inputs to FP8 and run:
import torch
from flash_attn import flash_attn_func
q = torch.randn(2, 4096, 32, 64, device='cuda', dtype=torch.float16)
k = torch.randn(2, 4096, 32, 64, device='cuda', dtype=torch.float16)
v = torch.randn(2, 4096, 32, 64, device='cuda', dtype=torch.float16)
# Convert to float8_e4m3 (FP8)
q_fp8 = q.to(torch.float8_e4m3fn)
k_fp8 = k.to(torch.float8_e4m3fn)
v_fp8 = v.to(torch.float8_e4m3fn)
# FlashAttention-3 automatically uses FP8 kernels on H100
out = flash_attn_func(q_fp8, k_fp8, v_fp8)
Pitfalls
- ImportError: cannot import flash_attn: Install with no-build-isolation flag:
pip install flash-attn --no-build-isolation. Or install CUDA toolkit first:conda install cuda -c nvidia. - Slower than expected (no speedup): Flash Attention benefits increase with sequence length. <512 tokens: Minimal speedup (10-20%). 512-2K tokens: 2-3x speedup. >2K tokens: 3-4x speedup. Check sequence length is sufficient.
- RuntimeError: CUDA error: Verify GPU supports Flash Attention.
torch.cuda.get_device_capability()should be ≥(7, 5) for Turing+. Volta (V100) is NOT supported. - Accuracy degradation: Check dtype is float16 or bfloat16 (not float32). Flash Attention uses float16/bfloat16 for speed. Float32 not supported.
Verification
- Verify speedup with profiling:
import torch
import torch.nn.functional as F
import torch.utils.benchmark as benchmark
def test_attention(use_flash):
q, k, v = [torch.randn(2, 8, 2048, 64, device='cuda', dtype=torch.float16) for _ in range(3)]
if use_flash:
with torch.backends.cuda.sdp_kernel(enable_flash=True):
return F.scaled_dot_product_attention(q, k, v)
else:
attn = (q @ k.transpose(-2, -1) / 8.0).softmax(dim=-1)
return attn @ v
t_flash = benchmark.Timer(stmt='test_attention(True)', globals=globals())
t_standard = benchmark.Timer(stmt='test_attention(False)', globals=globals())
print(f"Flash: {t_flash.timeit(100).mean:.3f}s")
print(f"Standard: {t_standard.timeit(100).mean:.3f}s")
# Expected: 2-4x speedup for sequences >512 tokens.
- Test accuracy matches baseline:
q, k, v = [torch.randn(1, 8, 512, 64, device='cuda', dtype=torch.float16) for _ in range(3)]
out_flash = F.scaled_dot_product_attention(q, k, v)
attn_weights = torch.softmax(q @ k.transpose(-2, -1) / 8.0, dim=-1)
out_standard = attn_weights @ v
diff = (out_flash - out_standard).abs().max()
print(f"Max difference: {diff:.6f}")
# Should be <1e-3 for float16
References
- Load
references/transformers-integration.mdWHEN you need to enable Flash Attention in HuggingFace Transformers (BERT, GPT, Llama models). - Load
references/benchmarks.mdWHEN you need detailed speed and memory comparisons across GPUs and sequence lengths.