Flash Attention: ускорить attention в 6 раз без потери точности
Введение: проблема attention
Трансформер attention вычисляется так:
Attention(Q, K, V) = softmax(Q @ K^T / sqrt(d)) @ V
Для seq_len=4096, head_dim=128:
Q @ K^T → (batch, heads, 4096, 4096)
4096 × 4096 × 4 байта = 64 MB на head
32 heads × 64 MB = 2 GB на слой!
Flash Attention — это алгоритм вычисления attention "out-of-core", который минимизирует обращения к DRAM (основной памяти), работая через быстрый SRAM (на чипе GPU).
Оригинальная статья: "FlashAttention: Fast and Memory-Efficient Exact Attention"
Почему обычный attention медленный?
Проблема 1: O(n²) промежуточная память
# Обычный attention (PyTorch)
def naive_attention(Q, K, V):
d_k = Q.shape[-1]
scores = Q @ K.transpose(-2, -3) / (d_k ** 0.5)
weights = torch.softmax(scores, dim=-1)
output = weights @ V
return output
# Для seq_len=4096, 32 heads:
# Промежуточная матрица: 4096² × 4 байта = 64 MB
# С градиентами: ~256 MB на слой
# Для 32 слоёв: ~8 GB!
Проблема 2: DRAM bottleneck
GPU имеет два уровня памяти:
GPU архитектура:
┌─────────────────────────────────┐
│ GPU Chip │
│ SRAM (L1 cache) │
│ ~100 MB/chip, ~5 TB/s │
│ ↕ HBM │
│ HBM (Global memory) │
│ ~80 GB (A100), ~2 TB/s │
└─────────────────────────────────┘
Обычный attention — 6 обращений к HBM:
1. Q из HBM → SRAM
2. K из HBM → SRAM
3. Q@K^T в HBM
4. softmax результат в HBM
5. V из HBM → SRAM
6. результат в HBM
Flash Attention — 2 обращения к HBM:
1. Чанки Q, K, V в SRAM
2. Вычисление + аккумулирование в SRAM
3. Результат ОДИН раз в HBM
Проблема 3: градиенты занимают двойную память
При backpropagation промежуточная матрица сохраняется для градиентов. Для seq_len=4096 это удваивает потребление памяти.
Как работает Flash Attention?
Ключевая идея: тильдо-разбиение (tiling)
Flash Attention разбивает Q, K, V на чанки, которые помещаются в SRAM:
Чанк 128 × 128 × 128 (head_dim):
128 × 128 × 4 байта = 64 KB — влезает в SRAM
Для каждого чанка Q_i:
Для каждого чанка K_j, V_j:
Вычисляем attention(Q_i, K_j, V_j) в SRAM
Аккумулируем результат с softmax normalization
Стадии вычисления
Flash Attention использует два прохода:
Проход 1 (forward):
Для каждого чанка K_j, V_j:
S_ij = Q_i @ K_j^T / sqrt(d)
M_i = max(M_i, max(S_ij)) ← локальный максимум
P_ij = exp(S_ij - M_i) ← нормализация
O_i += P_ij @ V_j ← аккумулирование
Проход 2 (градиенты):
Вычисляем градиенты без сохранения промежуточной матрицы
Используем те же чанки, но в обратном порядке
Почему это точно?
Flash Attention даёт точно такой же результат, как обычный attention (в пределах численной точности float16/bfloat16). Разница только в порядке вычислений, не в формуле.
Flash Attention 2
Flash Attention 2 (2023) оптимизирует ядро CUDA:
FA1: много маленьких thread block'ов
FA2: fewer, larger thread block'ов
better occupancy на GPU A100/H100
Ускорение:
FA1: 2-3x быстрее обычного
FA2: 4-6x быстрее обычного
Сравнение скоростей (seq_len=4096, A100):
Обычный PyTorch: 100% (базовая линия)
FlashAttention 1: 250% (2.5x быстрее)
FlashAttention 2: 500% (5x быстрее)
Flash Attention 3
Flash Attention 3 (2024) — для H100 и новее:
Оптимизации FA3:
1. Tensor Core utilization (WMMA)
2. Fusion of softmax
3. Async scheduling (CUDA graphs)
4. Block-scaled attention
Результат на H100:
В 7 раз быстрее PyTorch attention
1.4x быстрее FA2 на H100
Flash Attention 4
Flash Attention 4 (2025) — для Blackwell и Hopper:
FA4 нововведения:
1. FP8 поддержка
2. Better multi-GPU scaling
3. Optimized for Hopper Tensor Core V4
4. Reduced memory for long context
Для seq_len=131072:
FA4: 12 GB памяти
Обычный: 256 GB памяти
Практическое использование
Установка
pip install flash-attn --no-build-isolation
Или через Docker:
FROM nvidia/cuda:12.1.0-cudnn8-runtime-ubuntu22.04
RUN pip install flash-attn
Использование в коде
import torch
import flash_attn
# Простое использование
q = torch.randn(2, 32, 4096, 128, device='cuda')
k = torch.randn(2, 32, 4096, 128, device='cuda')
v = torch.randn(2, 32, 4096, 128, device='cuda')
# Flash Attention
out = flash_attn.flash_attn_qkvpacked_func(
torch.stack([q, k, v], dim=2),
dropout_p=0.0,
softmax_scale=128 ** -0.5
)
# Или через transformers
from transformers import AutoModelForCausalLM
model = AutoModelForCausalLM.from_pretrained(
"meta-llama/Llama-3-70B",
attn_implementation="flash_attention_2", # <-- ключевая строка
torch_dtype=torch.float16,
).cuda()
Сравнение производительности
import time
seq_lens = [1024, 2048, 4096, 8192]
batch_size = 1
head_dim = 128
n_heads = 32
for seq_len in seq_lens:
q = torch.randn(batch_size, n_heads, seq_len, head_dim, device='cuda')
k = q.clone()
v = q.clone()
# Обычный attention
torch.cuda.synchronize()
start = time.time()
for _ in range(10):
scores = q @ k.transpose(-2, -1) / (head_dim ** 0.5)
weights = torch.softmax(scores, dim=-1)
out1 = weights @ v
torch.cuda.synchronize()
naive_time = (time.time() - start) / 10
# Flash Attention
torch.cuda.synchronize()
start = time.time()
for _ in range(10):
out2 = flash_attn.flash_attn_func(q, k, v)
torch.cuda.synchronize()
flash_time = (time.time() - start) / 10
print(f"seq_len={seq_len:5d}: "
f"naive={naive_time*1000:6.1f}ms "
f"flash={flash_time*1000:5.1f}ms "
f"speedup={naive_time/flash_time:4.1f}x")
Пример вывода:
seq_len= 1024: naive= 12.3ms flash= 3.1ms speedup= 3.9x
seq_len= 2048: naive= 48.7ms flash= 10.2ms speedup= 4.8x
seq_len= 4096: naive= 195.2ms flash= 38.5ms speedup= 5.1x
seq_len= 8192: naive= 780.1ms flash= 142.3ms speedup= 5.5x
Когда Flash Attention НЕ работает
Не работает:
❌ CPU (нет CUDA)
❌ MPS (Apple Silicon — используйте native attention)
❌ seq_len < 64 (overhead тилинга больше выгоды)
❌ Очень маленькие модели (attention — не bottleneck)
Работает лучше всего:
✅ seq_len > 1024
✅ Большие модели (7B+)
✅ Обучение (где градиенты — половина памяти)
✅ Long-context модели (32K+)
Альтернативы
PagedAttention (vLLM)
PagedAttention — для inference serving:
- Разбивает KV cache на страницы
- Устраняет fragmentation памяти
- 2-4x больше batch size чем конкуренты
Используется в vLLM, FlexGen
Sparse Attention
Sparse Attention — только часть attention:
- Sliding window: attention только в окне 256 токенов
- Global tokens: несколько токенов attention ко всем
- Пример: Longformer, BigBird
Сложность: O(n × w) вместо O(n²)
где w — размер окна
Linear Attention
Linear Attention — переформулировка:
Attention(Q, K, V) = φ(Q) @ φ(K)^T @ V
где φ — преобразование (например, exp(Q))
Сложность: O(n × d) вместо O(n²)
Пример: Performer, Linformer
Системные требования
Flash Attention 2:
- GPU: NVIDIA с compute capability ≥ 8.0 (Ampere+)
- CUDA: 11.6+
- Память: минимум 8 GB (для 7B модели)
- Лучше: 24 GB+ (A10/A100/H100)
Flash Attention 3:
- GPU: NVIDIA Hopper (H100) или новее
- CUDA: 12.1+
- FP8 поддержка (H100+)
Flash Attention 4:
- GPU: NVIDIA Blackwell (B100/B200) или Hopper
- CUDA: 12.4+
- FP8 + FP4 поддержка
Итоги
Flash Attention — must-have для:
- Обучения больших моделей (экономия 2-4x памяти)
- Long-context inference (4096+ токенов)
- Serving больших batch'ей (больше параллелизма)
Рекомендации:
1. Всегда используйте flash-attn для seq_len > 1024
2. Для обучения — обязательно (экономия памяти)
3. Для inference — зависит от serving framework
4. Для Apple Silicon — используйте native attention