Attention Mechanisms: от Self-Attention до Flash Attention 2

opensourceaillmattentionarchitectureperformanceit
← Back to Blog

Введение: почему 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:

  1. Последовательность: нельзя параллелить по длине последовательности
  2. Затухание градиента: далеко стоящие слова не влияют на результат
  3. Фиксированная размерность скрытого состояния

CNN: локальные паттерны

CNN захватывают локальные паттерны через фильтры, но:

  1. Рецептивное поле растёт медленно (нужно много слоёв для глобального контекста)
  2. Фиксированный размер ядра ограничивает дальние зависимости

Решение: 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.