Ускорьте работу функции внимания на графическом процессоре с помощью cuDNN и TransformerEngine.

1. Введение

Обучение Jax на GPU. Лабораторная работа 5: Механизм внимания на GPU.

В этом практическом занятии вы сосредоточитесь на вычислении механизма внимания на графических процессорах NVIDIA. Вы начнете с простой реализации, замените ее встроенным механизмом объединенного внимания JAX, а затем активируете бэкенд NVIDIA cuDNN для той же операции. К концу вы узнаете, что меняется на каждом уровне и насколько быстрее становятся оптимизированные для графического процессора пути по мере роста задачи.

Что вы будете делать

  • Реализуйте масштабируемое внимание на основе скалярного произведения с нуля, используя базовые операции JAX.
  • Замените его на jax.nn.dot_product_attention , встроенный в JAX модуль ядра.
  • Принудительно запустите бэкенд cuDNN с implementation="cudnn" и добавьте причинно-следственное маскирование.
  • Измените длину последовательности сканирования и размер пакета, чтобы увидеть, где использование объединенных ядер приносит наибольшую выгоду.
  • Сравните формы многоголовочного внимания MHA, GQA и MQA.
  • Проведите тестирование производительности NVIDIA TransformerEngine и изучите его путь выполнения кода в режиме FP8.

Что вам понадобится

  • Проект Google Cloud с включенной оплатой и кредитами для семинара или резервированием, покрывающим использование графического процессора.
  • Квота на использование как минимум двух видеокарт NVIDIA L4 в выбранном вами регионе ( как проверить квоту на видеокарты )
  • Выполнение заданий 1–4 или работа в эквивалентной среде JAX GPU с поддержкой CUDA.
  • Знание матричного умножения и функции softmax. Опыт работы с CUDA или cuDNN не требуется.

Примерное время выполнения: 60 минут .

Как работает внимание

Масштабированное внимание с использованием скалярного произведения принимает три входных параметра — Запрос , Ключ и Значение — и вычисляет:

Attention(Q, K, V) = softmax(Q K^T / sqrt(d_k)) V

Коэффициент масштабирования 1 / sqrt(d_k) предотвращает чрезмерное увеличение скалярных произведений по мере увеличения размерности головки, что привело бы к перемещению функции softmax в области, где ее градиенты очень малы.

Четыре этапа, каждый из которых плавно переходит в следующий:

  1. ОценкаQ @ KT : скалярное произведение показывает, насколько каждая позиция запроса должна уделять внимание каждой позиции ключа.
  2. Scale + Softmaxsoftmax(scores / sqrt(d)) : масштабирование предотвращает исчезновение градиента, а softmax преобразует оценки в веса внимания, сумма которых равна 1.
  3. Attendweights @ V : взвешенная сумма векторов значений формирует результат для каждой позиции запроса.
  4. Выходные данные — имеют ту же форму, что и Q: каждая позиция запроса теперь содержит информацию из позиций, на которые она была обращена.

2. Прежде чем начать

Выберите свой проект

В консоли Google Cloud выберите или создайте проект с включенной функцией выставления счетов.

Открытая облачная оболочка

Чтобы запустить сеанс Cloud Shell , нажмите кнопку «Активировать Cloud Shell» (значок терминала в правом верхнем углу консоли), а затем укажите в качестве исполнителя свой проект:

gcloud config set project <YOUR_PROJECT_ID>

Подготовка среды для работы с графическим процессором.

Выполните следующие команды в Cloud Shell . В Codelab 1 подробно объясняется каждая команда, включая требования к квоте GPU и зоне.

gcloud services enable \
  container.googleapis.com \
  compute.googleapis.com \
  iam.googleapis.com \
  cloudresourcemanager.googleapis.com \
  logging.googleapis.com \
  monitoring.googleapis.com

git clone https://github.com/Google-Cloud-AI/partner-ai-nvidia.git
cd partner-ai-nvidia/05-workshops/jax-on-gpu/terraform

cp terraform.tfvars.example terraform.tfvars

Отредактируйте файл terraform.tfvars и установите project_id = " " Затем выполните настройку кластера и разверните project_id = " " :

terraform init
terraform apply
$(terraform output -raw get_credentials_command)

cd ..
kubectl apply -f deploy/jupyter.yaml

terraform apply занимает около 12 минут. После завершения дождитесь завершения работы Pod и LoadBalancer, затем считайте одноразовый токен JupyterLab из лога Pod:

kubectl get pod jax-jupyter -w         # wait for Running, then Ctrl+C
kubectl get svc jax-jupyter-svc -w     # wait for EXTERNAL-IP, then Ctrl+C
kubectl logs jax-jupyter | grep -o 'token=[a-z0-9]*' | head -1

Откройте http:// :8884 , вставьте токен и создайте новый блокнот Python 3 в /workspace . Каждый блок кода из этого практического занятия помещается в ячейку этого блокнота.

Установите все необходимое для этого практического занятия.

!pip install --quiet matplotlib flax

Настройте и проверьте графический процессор.

Импортируйте JAX и убедитесь, что в качестве бэкенда по умолчанию используется графический процессор (GPU). В этой ячейке также определены block_tree , show_table и show_bars — вспомогательные функции, которые используются на каждом последующем шаге для ожидания завершения работы устройства и отображения результатов.

import os
os.environ["LD_LIBRARY_PATH"] = "/usr/local/nvidia/lib64:" + os.environ.get("LD_LIBRARY_PATH", "")

import html
import math
import time
from functools import partial

from IPython.display import HTML, display
import matplotlib.pyplot as plt
import numpy as np

import jax
import jax.numpy as jnp


devices = jax.devices()
gpu_devices = [d for d in devices if d.platform == "gpu"]
device = gpu_devices[0] if gpu_devices else None

