Оптимизация 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
Ресурсы
- Статья: «FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness» (NeurIPS 2022)
- Статья: «FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning» (ICLR 2024)
- Блог: https://tridao.me/blog/2024/flash3/
- GitHub: https://github.com/Dao-AILab/flash-attention
- Документация PyTorch: https://pytorch.org/docs/stable/generated/torch.nn.functional.scaled_dot_product_attention.html