Attention Mechanisms: от Self-Attention до Flash Attention 2
Введение: почему Attention — это сердце трансформера
Каждый современный LLM — от GPT до Llama — построен на механизме Self-Attention. Без него модель не может связать слова в контексте, запомнить информацию из начала длинного текста или понять, какие слова в предложении важнее других.
В этой статье разберём:
- Что такое Self-Attention и как работает формула QKV
- Почему наивная реализация O(n²) по памяти и вычислениям
- Как работает Multi-Head Attention и почему "heads" — это не просто параллелизм
- Что такое Flash Attention и почему это прорыв 2022-2024 годов
- Как оптимизировать Attention для длинных контекстов (Linear Attention, Sparse Attention)
- Практические советы для выбора и настройки
Проблема: почему RNN и CNN не справились?
RNN: последовательная обработка
До трансформеров основной архитектурой для NLP были RNN (LSTM/GRU):
Токен 1 → RNN → hidden_1
Токен 2 → RNN → hidden_2
Токен 3 → RNN → hidden_3
Проблемы RNN:
- Последовательность: нельзя параллелить по длине последовательности
- Затухание градиента: далеко стоящие слова не влияют на результат
- Фиксированная размерность скрытого состояния
CNN: локальные паттерны
CNN захватывают локальные паттерны через фильтры, но:
- Рецептивное поле растёт медленно (нужно много слоёв для глобального контекста)
- Фиксированный размер ядра ограничивает дальние зависимости
Решение: Self-Attention
Self-Attention позволяет каждому токену взаимодействовать с каждым за один шаг:
Токен "Кот" → Attention → видит "сидел", "на", "ковре"
Токен "сидел" → Attention → видит "Кот", "на", "ковре"
Токен "ковре" → Attention → видит "Кот", "сидел", "на"
Все вычисления параллельны.
Self-Attention: формула и интуиция
QKV: Query, Key, Value
Каждый токен имеет три вектора:
- Query (Q): "что я ищу?"
- Key (K): "что я содержу?"
- Value (V): "что я передаю?"
import torch
import torch.nn as nn
import torch.nn.functional as F
class SelfAttention(nn.Module):
def __init__(self, d_model=512, d_head=64):
super().__init__()
self.d_model = d_model
self.d_head = d_head
self.n_heads = d_model // d_head
# Весовые матрицы для Q, K, V
self.W_q = nn.Linear(d_model, d_model)
self.W_k = nn.Linear(d_model, d_model)
self.W_v = nn.Linear(d_model, d_model)
# Output projection
self.W_o = nn.Linear(d_model, d_model)
def forward(self, x):
# x: [batch_size, seq_len, d_model]
batch_size, seq_len, _ = x.shape
# Проекция в Q, K, V
Q = self.W_q(x) # [batch, seq_len, d_model]
K = self.W_k(x) # [batch, seq_len, d_model]
V = self.W_v(x) # [batch, seq_len, d_model]
# Разделяем на heads: [batch, n_heads, seq_len, d_head]
Q = Q.view(batch_size, seq_len, self.n_heads, self.d_head).transpose(1, 2)
K = K.view(batch_size, seq_len, self.n_heads, self.d_head).transpose(1, 2)
V = V.view(batch_size, seq_len, self.n_heads, self.d_head).transpose(1, 2)
# Attention scores: Q @ K^T
scores = torch.matmul(Q, K.transpose(-2, -1)) / (self.d_head ** 0.5)
# [batch, n_heads, seq_len, seq_len]
# Softmax по строкам (каждый токен attention-weights ко всем)
attn_weights = F.softmax(scores, dim=-1)
# Weighted sum Values
output = torch.matmul(attn_weights, V)
# [batch, n_heads, seq_len, d_head]
# Concatenate heads
output = output.transpose(1, 2).contiguous()
output = output.view(batch_size, seq_len, self.d_model)
# Output projection
output = self.W_o(output)
return output
Пошаговая интуиция
Предложение: "Кот сидел на ковре"
Для токена "сидел":
Q_сидел = [0.2, -0.5, 0.8, ...] # "ищу подлежащее"
K_кот = [0.1, -0.4, 0.9, ...] # similarity = 0.92 ← ВЫСОКОЕ
K_сидел = [0.3, -0.6, 0.7, ...] # similarity = 0.78
K_на = [-0.1, -0.3, 0.5, ...] # similarity = 0.45
K_ковре = [0.0, -0.2, 0.6, ...] # similarity = 0.51
Attention weights: [0.42, 0.28, 0.12, 0.18]
Output = 0.42 * V_кот + 0.28 * V_сидел + 0.12 * V_на + 0.18 * V_ковре
Multi-Head Attention: почему "heads"?
Идея: разные heads = разные "ракурсы"
Один head может учиться захватывать синтаксические связи, другой — семантические, третий — позиционные:
class MultiHeadAttention(nn.Module):
def __init__(self, d_model=512, n_heads=8):
super().__init__()
self.n_heads = n_heads
self.d_head = d_model // n_heads
# Один общий линейный слой для всех heads
self.W_out = nn.Linear(d_model, d_model)
def forward(self, x, mask=None):
batch_size, seq_len, _ = x.shape
# Q, K, V: [batch, seq_len, d_model]
Q, K, V = self.W_q(x), self.W_k(x), self.W_v(x)
# Разделяем на heads
Q = Q.view(batch_size, seq_len, self.n_heads, self.d_head).transpose(1, 2)
K = K.view(batch_size, seq_len, self.n_heads, self.d_head).transpose(1, 2)
V = V.view(batch_size, seq_len, self.n_heads, self.d_head).transpose(1, 2)
# Scaled Dot-Product Attention
scores = torch.matmul(Q, K.transpose(-2, -1)) / (self.d_head ** 0.5)
# Mask: для decoder-а (causal mask)
if mask is not None:
scores = scores.masked_fill(mask == 0, -1e9)
attn_weights = F.softmax(scores, dim=-1)
output = torch.matmul(attn_weights, V)
# Concatenate heads
output = output.transpose(1, 2).contiguous()
output = output.view(batch_size, seq_len, self.d_model)
return self.W_out(output)
Что изучают разные heads?
Исследования (OpenAI, 2020) показывают:
Head 0: отслеживает согласование подлежащего и сказуемого
Head 1: захватывает притяжательные местоимения
Head 2: позиционный head (смотрит только на соседние токены)
Head 3: связывает отрицания с их областью
Head 4: семантическая роль (кто делает что)
Head 5: attention на начало последовательности
Head 6: attention на конец последовательности
Head 7: общий контекстный head (diffuse attention)
Causal Mask (Decoder-attention)
В decoder-е каждый токен видит только предыдущие:
Seq: [A, B, C, D, E]
Без mask:
A → [A, B, C, D, E] ← видит всё
B → [A, B, C, D, E] ← видит всё
С causal mask:
A → [A, 0, 0, 0, 0] ← видит только себя
B → [A, B, 0, 0, 0] ← видит A и себя
C → [A, B, C, 0, 0] ← видит A, B, себя
D → [A, B, C, D, 0] ← видит A, B, C, себя
E → [A, B, C, D, E] ← видит всё (это последний токен)
# Causal mask для decoder
def causal_mask(seq_len):
mask = torch.tril(torch.ones(seq_len, seq_len))
# mask[i, j] = 1 если j <= i, иначе 0
return mask
Сложность Attention: O(n²) проблема
Вычислительная сложность
Self-Attention для последовательности длины n:
1. Q @ K^T: O(n² × d) ← матричное умножение n×d × d×n = n²×d
2. Softmax: O(n²) ← для каждой пары (i, j)
3. Attn @ V: O(n² × d) ← n×n × n×d = n²×d
Итого: O(n² × d)
Память: O(n²) для attention matrix
Что это значит на практике?
Модель: Llama 3 8B, d_model = 4096
seq_len = 4K:
Attention matrix: 4096 × 4096 × 4 bytes = 64 MB
Compute: ~2 × 4096² × 4096 ≈ 110 GFLOPs (только attention!)
seq_len = 32K:
Attention matrix: 32768 × 32768 × 4 bytes = 4 GB
Compute: ~2 × 32768² × 4096 ≈ 8.6 TFLOPs
seq_len = 128K:
Attention matrix: 128K × 128K × 4 bytes = 64 GB ← не влезает в VRAM!
Compute: ~2 × 128K² × 4096 ≈ 136 TFLOPs
Проблема: при удвоении длины последовательности стоимость растёт в 4 раза.
Flash Attention: прорыв
Основная идея
Наивная реализация Attention:
1. Q @ K^T → scores [n×n] ← записываем в память
2. scores - max → subtract ← читаем scores
3. scores.exp() → exp ← читаем scores
4. sum(exp) → denominator ← читаем exp
5. exp / sum → output ← читаем exp и sum
6. output @ V → result ← читаем output, V
Для seq_len=32K attention matrix = 4 GB. Чтение/запись этой матрицы — узкое место.
Flash Attention (Dao et al., 2022):
- Разбивает input на tiles (блоки)
- Вычисляет attention по тайлам, не храня полную матрицу
- Использует tiled softmax (online softmax)
- Всё помещается в SRAM (очень быстрая память на чипе)
# Псевдокод Flash Attention
def flash_attention(Q, K, V):
# Q, K, V: [batch, n_heads, seq_len, d_head]
# d_head обычно 64-256 (маленькие блоки)
# Разбиваем на тайлы (например, 64×64)
tile_size = 64
for tile_i in range(0, seq_len, tile_size): # по строкам
for tile_j in range(0, seq_len, tile_size): # по столбцам
# Загружаем тайл K и V в SRAM
K_tile = K[:, :, tile_j:tile_j+tile_size, :]
V_tile = V[:, :, tile_j:tile_j+tile_size, :]
# Вычисляем attention для этого тайла
Q_tile = Q[:, :, tile_i:tile_i+tile_size, :]
scores = Q_tile @ K_tile.transpose(-2, -1) # маленький!
# Online softmax для тайла
# ... (с учётом предыдущих тайлов)
# Accumulate в output
output[:, :, tile_i:tile_i+tile_size, :] += attn_tile @ V_tile
Результат
seq_len = 32K, d_head = 128:
Наивная реализация:
VRAM для attention matrix: 4 GB
Время: 2.5 сек
Память: 4 GB (matrix) + 2 GB (activations) = 6 GB
Flash Attention 2:
VRAM для attention matrix: ~10 MB (tile)
Время: 0.8 сек (в 3 раза быстрее!)
Память: 10 MB (tile) + 2 GB (activations) = 2 GB
Flash Attention 3 (2024)
- Ещё больше оптимизаций: register-level tiling, fused softmax
- 1.3x быстрее чем FA2
- Только для NVIDIA H100+ (sm90)
Оптимизации для длинных контекстов
1. Linear Attention
Заменяем softmax на positive kernel:
def linear_attention(Q, K, V):
# Normalizing functions
Q = F.relu(Q)
K = F.relu(K)
# (Q @ K^T) @ V → Q @ (K^T @ V) ← меняем порядок!
KV = K.transpose(-2, -1) @ V # [batch, n_heads, d_head, d_head]
output = Q @ KV # [batch, n_heads, seq_len, d_head]
return output
Сложность: O(n × d²) вместо O(n² × d) Минус: качество обычно чуть хуже softmax attention
2. Sparse Attention
Каждый токен attention только к ограниченному числу других:
Dense Attention:
Токен 1 → все 32K токенов
Sparse Attention (sliding window 256):
Токен 1 → токены [1-128, 2-129, ..., 256]
Токен 1000 → токены [872-1000, 1001-1128]
Long-range tokens (every 128):
Токен 1 → также токены 1, 129, 257, 385, ...
# Sliding window attention
def sparse_attention(Q, K, V, window_size=256):
batch, heads, seq_len, d_head = Q.shape
# Local attention (sliding window)
local_scores = torch.matmul(Q, K.transpose(-2, -1)) / (d_head ** 0.5)
# Mask: только window_size соседей
mask = torch.ones(seq_len, seq_len, device=Q.device)
mask = torch.triu(mask, diagonal=1-window_size)
mask = torch.tril(mask, diagonal=window_size-1)
local_scores = local_scores.masked_fill(mask == 0, -1e9)
local_attn = F.softmax(local_scores, dim=-1)
return torch.matmul(local_attn, V)
Используется в: Longformer, BigBird, MPT-30B
3. KV Cache Compression
PagedAttention (vLLM)
Группирует токены в блоки и управляет KV cache как операционная система — памятью:
Обычный KV Cache:
Токен 1: K[0:64], V[0:64] → блок 0
Токен 2: K[64:128], V[64:128] → блок 1
...
Токен 32768: K[2097184:2097248] → блок 32767
Проблема: префиксы одинаковых промптов занимают отдельную память
PagedAttention:
Блок 0: K[0:64], V[0:64] → префикс "system: "
Блок 1: K[64:128], V[64:128] → префикс "system: "
Блок 2: K[128:192], V[128:192]→ user message 1
Блок 3: K[192:256], V[192:256]→ user message 2
Блоки 0,1 разделяются между запросами с одинаковым system prompt!
KV Cache Quantization
FP16 KV Cache:
Каждый K, V элемент: 2 байта
Для 32K seq_len, 32 layers, 32 heads, d_head=128:
32 × 32 × 128 × 32768 × 2 bytes = 8.4 GB
INT8 KV Cache:
Каждый K, V элемент: 1 байт
32 × 32 × 128 × 32768 × 1 byte = 4.2 GB
Сжатие: 2x без заметной потери качества
Rotary Embeddings (RoPE) и Attention
RoPE — способ добавить информацию о позиции в attention:
def apply_rope(x, cos, sin):
"""
Применяет Rotary Positional Embeddings к Q и K.
x: [batch, n_heads, seq_len, d_head]
cos, sin: positional embeddings
"""
x_reshaped = x.view(*x.shape[:-1], -1, 2) # [batch, n_heads, seq_len, d_head//2, 2]
x1 = x_reshaped[..., 0]
x2 = x_reshaped[..., 1]
sin = sin.unsqueeze(1).unsqueeze(1)
cos = cos.unsqueeze(1).unsqueeze(1)
o1 = x1 * cos - x2 * sin
o2 = x1 * sin + x2 * cos
return o1.view(*x.shape), o2.view(*x.shape)
RoPE встраивается до вычисления attention scores, позволяя attention weights зависеть от относительных позиций токенов.
Практические советы
Как выбрать Attention implementation?
NVIDIA GPU (Ampere+):
→ Flash Attention 2 (через torch.backends.cuda.sdp_kernel)
NVIDIA H100+:
→ Flash Attention 3 (если поддерживает)
AMD GPU / CPU:
→ PyTorch native scaled_dot_product_attention
Длинный контекст (>32K):
→ Sparse Attention или Linear Attention
Памяти впритык:
→ Flash Attention (минимальный memory footprint)
Production (vLLM):
→ PagedAttention (автоматически)
Включение Flash Attention в PyTorch
import torch
# Включение Flash Attention 2
with torch.backends.cuda.sdp_kernel(
enable_flash=True, # Flash Attention 2
enable_math=False, # Отключаем PyTorch fallback
enable_mem_efficient=False # Отключаем memory efficient
):
output = attn(Q, K, V)
Проверка, что Flash Attention работает
import torch
# Проверяем поддержку
print(torch.cuda.is_available()) # True
print(torch.cuda.get_device_capability()) # (8, 0) для A100, (8, 6) для 3090
# Для sm 8.0+ (Ampere) и выше — Flash Attention 2 поддерживается
major, minor = torch.cuda.get_device_capability()
if major >= 8:
print("Flash Attention 2 supported!")
Итоги
- Self-Attention — основа всех современных LLM, позволяет каждому токену взаимодействовать с каждым
- Формула: Attention(Q, K, V) = softmax(QK^T / √d) V
- Multi-Head Attention: несколько attention "ракурсов" параллельно
- Сложность O(n²) по памяти и вычислениям — главный bottleneck
- Flash Attention: tiled computation, O(n) памяти вместо O(n²), в 2-3x быстрее
- Для длинных контекстов: Sparse Attention, Linear Attention, KV Cache compression
- В production: vLLM с PagedAttention для эффективного управления памятью
Attention — не просто механизм, а целая экосистема оптимизаций. Понимание того, как он работает, критично для эффективного запуска LLM.