print(f"JAX version:     {jax.__version__}")
print(f"Default backend: {jax.default_backend()}")
print(f"Devices:         {devices}")

assert gpu_devices, f"This lab assumes a GPU backend. Available devices: {devices}"
print(f"Using GPU:       {device}")


def block_tree(tree):
    """Wait until a PyTree of JAX arrays is ready on device."""
    return jax.block_until_ready(tree)


def show_table(headers, rows, title=None, aligns=None):
    """Render rows as an HTML table."""
    aligns = aligns or ["left"] * len(headers)
    parts = ["<div style='font-family: system-ui; max-width: 980px;'>"]
    if title:
        parts.append(f"<h4 style='margin: 0 0 8px 0;'>{html.escape(title)}</h4>")
    parts.append("<table style='border-collapse: collapse; width: 100%; font-size: 13px;'>")
    parts.append("<thead><tr>")
    for h, a in zip(headers, aligns):
        parts.append(
            f"<th style='text-align:{a}; border-bottom:1px solid #d0d7de; padding:6px;'>"
            f"{html.escape(str(h))}</th>"
        )
    parts.append("</tr></thead><tbody>")
    for row in rows:
        parts.append("<tr>")
        for cell, a in zip(row, aligns):
            parts.append(
                f"<td style='text-align:{a}; border-bottom:1px solid #eef1f4; padding:6px;'>"
                f"{html.escape(str(cell))}</td>"
            )
        parts.append("</tr>")
    parts.append("</tbody></table></div>")
    display(HTML("".join(parts)))


def show_bars(rows, title, unit="", lower_is_better=False):
    """Render (label, value) pairs as a horizontal bar chart in HTML."""
    max_value = max(float(value) for _, value in rows) or 1.0
    color = "#1a7f37" if not lower_is_better else "#0969da"
    parts = ["<div style='font-family: Arial, sans-serif; max-width: 760px;'>"]
    parts.append(f"<h4 style='margin: 0 0 8px 0;'>{html.escape(title)}</h4>")
    for label, value in rows:
        width = max(3, 100 * float(value) / max_value)
        parts.append(
            "<div style='display:grid; grid-template-columns: 190px 1fr 130px; gap: 8px; "
            "align-items:center; margin: 6px 0;'>"
            f"<div style='font-size:13px;'>{html.escape(str(label))}</div>"
            "<div style='background:#f6f8fa; border-radius:6px; overflow:hidden; height:22px;'>"
            f"<div style='height:22px; width:{width:.1f}%; background:{color};'></div></div>"
            f"<div style='font-size:13px; font-variant-numeric: tabular-nums;'>{float(value):,.1f} {html.escape(unit)}</div>"
            "</div>"
        )
    parts.append(
        f"<div style='font-size:12px; color:#57606a;'>"
        f"{'Lower' if lower_is_better else 'Higher'} is better.</div></div>"
    )
    display(HTML("".join(parts)))

Вы должны увидеть версию JAX, gpu в качестве бэкенда по умолчанию, список устройств CUDA и графический процессор, который используется в остальной части практического задания.

3. Создайте тестовые массивы Q, K и V.

В ядрах механизма внимания с объединенными массивами учитывается расположение входных массивов в памяти, поэтому перед написанием кода механизма внимания необходимо создать массивы Q, K и V в той структуре, которую ожидает JAX. jax.nn.dot_product_attention ожидает следующую структуру:

Тусклый

Значение

Наш вариант по умолчанию

Б

Размер партии

4

Т

Длина последовательности запроса

128

С

Длина последовательности ключ/значение

128 (то же, что и T для самовнимания)

Н

Количество точек внимания

8

ЧАС

Размеры на человека

64

В упрощенной реализации инициализация происходит с использованием float32 , сгенерированных случайными значениями.

BATCH = 4
SEQ_LEN = 128
NUM_HEADS = 8
HEAD_DIM = 64

key = jax.random.key(0)
k1, k2, k3 = jax.random.split(key, 3)

q = jax.random.normal(k1, (BATCH, SEQ_LEN, NUM_HEADS, HEAD_DIM), dtype=jnp.float32)
k = jax.random.normal(k2, (BATCH, SEQ_LEN, NUM_HEADS, HEAD_DIM), dtype=jnp.float32)
v = jax.random.normal(k3, (BATCH, SEQ_LEN, NUM_HEADS, HEAD_DIM), dtype=jnp.float32)

q, k, v = jax.device_put((q, k, v), device)

show_table(
    ["Array", "Shape", "Dtype", "Layout"],
    [
        ("Q (query)", q.shape, q.dtype, "(B, T, N, H)"),
        ("K (key)", k.shape, k.dtype, "(B, S, N, H)"),
        ("V (value)", v.shape, v.dtype, "(B, S, N, H)"),
    ],
    title="Attention inputs on GPU",
)

Вы должны увидеть таблицу, в которой для каждого массива будет отдельная строка, указывающая на форму (4, 128, 8, 64) и тип данных float32 . Вызов jax.device_put привязывает все три массива к графическому процессору, выбранному вами в ячейке настройки, поэтому в последующих тестах производительности ничего не будет измерять передачу данных с хоста.

4. Внедрить механизм внимания с нуля.

Прежде чем использовать объединенное ядро, запишите формулу напрямую, чтобы точно видеть, что именно оно заменяет. Данная реализация выполняет четыре шага в указанном порядке: вычисление оценок, масштабирование, softmax, умножение на значения. Вы транспонируете измерения заголовка и последовательности, чтобы умножение матриц работало по оси последовательности.

Он работает корректно, но запускает три отдельных ядра GPU для умножения матрицы оценок, функции softmax и умножения матрицы значений, а также материализует полную матрицу весов внимания (B, N, T, S) в памяти GPU.

