Перейти к основному содержимому

Оптимизация Flash Attention

Оптимизирует внимание трансформеров с помощью Flash Attention, обеспечивая ускорение в 2–4 раза и снижение потребления памяти в 10–20 раз. Используйте при обучении или запуске трансформеров с длинными последовательностями (>512 токенов), при возникновении проблем с памятью GPU при работе с вниманием или когда требуется более быстрый инференс. Поддерживает нативный SDPA PyTorch, библиотеку flash-attn, H100 FP8 и скользящее окно внимания.

Метаданные навыка​

ИсточникОпционально — установка: vibeos skills install official/mlops/flash-attention
Путьoptional-skills/mlops/flash-attention
Версия1.0.0
АвторOrchestra Research
ЛицензияMIT
Зависимостиflash-attn, torch, transformers
Платформыlinux, macos
ТегиОптимизация, Flash Attention, Оптимизация внимания, Эффективность памяти, Оптимизация скорости, Длинный контекст, PyTorch, SDPA, H100, FP8, Трансформеры

Справочник: полный SKILL.md​

к сведению

Ниже приведено полное описание навыка, которое VibeOS загружает при его активации. Агент видит эти инструкции, когда навык активен.

Flash Attention — быстрое эффективное по памяти внимание

Быстрый старт​

Flash Attention обеспечивает ускорение в 2–4 раза и снижение потребления памяти в 10–20 раз для внимания трансформеров за счет IO-ориентированной тайлинга и рекомпьютации.

Нативный PyTorch (проще всего, PyTorch 2.2+):

import torch
import torch.nn.functional as F

q = torch.randn(2, 8, 512, 64, device='cuda', dtype=torch.float16) # [batch, heads, seq, dim]
k = torch.randn(2, 8, 512, 64, device='cuda', dtype=torch.float16)
v = torch.randn(2, 8, 512, 64, device='cuda', dtype=torch.float16)

# Автоматически использует Flash Attention, если доступно
out = F.scaled_dot_product_attention(q, k, v)

Библиотека flash-attn (больше возможностей):

pip install flash-attn --no-build-isolation
from flash_attn import flash_attn_func

# q, k, v: [batch, seqlen, nheads, headdim]
out = flash_attn_func(q, k, v, dropout_p=0.0, causal=True)

Типовые рабочие процессы​

Рабочий процесс 1: Включение в существующей модели PyTorch​

Скопируйте этот чек-лист:

Интеграция Flash Attention:
- [ ] Шаг 1: Проверьте версию PyTorch (≥2.2)
- [ ] Шаг 2: Включите бэкенд Flash Attention
- [ ] Шаг 3: Проверьте ускорение с помощью профилирования
- [ ] Шаг 4: Убедитесь, что точность совпадает с базовой

Шаг 1: Проверьте версию PyTorch

python -c "import torch; print(torch.__version__)"
# Должно быть ≥2.2.0

Если <2.2, обновите:

pip install --upgrade torch

Шаг 2: Включите бэкенд Flash Attention

Замените стандартное внимание:

# До (стандартное внимание)
attn_weights = torch.softmax(q @ k.transpose(-2, -1) / math.sqrt(d_k), dim=-1)
out = attn_weights @ v

# После (Flash Attention)
import torch.nn.functional as F
out = F.scaled_dot_product_attention(q, k, v, attn_mask=mask)

Принудительное использование бэкенда Flash Attention:

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)

Шаг 3: Проверьте ускорение с помощью профилирования

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"Стандартное: {t_standard.timeit(100).mean:.3f}s")

Ожидаемый результат: ускорение в 2–4 раза для последовательностей длиннее 512 токенов.

Шаг 4: Убедитесь, что точность совпадает с базовой

# Сравнение выходов
q, k, v = [torch.randn(1, 8, 512, 64, device='cuda', dtype=torch.float16) for _ in range(3)]

# Flash Attention
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"Максимальная разница: {diff:.6f}")
# Должно быть <1e-3 для float16

Рабочий процесс 2: Использование библиотеки flash-attn для продвинутых функций​

Для multi-query attention, скользящего окна или H100 FP8.

Скопируйте этот чек-лист:

Настройка библиотеки flash-attn:
- [ ] Шаг 1: Установите библиотеку flash-attn
- [ ] Шаг 2: Измените код внимания
- [ ] Шаг 3: Включите продвинутые функции
- [ ] Шаг 4: Проведите бенчмарк производительности

Шаг 1: Установите библиотеку flash-attn

# NVIDIA GPU (CUDA 12.0+)
pip install flash-attn --no-build-isolation

# Проверка установки
python -c "from flash_attn import flash_attn_func; print('Успешно')"

Шаг 2: Измените код внимания

from flash_attn import flash_attn_func

# Вход: [batch_size, seq_len, num_heads, head_dim]
# Транспонируйте из [batch, heads, seq, dim] при необходимости
q = q.transpose(1, 2) # [batch, seq, heads, dim]
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), # Без скользящего окна
softmax_scale=None # Автомасштабирование
)

out = out.transpose(1, 2) # Обратно в [batch, heads, seq, dim]

Шаг 3: Включите продвинутые функции

Multi-query attention (общие K/V для всех голов):

