Моделі та платформи ШІ

Flash Attention: революційна техніка для підвищення ефективності трансформерів

mm
Додайте Unite.AI до бажаних джерел у Google
div]:bg-bg-300 [&_pre]:-mr-4 md:[&_pre]:-mr-9″>
_*]:min-w-0″>

Як трансформерні моделі зростають у розмірі та складності, вони стикаються з значними проблемами щодо обчислювальної ефективності та використання пам’яті, особливо при роботі з довгими послідовностями. Flash Attention – це техніка оптимізації, яка обіцяє революціонізувати спосіб реалізації та масштабування механізмів уваги в моделях трансформерів.

У цьому комплексному керівництві ми глибоко вивчимо Flash Attention, досліджуючи його основні концепції, деталі реалізації та глибокий вплив на область машинного навчання.

Проблема: увага дорога

Перед тим, як ми зануримося у рішення, давайте спочатку зрозуміємо проблему, яку Flash Attention намагається вирішити. Механізм уваги, хоча й потужний, має значну обчислювальну вартість, особливо для довгих послідовностей.

Стандартна увага: швидкий огляд

Стандартний механізм уваги в моделях трансформерів можна підсумувати наступною формулою:

Увага(Q, K, V) = softmax(QK^T / √d) V

Де Q, K і V – матриці запиту, ключа та значення відповідно, а d – розмірність векторів ключа.

Хоча ця формула елегантна, її реалізація призводить до кількох неефективностей:

  1. Бутільня з пам’яттю: проміжна матриця уваги (QK^T) має розмір N x N, де N – довжина послідовності. Для довгих послідовностей це може швидко виснажити доступну пам’ять GPU.
  2. Надмірний доступ до пам’яті: у стандартних реалізаціях матриця уваги обчислюється, зберігається в пам’яті з високою пропускною здатністю (HBM) і потім читается для операції softmax. Цей надмірний доступ до пам’яті є основною бутільнею.
  3. Недооцінка обчислювальної потужності GPU: сучасні GPU мають значно більше обчислювальної потужності (FLOPS), ніж пропускна здатність пам’яті. Стандартна реалізація уваги обмежена пам’яттю, залишаючи значну частину обчислювальної потужності GPU невикористаною.

Давайте проілюструємо це простим фрагментом коду Python, який показує стандартну реалізацію уваги:

</pre>
import torch

<p>def standard_attention(Q, K, V):
# Q, K, V shape: (batch_size, seq_len, d_model)
d_k = K.size(-1)
scores = torch.matmul(Q, K.transpose(-2, -1)) / torch.sqrt(torch.tensor(d_k))
attention_weights = torch.softmax(scores, dim=-1)
return torch.matmul(attention_weights, V)</p>

Ця реалізація, хоча й прямолінійна, страждає від неефективностей, згаданих вище. Тензор scores, який має форму (batch_size, seq_len, seq_len), може стати надзвичайно великим для довгих послідовностей.

Введення Flash Attention

Flash Attention, представлений Tri Dao і колегами у їхній роботі 2022 року, є підходом до обчислення уваги, який суттєво зменшує використання пам’яті та покращує обчислювальну ефективність. Ключові ідеї, що стоять за Flash Attention, полягають у:

  1. Тайлінг: розбиття великої матриці уваги на менші тайли, які поміщаються у швидку пам’ять SRAM.
  2. Переобчислення: замість збереження всієї матриці уваги, переобчислення частини її під час зворотнього проходу.
  3. IO-інформаційна реалізація: оптимізація алгоритму для мінімалізації руху даних між різними рівнями ієрархії пам’яті GPU.

Алгоритм Flash Attention

У своєму ядрі Flash Attention переосмислює спосіб обчислення механізму уваги. Замість обчислення всієї матриці уваги одразу, він обробляє її блоками, використовуючи ієрархію пам’яті сучасних GPU.

Ось високорівневий огляд алгоритму:

  1. Вхід: матриці Q, K, V у пам’яті HBM (High Bandwidth Memory) і на кристалі SRAM розміром M.
  2. Розміри блоків обчислюються на основі доступної SRAM.
  3. Ініціалізація матриці виводу O та допоміжних векторів l і m.
  4. Алгоритм розбиває вхідні матриці на блоки для розміщення у SRAM.
  5. Два вкладені цикли обробляють ці блоки:
    • Зовнішній цикл завантажує блоки K і V
    • Внутрішній цикл завантажує блоки Q і виконує обчислення
  6. Обчислення на кристалі включають матричне множення, softmax і обчислення виводу.
  7. Результати записуються назад у HBM після обробки кожного блоку.

Цей блоковий обчислювальний процес дозволяє Flash Attention підтримувати значно менший відбиток пам’яті, продовжуючи обчислювати точну увагу.

Математика за Flash Attention