def naive_attention(q, k, v):
    """Scaled dot-product attention from scratch."""
    scale = 1.0 / math.sqrt(q.shape[-1])

    # (B, T, N, H) to (B, N, T, H) so matmul runs over the T/S axis
    q_t = jnp.transpose(q, (0, 2, 1, 3))
    k_t = jnp.transpose(k, (0, 2, 1, 3))
    v_t = jnp.transpose(v, (0, 2, 1, 3))

    # Score: (B, N, T, H) @ (B, N, H, S) to (B, N, T, S)
    scores = jnp.matmul(q_t, jnp.transpose(k_t, (0, 1, 3, 2))) * scale
    weights = jax.nn.softmax(scores, axis=-1)

    # Attend: (B, N, T, S) @ (B, N, S, H) to (B, N, T, H)
    out_t = jnp.matmul(weights, v_t)

    # Back to (B, T, N, H)
    return jnp.transpose(out_t, (0, 2, 1, 3))


naive_out = block_tree(naive_attention(q, k, v))

show_table(
    ["", "Value"],
    [
        ("Output shape", str(naive_out.shape)),
        ("Output dtype", str(naive_out.dtype)),
    ],
    title="Naive attention",
)

Вы должны увидеть выходные данные в формате (4, 128, 8, 64) и с типом данных float32 что соответствует формату Q, как и обещал шаг Output формулы.

5. Переключиться на dot_product_attention

Функция jax.nn.dot_product_attention объединяет этапы оценки, масштабирования, softmax и внимания в одну операцию. JAX и XLA затем могут оптимизировать схему доступа к памяти. В частности, они могут избежать материализации полной матрицы весов внимания, когда последовательность длинная.

При значении implementation=None по умолчанию JAX автоматически выбирает наилучший доступный бэкенд. На графическом процессоре с доступным cuDNN и совместимыми входными данными он может уже использовать cuDNN. На другом оборудовании он переключается на XLA.

sdpa_out = block_tree(jax.nn.dot_product_attention(q, k, v))

max_diff = float(jnp.max(jnp.abs(naive_out - sdpa_out)))

show_table(
    ["", "Value"],
    [
        ("Output shape", str(sdpa_out.shape)),
        ("Output dtype", str(sdpa_out.dtype)),
        ("Max |naive − SDPA|", f"{max_diff:.2e}"),
        ("Outputs close (atol=1e-3)", str(bool(jnp.allclose(naive_out, sdpa_out, atol=1e-3)))),
    ],
    title="JAX SDPA vs naive",
)

Вы должны увидеть ту же форму и тип данных, что и в наивной версии, небольшую максимальную абсолютную разницу и True для проверки близости atol=1e-3 .

6. Принудительно запустите бэкэнд cuDNN fused.

Позволить JAX выбирать бэкенд удобно, но извне невозможно определить, какое именно ядро ​​было запущено. Установка implementation="cudnn" принудительно запускает ядра NVIDIA cuDNN с механизмом внимания. Это оптимизированные вручную ядра для графических процессоров, которые объединяют все вычисления, связанные с механизмом внимания, в один запуск ядра с оптимизированными шаблонами доступа к памяти.

Обратите внимание, что механизм внимания cuDNN имеет аппаратные и структурные требования, такие как вычислительные возможности графического процессора, которые должны быть >= 8.0 (Ampere или новее) или иметь входные типы данных float16 или bfloat16 . Если эти требования не выполняются и вы устанавливаете implementation="cudnn" , JAX выдает ошибку, а не молча возвращается к исходному состоянию. Именно поэтому приведенный ниже код преобразует тип данных в bfloat16 и оборачивает вызов в try / except : он устанавливает HAS_CUDNN_SDPA , чтобы каждый последующий шаг знал, доступен ли путь к cuDNN на этом компьютере.

q_bf16 = q.astype(jnp.bfloat16)
k_bf16 = k.astype(jnp.bfloat16)
v_bf16 = v.astype(jnp.bfloat16)

HAS_CUDNN_SDPA = False

try:
    cudnn_out = block_tree(
        jax.nn.dot_product_attention(q_bf16, k_bf16, v_bf16, implementation="cudnn")
    )
    HAS_CUDNN_SDPA = True

    xla_bf16_out = block_tree(
        jax.nn.dot_product_attention(q_bf16, k_bf16, v_bf16, implementation="xla")
    )
    max_diff = float(jnp.max(jnp.abs(
        cudnn_out.astype(jnp.float32) - xla_bf16_out.astype(jnp.float32)
    )))

    show_table(
        ["", "Value"],
        [
            ("Output shape", str(cudnn_out.shape)),
            ("Output dtype", str(cudnn_out.dtype)),
            ("Max |cuDNN − XLA| (both bf16)", f"{max_diff:.2e}"),
            ("Outputs close (rtol=1e-2, atol=1e-2)", str(bool(jnp.allclose(cudnn_out, xla_bf16_out, rtol=1e-2, atol=1e-2)))),
        ],
        title="cuDNN fused attention",
    )

except Exception as e:
    print(f"cuDNN SDPA not available on this GPU: {e}")
    print("Continuing with XLA backend only.")

Если ваш графический процессор не может запустить cuDNN SDPA, код выводит причину, и выполнение практического задания продолжается на бэкенде XLA.

7. Добавить причинно-следственную маскировку.

В авторегрессионных моделях (декодерах в стиле GPT) каждая позиция может учитывать только более ранние позиции. Установка параметра is_causal=True применяет эту нижнетреугольную маску внутри объединенного ядра, поэтому вам не нужно самостоятельно создавать матрицу маски.

