Flash Attention: ускорить attention в 6 раз без потери точности

opensourceaillmoptimizationgpuit
← Back to Blog

Введение: проблема 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

Ссылки