from flash_attn import flash_attn_func

# q: [batch, seq, num_q_heads, dim]
# k, v: [batch, seq, num_kv_heads, dim] # Меньше KV-голов
out = flash_attn_func(q, k, v) # Автоматически обрабатывает MQA

Скользящее окно внимания (локальное внимание):

# Внимание только к окну из 256 токенов до/после
out = flash_attn_func(
q, k, v,
window_size=(256, 256), # Окно (слева, справа)
causal=True
)

Шаг 4: Проведите бенчмарк производительности

import torch
from flash_attn import flash_attn_func
import time

q, k, v = [torch.randn(4, 4096, 32, 64, device='cuda', dtype=torch.float16) for _ in range(3)]

# Прогрев
for _ in range(10):
_ = flash_attn_func(q, k, v)

# Бенчмарк
torch.cuda.synchronize()
start = time.time()
for _ in range(100):
out = flash_attn_func(q, k, v)
torch.cuda.synchronize()
end = time.time()

print(f"Время на итерацию: {(end-start)/100*1000:.2f}ms")
print(f"Выделено памяти: {torch.cuda.max_memory_allocated()/1e9:.2f}GB")

Рабочий процесс 3: Оптимизация H100 FP8 (FlashAttention-3)​

Для максимальной производительности на GPU H100.

Настройка FP8:
- [ ] Шаг 1: Убедитесь, что доступен GPU H100
- [ ] Шаг 2: Установите flash-attn с поддержкой FP8
- [ ] Шаг 3: Преобразуйте входные данные в FP8
- [ ] Шаг 4: Запустите с FP8 вниманием

Шаг 1: Убедитесь, что доступен GPU H100

nvidia-smi --query-gpu=name --format=csv
# Должно показывать «H100» или «H800»

Шаг 2: Установите flash-attn с поддержкой FP8

pip install flash-attn --no-build-isolation
# Поддержка FP8 включена для H100

Шаг 3: Преобразуйте входные данные в FP8

import torch

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)

# Преобразование в float8_e4m3 (FP8)
q_fp8 = q.to(torch.float8_e4m3fn)
k_fp8 = k.to(torch.float8_e4m3fn)
v_fp8 = v.to(torch.float8_e4m3fn)

Шаг 4: Запустите с FP8 вниманием

from flash_attn import flash_attn_func

# FlashAttention-3 автоматически использует FP8-ядра на H100
out = flash_attn_func(q_fp8, k_fp8, v_fp8)
# Результат: ~1.2 PFLOPS, в 1.5–2 раза быстрее FP16

Когда использовать и альтернативы​

Используйте Flash Attention, когда:

  • Обучаете трансформеры с последовательностями >512 токенов
  • Запускаете инференс с длинным контекстом (>2K токенов)
  • Память GPU ограничена (OOM при стандартном внимании)
  • Нужно ускорение в 2–4 раза без потери точности
  • Используете PyTorch 2.2+ или можете установить flash-attn

Вместо этого используйте альтернативы:

  • Стандартное внимание: Последовательности <256 токенов (оверхед не оправдан)
  • xFormers: Нужны другие варианты внимания (не только скорость)
  • Memory-efficient attention: Инференс на CPU (Flash Attention требует GPU)

Частые проблемы​

Проблема: ImportError: cannot import flash_attn

Установите с флагом no-build-isolation:

pip install flash-attn --no-build-isolation

Или сначала установите CUDA toolkit:

conda install cuda -c nvidia
pip install flash-attn --no-build-isolation

Проблема: Медленнее, чем ожидалось (нет ускорения)

Преимущества Flash Attention возрастают с длиной последовательности:

  • <512 токенов: Минимальное ускорение (10–20%)
  • 512–2K токенов: Ускорение в 2–3 раза
  • 2K токенов: Ускорение в 3–4 раза

Проверьте, достаточна ли длина последовательности.

Проблема: RuntimeError: CUDA error

Убедитесь, что GPU поддерживает Flash Attention:

import torch
print(torch.cuda.get_device_capability())
# Должно быть ≥(7, 5) для Turing+

Flash Attention требует:

  • Ampere (A100, A10): ✅ Полная поддержка
  • Turing (T4): ✅ Поддерживается
  • Volta (V100): ❌ Не поддерживается

Проблема: Снижение точности

Проверьте, что dtype — float16 или bfloat16 (не float32):

q = q.to(torch.float16)  # Или torch.bfloat16

Flash Attention использует float16/bfloat16 для скорости. Float32 не поддерживается.

Продвинутые темы​

Интеграция с HuggingFace Transformers: См. references/transformers-integration.md для включения Flash Attention в модели BERT, GPT, Llama.

Бенчмарки производительности: См. references/benchmarks.md для подробных сравнений скорости и памяти на разных GPU и длинах последовательностей.

Требования к оборудованию​

  • GPU: NVIDIA Ampere+ (A100, A10, A30) или AMD MI200+
  • VRAM: Как для стандартного внимания (Flash Attention не увеличивает память)
  • CUDA: 12.0+ (минимум 11.8)
  • PyTorch: 2.2+ для нативной поддержки

Не поддерживается: V100 (Volta), инференс на CPU

Ресурсы​