Приведённый ниже код запускает механизм внимания дважды, с маской и без неё, и сравнивает две позиции, чтобы показать, что изменила маска.

causal_out = block_tree(
    jax.nn.dot_product_attention(q, k, v, is_causal=True)
)

# With causal masking, the last position attends to all positions.
# The first position attends only to itself.
nocausal_out = block_tree(
    jax.nn.dot_product_attention(q, k, v, is_causal=False)
)

# First position should differ
first_pos_diff = float(jnp.max(jnp.abs(causal_out[:, 0] - nocausal_out[:, 0])))
# Last position should be the same
last_pos_diff = float(jnp.max(jnp.abs(causal_out[:, -1] - nocausal_out[:, -1])))

show_table(
    ["Position", "Max diff (causal vs full)", "Expected"],
    [
        ("First (t=0)", f"{first_pos_diff:.4f}", "Large — causal restricts to self only"),
        ("Last (t=T-1)", f"{last_pos_diff:.2e}", "~0 — attends to all positions either way"),
    ],
    title="Causal masking effect on attention output",
)

Вы должны увидеть большую разницу на первой позиции и почти нулевую разницу на последней. Эта асимметрия обусловлена ​​тем, что маска с позицией 0 теряет все ключи, кроме своего собственного, в то время как на последней позиции уже видна вся последовательность.

8. Запуск нескольких сравнительных тестов.

Теперь у вас есть три способа вычислить одну и ту же функцию. Какой из них выбрать, полностью зависит от формы задачи, и вы можете это вычислить.

Засеките время для вариантов при стандартной форме.

def benchmark_attention(fn, q, k, v, warmup=3, repeats=50):
    """Time an attention function. Returns median milliseconds per call."""
    jit_fn = jax.jit(fn)

    for _ in range(warmup):
        block_tree(jit_fn(q, k, v))

    times = []
    for _ in range(repeats):
        start = time.perf_counter()
        block_tree(jit_fn(q, k, v))
        times.append((time.perf_counter() - start) * 1000)

    return np.median(times)


t_naive = benchmark_attention(naive_attention, q, k, v)
t_sdpa = benchmark_attention(
    lambda q, k, v: jax.nn.dot_product_attention(q, k, v, implementation="xla"),
    q, k, v,
)

results = [
    ("Naive (matmul + softmax + matmul)", f"{t_naive:.2f}"),
    ("SDPA (XLA, float32)", f"{t_sdpa:.2f}"),
]
bar_data = [
    ("Naive", t_naive),
    ("SDPA XLA f32", t_sdpa),
]

if HAS_CUDNN_SDPA:
    t_sdpa_bf16 = benchmark_attention(
        lambda q, k, v: jax.nn.dot_product_attention(q, k, v, implementation="xla"),
        q_bf16, k_bf16, v_bf16,
    )
    t_cudnn = benchmark_attention(
        lambda q, k, v: jax.nn.dot_product_attention(q, k, v, implementation="cudnn"),
        q_bf16, k_bf16, v_bf16,
    )
    results.append(("SDPA (XLA, bfloat16)", f"{t_sdpa_bf16:.2f}"))
    results.append(("SDPA (cuDNN, bfloat16)", f"{t_cudnn:.2f}"))
    bar_data.append(("SDPA XLA bf16", t_sdpa_bf16))
    bar_data.append(("SDPA cuDNN bf16", t_cudnn))

show_table(
    ["Implementation", "Median ms/call"],
    results,
    title=f"Attention timing — B={BATCH}, T={SEQ_LEN}, N={NUM_HEADS}, H={HEAD_DIM}",
    aligns=["left", "right"],
)
show_bars(bar_data, "Attention latency (ms per call)", "ms", lower_is_better=True)

Для такого небольшого размера задачи все реализации очень быстры и в основном ограничены накладными расходами, поэтому наивная версия конкурентоспособна. cuDNN bf16 должна быть немного быстрее, но преимущество объединенного внимания обычно становится более очевидным при большей длине последовательности, где важно избегать полной матрицы внимания.

Просканируйте длину последовательности

Преимущество объединенных ядер внимания возрастает с увеличением длины последовательности. Простейшая реализация материализует матрицу внимания (B, N, T, S) в памяти графического процессора, что занимает O(T²) памяти. Объединенные ядра, такие как cuDNN FlashAttention, разбивают вычисления на части, поэтому они никогда не материализуют всю матрицу целиком, сохраняя объем памяти на уровне O(T).

Приведенный ниже анализ времени выполнения каждой реализации охватывает последовательности длиной от 64 до 1024 символов.

SEQ_LENS = [64, 128, 256, 512, 1024]

sweep_results = []
for sl in SEQ_LENS:
    rk = jax.random.key(sl)
    rk1, rk2, rk3 = jax.random.split(rk, 3)

    q_s = jax.random.normal(rk1, (BATCH, sl, NUM_HEADS, HEAD_DIM), dtype=jnp.float32)
    k_s = jax.random.normal(rk2, (BATCH, sl, NUM_HEADS, HEAD_DIM), dtype=jnp.float32)
    v_s = jax.random.normal(rk3, (BATCH, sl, NUM_HEADS, HEAD_DIM), dtype=jnp.float32)
    q_s, k_s, v_s = jax.device_put((q_s, k_s, v_s), device)

    q_sb = q_s.astype(jnp.bfloat16)
    k_sb = k_s.astype(jnp.bfloat16)
    v_sb = v_s.astype(jnp.bfloat16)

    row = {"seq_len": sl}

    row["naive_ms"] = benchmark_attention(
        naive_attention,
        q_s, k_s, v_s,
        warmup=2,
        repeats=20,
    )

    row["sdpa_xla_f32_ms"] = benchmark_attention(
        lambda q, k, v: jax.nn.dot_product_attention(q, k, v, implementation="xla"),
        q_s, k_s, v_s,
        warmup=2,
        repeats=20,
    )

    row["sdpa_xla_bf16_ms"] = benchmark_attention(
        lambda q, k, v: jax.nn.dot_product_attention(q, k, v, implementation="xla"),
        q_sb, k_sb, v_sb,
        warmup=2,
        repeats=20,
    )

    if HAS_CUDNN_SDPA:
        row["cudnn_bf16_ms"] = benchmark_attention(
            lambda q, k, v: jax.nn.dot_product_attention(q, k, v, implementation="cudnn"),
            q_sb, k_sb, v_sb,
            warmup=2,
            repeats=20,
        )

    sweep_results.append(row)