Ключ до того, щоб зробити Flash Attention працездатним, полягає у математичному трюці, який дозволяє обчислювати softmax у блоковому порядку. Робота вводить дві ключові формули:

  1. Декомпозиція softmax:
    softmax(x) = exp(x - m) / Σexp(x - m)

    де m – максимальне значення у x.

  2. Об’єднання softmax:
    softmax(x ∪ y) = softmax(softmax(x) * e^(m_x - m), softmax(y) * e^(m_y - m))

    де m = max(m_x, m_y)

Ці формули дозволяють Flash Attention обчислювати часткові результати softmax для кожного блоку та потім правильно об’єднувати їх для отримання кінцевого результату.

Деталі реалізації

Давайте зануримося у спрощену реалізацію Flash Attention, щоб проілюструвати його основні концепції:

import torch

<p>def flash_attention(Q, K, V, block_size=256):
batch_size, seq_len, d_model = Q.shape</p>

<p># Ініціалізація матриці виводу та допоміжних векторів
O = torch.zeros_like(Q)
L = torch.zeros((batch_size, seq_len, 1))
M = torch.full((batch_size, seq_len, 1), float(&#039;-inf&#039;))</p>

<p>for i in range(0, seq_len, block_size):
Q_block = Q[:, i:i+block_size, :]</p>

<p>for j in range(0, seq_len, block_size):
K_block = K[:, j:j+block_size, :]
V_block = V[:, j:j+block_size, :]</p>

<p># Обчислення оцінок уваги для цього блоку
S_block = torch.matmul(Q_block, K_block.transpose(-2, -1)) / (d_model ** 0.5)</p>

<p># Оновлення максимального значення
M_new = torch.maximum(M[:, i:i+block_size], S_block.max(dim=-1, keepdim=True).values)</p>

<p># Обчислення експонент
exp_S = torch.exp(S_block - M_new)
exp_M_diff = torch.exp(M[:, i:i+block_size] - M_new)</p>

<p># Оновлення суми
L_new = exp_M_diff * L[:, i:i+block_size] + exp_S.sum(dim=-1, keepdim=True)</p>

<p># Обчислення виводу для цього блоку
O[:, i:i+block_size] = (
exp_M_diff * O[:, i:i+block_size] +
torch.matmul(exp_S, V_block)
) / L_new</p>

<p># Оновлення допоміжних векторів
L[:, i:i+block_size] = L_new
M[:, i:i+block_size] = M_new</p>

return O

Ця реалізація, хоча й спрощена, захоплює суть Flash Attention. Вона обробляє вхідні дані блоками, підтримуючи допоміжні вектори (M і L) для правильного обчислення softmax через всі блоки.

Вплив Flash Attention

Введення Flash Attention мало глибокий вплив на область машинного навчання, особливо для великих мовних моделей та довгих контекстів. Деякі ключові переваги включають:

  1. Зменшення використання пам’яті: Flash Attention зменшує складність пам’яті з O(N^2) до O(N), де N – довжина послідовності. Це дозволяє обробляти значно довші послідовності з тією ж апаратурою.
  2. Покращення швидкості: Зміна руху даних та краще використання обчислювальних можливостей GPU дозволяє досягти значних прискорень. Автори повідомляють про прискорення до 3 разів під час навчання для GPT-2 порівняно зі стандартними реалізаціями.
  3. Точне обчислення: На відміну від деяких інших оптимізацій уваги, Flash Attention обчислює точну увагу, а не наближення.
  4. Масштабованість: Зменшений відбиток пам’яті дозволяє масштабуватися до значно довших послідовностей, потенційно до мільйонів токенів.

Практичний вплив

Вплив Flash Attention простягається за межі академічних досліджень. Його швидко прийняли у багатьох популярних бібліотеках машинного навчання та моделях:

  • Бібліотека Hugging Face Transformers: Популярна бібліотека Transformers включила Flash Attention, дозволяючи користувачам легко використати його переваги.
  • GPT-4 і далі: Хоча не підтверджено, існує припущення, що передові мовні моделі, такі як GPT-4, можуть використовувати техніки, подібні до Flash Attention, для обробки довгих контекстів.
  • Моделі довгих контекстів: Flash Attention дозволив створити нове покоління моделей, здатних обробляти надзвичайно довгі контексти, такі як моделі, які можуть обробляти цілі книги або довгі відео.

FlashAttention: недавні розробки

Стандартна увага проти Flash Attention

Стандартна увага проти Flash Attention

FlashAttention-2

Розробники оригінального Flash Attention представили FlashAttention-2 у 2023 році. Ця оновлена версія включає кілька покращень:

  1. Додаткова оптимізація: FlashAttention-2 досягає ще кращого використання GPU, досягнувши до 70% теоретичного піку FLOPS на GPU A100.
  2. Покращений зворотній проход: Зворотній проход оптимізований для досягнення майже тієї ж швидкості, що і прямий проход, що призводить до значних прискорень під час навчання.
  3. Підтримка різних варіантів уваги: FlashAttention-2 розширює підтримку різних варіантів уваги, включаючи групову увагу запитів і багаторазову увагу.

FlashAttention-3

Випущений у 2024 році, FlashAttention-3 представляє останнє досягнення у цій лінії досліджень. Він вводить кілька нових технік для подальшого покращення продуктивності:

  1. Асинхронне обчислення: Використання асинхронної природи нових інструкцій GPU для перекриття різних обчислень.
  2. Підтримка FP8: Використання низькопрецізного обчислення FP8 для ще швидшої обробки.
  3. Некогерентна обробка: Техніка для зменшення похибки квантування при використанні низькопрецізних форматів.

Ось спрощений приклад того, як FlashAttention-3 може використати асинхронне обчислення:

import torch
from torch.cuda.amp import autocast

<p>def flash_attention_3(Q, K, V, block_size=256):
with autocast(dtype=torch.float8): # Використання FP8 для обчислення
# ... (podobно до попередньої реалізації)</p>

<p># Асинхронне обчислення приклад
with torch.cuda.stream(torch.cuda.Stream()):
# Обчислення GEMM асинхронно
S_block = torch.matmul(Q_block, K_block.transpose(-2, -1)) / (d_model ** 0.5)</p>

<p># Тим часом, на потоці за замовчуванням:
# Підготовка до обчислення softmax</p>

<p># Синхронізація потоків
torch.cuda.synchronize()</p>

<p># Продовження з softmax і обчисленням виводу
# ...</p>

return O

Цей фрагмент коду ілюструє, як FlashAttention-3 може використати асинхронне обчислення та точність FP8. Зверніть увагу, що це спрощений приклад, а фактична реалізація буде значно складнішою та залежатиме від апаратного забезпечення.

Реалізація Flash Attention у ваших проєктах

Якщо ви зацікавлені у використанні Flash Attention у своїх проєктах, у вас є кілька варіантів:

  1. Використання існуючих бібліотек: Багато популярних бібліотек, таких як Hugging Face Transformers, вже включають реалізації Flash Attention. Оновлення до останньої версії та активація відповідних прапорів можуть бути достатніми.
  2. Кастомна реалізація: Для більшої гнучкості або спеціальних випадків використання ви можете реалізувати Flash Attention самостійно. Бібліотека xformers надає хорошу референсну реалізацію.
  3. Оптимізація для апаратного забезпечення: Якщо ви працюєте з конкретним апаратним забезпеченням (наприклад, GPU NVIDIA H100), ви можете використати апаратно-орієнтовані оптимізації для максимальної продуктивності.

Ось приклад того, як ви можете використовувати Flash Attention з бібліотекою Hugging Face Transformers:

from transformers import AutoModel, AutoConfig

<p># Активація Flash Attention
config = AutoConfig.from_pretrained(&quot;bert-base-uncased&quot;)
config.use_flash_attention = True</p>

<p># Завантаження моделі з Flash Attention
model = AutoModel.from_pretrained(&quot;bert-base-uncased&quot;, config=config)</p>

<p># Використання моделі як звичайно
# ...

Виклики та майбутні напрямки

Хоча Flash Attention зробив значний крок вперед у підвищенні ефективності механізмів уваги, залишаються виклики та напрямки для майбутніх досліджень:

  1. Апаратна специфікація: Поточні реалізації часто оптимізовані для конкретних архітектур GPU. Загальне застосування цих оптимізацій на різних апаратних платформах залишається викликом.
  2. Інтеграція з іншими техніками: Об’єднання Flash Attention з іншими методами оптимізації, такими як прунінг, квантування та стиснення моделей, є активною областю досліджень.
  3. Розширення на інші області: Хоча Flash Attention показав великий успіх у NLP, його розширення на інші області, такі як комп’ютерне бачення та багатомодальні моделі, продовжується.
  4. Теоретичне розуміння: Глибше розуміння того, чому Flash Attention працює так добре, може привести до ще більш потужних оптимізацій.

Висновок

Flash Attention революціонізував спосіб реалізації механізмів уваги в моделях трансформерів, суттєво покращуючи обчислювальну ефективність та зменшуючи використання пам’яті без втрат точності.

Як ми дослідили у цій статті, вплив Flash Attention виходить за межі простої оптимізації. Він дозволив розробити потужніші та ефективніші моделі.

Я провів останні п'ять років, занурючись у захопливий світ машинного навчання та глибокого навчання. Моя пристрасть та експертиза привели мене до внеску у понад 50 різних проектів програмної інженерії, з особливим акцентом на AI/ML. Моя тривала цікавість також привела мене до природної обробки мови, галузі, яку я бажаю дослідити далі.