headers = ["Seq len", "Naive (ms)", "SDPA XLA f32 (ms)", "SDPA XLA bf16 (ms)"]
if HAS_CUDNN_SDPA:
    headers.append("cuDNN bf16 (ms)")

table_rows = []
for r in sweep_results:
    row = [
        r["seq_len"],
        f"{r['naive_ms']:.2f}",
        f"{r['sdpa_xla_f32_ms']:.2f}",
        f"{r['sdpa_xla_bf16_ms']:.2f}",
    ]

    if HAS_CUDNN_SDPA:
        row.append(f"{r['cudnn_bf16_ms']:.2f}")

    table_rows.append(row)

show_table(
    headers,
    table_rows,
    title=f"Sequence-length sweep — B={BATCH}, N={NUM_HEADS}, H={HEAD_DIM}",
    aligns=["right"] * len(headers),
)

Процесс обхода занимает некоторое время, поскольку для каждой длины последовательности компилируются три или четыре отдельных варианта, прежде чем будет измерено их время выполнения. В итоге вы должны получить одну строку в таблице для каждой длины последовательности от 64 до 1024.

На этом графике сравнивается изменение задержки внимания в зависимости от длины последовательности для наивной реализации, JAX/XLA SDPA в форматах float32 и bfloat16, а также для объединенного внимания cuDNN в формате bfloat16.

fig, ax = plt.subplots(figsize=(8, 5))
seq_lens = [r["seq_len"] for r in sweep_results]

ax.plot(
    seq_lens,
    [r["naive_ms"] for r in sweep_results],
    "o-",
    label="Naive",
    color="#d1242f",
)

ax.plot(
    seq_lens,
    [r["sdpa_xla_f32_ms"] for r in sweep_results],
    "s-",
    label="SDPA XLA f32",
    color="#0969da",
)

ax.plot(
    seq_lens,
    [r["sdpa_xla_bf16_ms"] for r in sweep_results],
    "d-",
    label="SDPA XLA bf16",
    color="#8250df",
)

if HAS_CUDNN_SDPA:
    ax.plot(
        seq_lens,
        [r["cudnn_bf16_ms"] for r in sweep_results],
        "^-",
        label="cuDNN bf16",
        color="#1a7f37",
    )

ax.set_xlabel("Sequence length")
ax.set_ylabel("Median ms per call")
ax.set_title("Attention latency vs sequence length")
ax.legend()
ax.grid(True, alpha=0.25)
ax.set_xticks(seq_lens)

fig.tight_layout()
plt.show()

Вы должны видеть одну линию на каждую реализацию. И вы должны видеть разрыв, который увеличивается по мере удлинения последовательности.

Проведите перебор размера партии.

Более крупные пакеты данных компенсируют накладные расходы на запуск ядра и повышают эффективность использования графического процессора до тех пор, пока память графического процессора не станет узким местом. В приведенном ниже примере длина последовательности остается фиксированной на уровне 256, а размер пакета данных изменяется.

BATCH_SIZES = [1, 2, 4, 8, 16]
SWEEP_SEQ = 256

batch_results = []
for bs in BATCH_SIZES:
    rk = jax.random.key(bs + 100)
    rk1, rk2, rk3 = jax.random.split(rk, 3)

    q_b = jax.random.normal(rk1, (bs, SWEEP_SEQ, NUM_HEADS, HEAD_DIM), dtype=jnp.float32)
    k_b = jax.random.normal(rk2, (bs, SWEEP_SEQ, NUM_HEADS, HEAD_DIM), dtype=jnp.float32)
    v_b = jax.random.normal(rk3, (bs, SWEEP_SEQ, NUM_HEADS, HEAD_DIM), dtype=jnp.float32)
    q_b, k_b, v_b = jax.device_put((q_b, k_b, v_b), device)

    q_bb = q_b.astype(jnp.bfloat16)
    k_bb = k_b.astype(jnp.bfloat16)
    v_bb = v_b.astype(jnp.bfloat16)

    row = {"batch": bs}

    row["sdpa_xla_f32_ms"] = benchmark_attention(
        lambda q, k, v: jax.nn.dot_product_attention(q, k, v, implementation="xla"),
        q_b, k_b, v_b,
        warmup=2,
        repeats=20,
    )

    row["sdpa_xla_bf16_ms"] = benchmark_attention(
        lambda q, k, v: jax.nn.dot_product_attention(q, k, v, implementation="xla"),
        q_bb, k_bb, v_bb,
        warmup=2,
        repeats=20,
    )

    if HAS_CUDNN_SDPA:
        row["cudnn_bf16_ms"] = benchmark_attention(
            lambda q, k, v: jax.nn.dot_product_attention(q, k, v, implementation="cudnn"),
            q_bb, k_bb, v_bb,
            warmup=2,
            repeats=20,
        )

    batch_results.append(row)


headers = ["Batch size", "SDPA XLA f32 (ms)", "SDPA XLA bf16 (ms)"]
if HAS_CUDNN_SDPA:
    headers.append("cuDNN bf16 (ms)")

table_rows = []
for r in batch_results:
    row = [
        r["batch"],
        f"{r['sdpa_xla_f32_ms']:.2f}",
        f"{r['sdpa_xla_bf16_ms']:.2f}",
    ]

    if HAS_CUDNN_SDPA:
        row.append(f"{r['cudnn_bf16_ms']:.2f}")

    table_rows.append(row)

show_table(
    headers,
    table_rows,
    title=f"Batch-size sweep — T={SWEEP_SEQ}, N={NUM_HEADS}, H={HEAD_DIM}",
    aligns=["right"] * len(headers),
)


fig, ax = plt.subplots(figsize=(8, 5))
batches = [r["batch"] for r in batch_results]

ax.plot(
    batches,
    [r["sdpa_xla_f32_ms"] for r in batch_results],
    "s-",
    label="SDPA XLA f32",
    color="#0969da",
)

ax.plot(
    batches,
    [r["sdpa_xla_bf16_ms"] for r in batch_results],
    "d-",
    label="SDPA XLA bf16",
    color="#8250df",
)

if HAS_CUDNN_SDPA:
    ax.plot(
        batches,
        [r["cudnn_bf16_ms"] for r in batch_results],
        "^-",
        label="cuDNN bf16",
        color="#1a7f37",
    )

ax.set_xlabel("Batch size")
ax.set_ylabel("Median ms per call")
ax.set_title("Attention latency vs batch size")
ax.legend()
ax.grid(True, alpha=0.25)
ax.set_xticks(batches)

fig.tight_layout()
plt.show()

В результате вы получите таблицу и график, содержащие по одной записи на каждый размер пакета от 1 до 16. Простейшая реализация в этом анализе исключена, поэтому оба результата сравнивают только три объединенных пути.

9. Сравните MHA, GQA и MQA.

До сих пор каждая головка хранила свои собственные ключи и значения. Но декодеры часто нарушают эту симметрию, потому что во время вывода в памяти доминирует кэш ключ-значение.

Многоголовочное внимание (MHA) предоставляет каждой голове собственные проекции Q, K и V. Внимание с групповыми запросами (GQA) и внимание с множественными запросами (MQA) уменьшают количество голов KV для экономии памяти и вычислительных ресурсов во время вывода.

jax.nn.dot_product_attention обрабатывает все три параметра, при этом сигналы KV передаются автоматически, когда K < N.

rk = jax.random.key(42)
rk1, rk2, rk3, rk4, rk5 = jax.random.split(rk, 5)

q_mha = jax.random.normal(rk1, (2, 64, 8, 64), dtype=jnp.float32)

# MHA: 8 KV heads
k_mha = jax.random.normal(rk2, (2, 64, 8, 64), dtype=jnp.float32)
v_mha = jax.random.normal(rk3, (2, 64, 8, 64), dtype=jnp.float32)

# GQA: 2 KV heads (each shared by 4 query heads)
k_gqa = jax.random.normal(rk2, (2, 64, 2, 64), dtype=jnp.float32)
v_gqa = jax.random.normal(rk3, (2, 64, 2, 64), dtype=jnp.float32)

# MQA: 1 KV head (shared by all 8 query heads)
k_mqa = jax.random.normal(rk4, (2, 64, 1, 64), dtype=jnp.float32)
v_mqa = jax.random.normal(rk5, (2, 64, 1, 64), dtype=jnp.float32)

out_mha = block_tree(jax.nn.dot_product_attention(q_mha, k_mha, v_mha))
out_gqa = block_tree(jax.nn.dot_product_attention(q_mha, k_gqa, v_gqa))
out_mqa = block_tree(jax.nn.dot_product_attention(q_mha, k_mqa, v_mqa))

show_table(
    ["Pattern", "Q shape", "K shape", "V shape", "Output shape"],
    [
        ("MHA", q_mha.shape, k_mha.shape, v_mha.shape, out_mha.shape),
        ("GQA", q_mha.shape, k_gqa.shape, v_gqa.shape, out_gqa.shape),
        ("MQA", q_mha.shape, k_mqa.shape, v_mqa.shape, out_mqa.shape),
    ],
    title="Multi-head attention variants — all should produce the same output shape",
)

Все три строки должны выдавать одинаковую форму выходных данных, (2, 64, 8, 64) , даже несмотря на то, что формы K и V уменьшаются с 8 головок до 2 к 1. В этом и суть: можно сократить кэш KV, не меняя ничего после механизма внимания.

10. Сравнительный анализ NVIDIA TransformerEngine и FP8.

NVIDIA TransformerEngine предоставляет модули интегрированного внимания, оптимизированные для графических процессоров NVIDIA. Интеграция с JAX использует модули в стиле Flax Linen (а не NNX), поэтому модуль инициализируется один раз, а затем применяется со своими переменными.

Длина последовательности развертки с помощью TransformerEngine

Эта ячейка проверяет доступность NVIDIA TransformerEngine, а затем проводит сравнительный анализ DotProductAttention для причинно-следственных связей bf16 на последовательностях различной длины с использованием JAX SDPA с бэкэндами XLA и cuDNN на той же рабочей нагрузке.

TE_SEQ_LENS = [128, 256, 512, 1024, 2048]
TE_BATCH = BATCH

HAS_TE = False

try:
    import transformer_engine.jax as te
    import transformer_engine.jax.flax as te_flax
    HAS_TE = True
except ImportError:
    print("TransformerEngine not installed — skipping TE sections.")

if HAS_TE:
    te_results = []

    for sl in TE_SEQ_LENS:
        rk = jax.random.key(sl + 1000)
        rk1, rk2, rk3 = jax.random.split(rk, 3)

        q_te = jax.random.normal(
            rk1, (TE_BATCH, sl, NUM_HEADS, HEAD_DIM), dtype=jnp.bfloat16
        )
        k_te = jax.random.normal(
            rk2, (TE_BATCH, sl, NUM_HEADS, HEAD_DIM), dtype=jnp.bfloat16
        )
        v_te = jax.random.normal(
            rk3, (TE_BATCH, sl, NUM_HEADS, HEAD_DIM), dtype=jnp.bfloat16
        )
        q_te, k_te, v_te = jax.device_put((q_te, k_te, v_te), device)

        te_attention = te_flax.DotProductAttention(
            head_dim=HEAD_DIM,
            num_attention_heads=NUM_HEADS,
            num_gqa_groups=NUM_HEADS,
            attn_mask_type="causal",
            transpose_batch_sequence=False,
        )

        te_vars = te_attention.init(
            jax.random.key(0),
            q_te,
            k_te,
            v_te,
            deterministic=True,
        )

        def te_fn(q, k, v):
            return te_attention.apply(te_vars, q, k, v, deterministic=True)

        row = {"seq_len": sl}

        row["sdpa_xla_bf16_ms"] = benchmark_attention(
            lambda q, k, v: jax.nn.dot_product_attention(
                q, k, v, implementation="xla", is_causal=True
            ),
            q_te, k_te, v_te,
            warmup=2,
            repeats=20,
        )

        if HAS_CUDNN_SDPA:
            row["sdpa_cudnn_bf16_ms"] = benchmark_attention(
                lambda q, k, v: jax.nn.dot_product_attention(
                    q, k, v, implementation="cudnn", is_causal=True
                ),
                q_te, k_te, v_te,
                warmup=2,
                repeats=20,
            )

        row["te_bf16_ms"] = benchmark_attention(
            te_fn,
            q_te, k_te, v_te,
            warmup=2,
            repeats=20,
        )

        te_results.append(row)


    headers = ["Seq len", "SDPA XLA bf16 causal (ms)"]
    if HAS_CUDNN_SDPA:
        headers.append("SDPA cuDNN bf16 causal (ms)")
    headers.append("TE DotProductAttention bf16 causal (ms)")

    table_rows = []
    for r in te_results:
        row = [
            r["seq_len"],
            f"{r['sdpa_xla_bf16_ms']:.2f}",
        ]

        if HAS_CUDNN_SDPA:
            row.append(f"{r['sdpa_cudnn_bf16_ms']:.2f}")

        row.append(f"{r['te_bf16_ms']:.2f}")
        table_rows.append(row)

    show_table(
        headers,
        table_rows,
        title=f"TransformerEngine sequence-length sweep — B={TE_BATCH}, N={NUM_HEADS}, H={HEAD_DIM}",
        aligns=["right"] * len(headers),
    )


    fig, ax = plt.subplots(figsize=(8, 5))
    seq_lens = [r["seq_len"] for r in te_results]

    ax.plot(
        seq_lens,
        [r["sdpa_xla_bf16_ms"] for r in te_results],
        "d-",
        label="SDPA XLA bf16 causal",
        color="#8250df",
    )

    if HAS_CUDNN_SDPA:
        ax.plot(
            seq_lens,
            [r["sdpa_cudnn_bf16_ms"] for r in te_results],
            "^-",
            label="SDPA cuDNN bf16 causal",
            color="#1a7f37",
        )

    ax.plot(
        seq_lens,
        [r["te_bf16_ms"] for r in te_results],
        "o-",
        label="TE DotProductAttention bf16 causal",
        color="#d1242f",
    )

    ax.set_xlabel("Sequence length")
    ax.set_ylabel("Median ms per call")
    ax.set_title("Causal attention latency vs sequence length")
    ax.legend()
    ax.grid(True, alpha=0.25)
    ax.set_xticks(seq_lens)

    fig.tight_layout()
    plt.show()

На каждую длину последовательности от 128 до 2048 символов должна приходиться одна строка таблицы и одна точка графика. Этот бенчмарк bf16 сравнивает TransformerEngine с JAX SDPA на одной и той же задаче с причинно-следственным вниманием, но полный потенциал производительности TransformerEngine обычно проявляется на графических процессорах Hopper и Blackwell, когда доступен автопреобразование FP8.

Проверьте путь FP8.

На графических процессорах Hopper (с вычислительной мощностью >= 9.0, например, H100) TransformerEngine может запускать механизм внимания в FP8 для повышения пропускной способности. В FP8 используется алгоритм DelayedScaling , который отслеживает историю абсолютных максимумов для каждого тензора для вычисления динамических коэффициентов масштабирования.

  • Формат E4M3 для прямого прохода (4 бита экспоненты, 3 бита мантиссы)
  • Формат E5M2 для обратного прохода (5 бит экспоненты, 2 бита мантиссы)

Если графический процессор не поддерживает FP8, в этой ячейке показано, как будет выглядеть код без его запуска.

if HAS_TE:
    from transformer_engine.common.recipe import DelayedScaling, Format

    gpu_name = f"{device} {getattr(device, 'device_kind', '')}".lower()
    HAS_FP8 = any(
        tag in gpu_name
        for tag in ["h100", "h200", "b100", "b200", "gb200", "blackwell"]
    )

    fp8_recipe = DelayedScaling(
        margin=0,
        fp8_format=Format.HYBRID,
        amax_history_len=1024,
        amax_compute_algo="max",
    )

    if HAS_FP8:
        FP8_SEQ_LEN = 2048
        FP8_BATCH = BATCH

        rk = jax.random.key(9000)
        rk1, rk2, rk3 = jax.random.split(rk, 3)

        q_fp8 = jax.random.normal(
            rk1, (FP8_BATCH, FP8_SEQ_LEN, NUM_HEADS, HEAD_DIM), dtype=jnp.bfloat16
        )
        k_fp8 = jax.random.normal(
            rk2, (FP8_BATCH, FP8_SEQ_LEN, NUM_HEADS, HEAD_DIM), dtype=jnp.bfloat16
        )
        v_fp8 = jax.random.normal(
            rk3, (FP8_BATCH, FP8_SEQ_LEN, NUM_HEADS, HEAD_DIM), dtype=jnp.bfloat16
        )
        q_fp8, k_fp8, v_fp8 = jax.device_put((q_fp8, k_fp8, v_fp8), device)

        fp8_attention = te_flax.DotProductAttention(
            head_dim=HEAD_DIM,
            num_attention_heads=NUM_HEADS,
            num_gqa_groups=NUM_HEADS,
            attn_mask_type="causal",
            transpose_batch_sequence=False,
        )

        bf16_vars = fp8_attention.init(
            jax.random.key(0),
            q_fp8,
            k_fp8,
            v_fp8,
            deterministic=True,
        )

        bf16_out = block_tree(
            fp8_attention.apply(
                bf16_vars,
                q_fp8,
                k_fp8,
                v_fp8,
                deterministic=True,
            )
        )

        with te.autocast(enabled=True, recipe=fp8_recipe):
            fp8_vars = fp8_attention.init(
                jax.random.key(1),
                q_fp8,
                k_fp8,
                v_fp8,
                deterministic=True,
            )
            fp8_out = block_tree(
                fp8_attention.apply(
                    fp8_vars,
                    q_fp8,
                    k_fp8,
                    v_fp8,
                    deterministic=True,
                )
            )

        max_diff_fp8 = float(jnp.max(jnp.abs(
            bf16_out.astype(jnp.float32) - fp8_out.astype(jnp.float32)
        )))

        show_table(
            ["", "Value"],
            [
                ("GPU", getattr(device, "device_kind", str(device))),
                ("Input dtype", str(q_fp8.dtype)),
                ("bf16 output dtype", str(bf16_out.dtype)),
                ("FP8 autocast output dtype", str(fp8_out.dtype)),
                ("Output shape", str(fp8_out.shape)),
                ("Max |TE bf16 - TE FP8 autocast|", f"{max_diff_fp8:.2e}"),
            ],
            title="FP8 attention with TransformerEngine",
        )

    else:
        show_table(
            ["", "Value"],
            [
                ("GPU", getattr(device, "device_kind", str(device))),
                ("FP8 support", "No detected support; requires Hopper/Blackwell-class GPU"),
            ],
            title="FP8 attention — not available on this GPU",
        )

        print()
        print("The FP8 path uses TransformerEngine autocast:")
        print()
        print("  with te.autocast(enabled=True, recipe=fp8_recipe):")
        print("      out = fp8_attention.apply(vars, q, k, v, deterministic=True)")

else:
    print("TransformerEngine not available — FP8 section skipped.")

На уровне L4 вы должны увидеть таблицу с указанием вашего графического процессора и сообщением об отсутствии поддержки FP8, за которой следуют две строки с вызовом te.autocast , который вы бы использовали на графическом процессоре Hopper. Сохраните этот фрагмент: это единственное изменение, которое требуется для поддержки FP8 в месте вызова.

11. Уборка

Удалите рабочую нагрузку Jupyter, включая балансировщик нагрузки и постоянный том:

kubectl delete -f deploy/jupyter.yaml

Уничтожьте кластер, пул узлов, VPC и учетную запись службы:

cd terraform
terraform destroy

При появлении запроса введите yes , затем подтвердите, что ничего не осталось:

gcloud container clusters list
gcloud compute instances list

Оба поля должны быть пустыми для этого проекта. Если вы создали проект только для этой серии, вы можете удалить весь проект из консоли Cloud .

12. Поздравляем!

Вы перешли от написанного вручную алгоритма внимания к оптимизированным для графического процессора объединенным ядрам и измерили разницу на реальном оборудовании.

Что вы узнали

  • Наивный механизм внимания работает, но запускает несколько ядер GPU и материализует всю матрицу внимания в памяти.
  • jax.nn.dot_product_attention объединяет вычисления в одну операцию, а при implementation=None JAX автоматически выбирает наилучший бэкенд.
  • implementation="cudnn" принудительно использует ядра интегрированного внимания NVIDIA cuDNN, которые работают быстрее всего на длинных последовательностях, и требует входных данных в bfloat16 или float16 и вычислительных мощностей версии 8.0 или выше.
  • Маскирование причинно-следственных связей с is_causal=True встроено в ядро ​​fused — ручная матрица маскирования не требуется.
  • GQA и MQA уменьшают количество заголовков ключ-значение для экономии памяти во время вывода, а dot_product_attention автоматически обрабатывает широковещательную рассылку.
  • TransformerEngine обеспечивает объединенное внимание с опциональной точностью FP8 на графических процессорах Hopper.

Следующие шаги

  • В практическом занятии Codelab 6: Масштабирование обучения JAX на нескольких графических процессорах будет продемонстрировано разделение массивов данных между слоями L4 и параллельное выполнение этапов обучения.
  • Измените HEAD_DIM (попробуйте 32, 64, 128) и понаблюдайте, какие значения принимает cuDNN и как меняются временные параметры.
  • Увеличьте значение SEQ_LEN (попробуйте 256, 512, 1024, 2048) и понаблюдайте за использованием памяти и коэффициентом ускорения cuDNN.
  • Повторите сравнение MHA, GQA и MQA с K=2 и K=1 KV-заголовками против 8 запросов и проанализируйте экономию KV-кэша для вывода результатов.

Справочная документация