1. Введение

В практическом занятии "Управление компиляцией JAX с помощью jax.jit" вы узнали, что компиляция является одной из причин медленной работы JAX-программ, а также как использовать jax.jit для предотвращения повторной компиляции. Но компиляция — лишь одна из возможных причин.
Когда нагрузка на графический процессор низкая, код Python почти никогда не указывает на истинную причину. Это может быть компиляция, измерение времени, которое не дождалось обработки графическим процессором, передача данных между хостом и устройством, скрытая в логах, слишком маленький пакет данных, не способный обеспечить достаточную загрузку графического процессора, или нехватка памяти.
В этом практическом занятии вы начнете измерять эти аспекты. Вы проанализируете реальный этап обучения, прочитаете трассировку в XProf и перейдете к временной шкале CUDA с помощью Nsight Systems.
Что вы будете делать
- Разделите время компиляции первого вызова от времени выполнения, кэшированного в кэше, и честно рассчитайте время выполнения JAX с помощью
block_until_ready() - Захват трассировки JAX-профилировщика с помощью
jax.profiler.traceи именованных аннотаций шагов. - Откройте этот трассировочный файл в XProf через переадресацию портов Cloud Shell, используя TensorBoard в качестве альтернативного интерфейса.
- Проведите диагностику четырех распространенных причин замедления работы: слишком много мелких операций, передача данных между хостом и устройством, небольшие пакеты данных и нехватка памяти.
- Создавайте временную шкалу на уровне CUDA с помощью Nsight Systems , используя диапазоны NVTX для обозначения интересующей вас области.
- Сведите итоги отчета Nsight в блокноте с помощью
nsys stats
Что вам понадобится
- Проект Google Cloud с включенной оплатой и кредитами для семинара или резервированием, покрывающим использование графического процессора.
- Квота на использование как минимум двух видеокарт NVIDIA L4 в выбранном вами регионе ( как проверить квоту на видеокарты )
- Выполнение лабораторных работ 1 и 2 или работа в эквивалентной среде JAX GPU.
- При желании вы можете установить Nsight Systems на свой компьютер, чтобы открыть отчет CUDA в графическом интерфейсе. Это бесплатная программа, и в ходе практического занятия также будет выведено текстовое резюме этого отчета.
Примерное время выполнения: 70 минут .
Процесс профилирования
Эффективный подход к отладке производительности JAX заключается в переходе от простых проверок к более сложным инструментам.
Сначала выполните замеры времени, блокируя выполнение до завершения работы на графическом процессоре. Затем запишите трассировку JAX-профилировщика, чтобы увидеть компиляцию, активность хоста и выполнение на устройстве одновременно. Далее откройте эту трассировку в XProf или TensorBoard, чтобы просмотреть временные шкалы, память, графики и статистику операций. Когда вам потребуется представление на уровне CUDA, используйте Nsight Systems для просмотра потоков, ядер, копирования памяти, вызовов библиотек и обмена данными.
На что обратить внимание
При открытии трассировки не пытайтесь понять каждое событие сразу. Начните с поиска нескольких распространенных визуальных закономерностей.
Показатели компиляции указывают на то, уходит ли время на компиляцию XLA, а не на выполнение. Пустые промежутки в строках GPU часто означают, что хост не передает данные на устройство достаточно быстро. Активность передачи данных может указывать на случайную синхронизацию хоста и устройства, например, на запись в журнал с использованием float(loss) . Пиковые значения использования памяти помогают выявить пакеты данных, активации или временные буферы, которые приближают GPU к пределу его возможностей.
2. Прежде чем начать
Выберите свой проект
В консоли Google Cloud выберите или создайте проект с включенной функцией выставления счетов.
Открытая облачная оболочка
Чтобы запустить сеанс Cloud Shell , нажмите кнопку «Активировать Cloud Shell» (значок терминала в правом верхнем углу консоли), а затем укажите в качестве исполнителя свой проект:
gcloud config set project <YOUR_PROJECT_ID>
Этот практический пример выполняется в той же среде, что и практический пример 1: Запуск вашей первой программы JAX на графических процессорах NVIDIA с использованием GKE . Если ваш кластер GKE и под JupyterLab все еще работают, перейдите к разделу «Установка необходимых компонентов для этого практического примера ». В противном случае, подготовьте среду сейчас.
Подготовка среды для работы с графическим процессором.
Выполните следующие действия в Cloud Shell .
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 занимает около 10 минут. После завершения дождитесь завершения работы 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:// , вставьте токен и создайте новый блокнот Python 3 в /workspace . Каждый блок кода из этого практического занятия помещается в ячейку этого блокнота.
Установите все необходимое для этого практического занятия.
!pip install --quiet xprof nvtx
nsys , профилировщик командной строки от Nsight Systems, уже входит в состав контейнера NVIDIA JAX, поэтому для дальнейшего использования CUDA ничего устанавливать не нужно. Добавьте tensorboard к команде установки, если вы предпочитаете использовать вкладку профилирования TensorBoard вместо автономного XProf в качестве средства просмотра трассировки.
Настройте и проверьте графический процессор.
Импортируйте инструменты, проверьте графический процессор и определите небольшие вспомогательные функции отображения, которые используются в остальной части этого практического задания. Если этот этап завершится неудачей, исправьте окружение, прежде чем продолжить; профилирование резервного процессора может ввести в заблуждение.
import csv
import importlib.util
import io
import os
import pathlib
import shutil
import socket
import subprocess
import sys
import tempfile
import textwrap
import time
from IPython.display import HTML, Javascript, display
import jax
import jax.numpy as jnp
import numpy as np
def require_executable(name):
"""Look up `name` on PATH and assert it's found; returns the absolute path or fails fast."""
path = shutil.which(name)
assert path, f"Required executable '{name}' not found on PATH."
return path
XPROF_BIN = require_executable("xprof")
NSYS_BIN = require_executable("nsys")
TENSORBOARD_BIN = shutil.which("tensorboard")
assert importlib.util.find_spec("nvtx"), "Required Python package 'nvtx' is missing. Install with: pip install nvtx"
def show_bars(rows, title, unit="", lower_is_better=False):
"""Render (label, value) pairs as a horizontal bar chart in HTML, scaled to the largest value."""
max_value = max(float(value) for _, value in rows) or 1.0
html = ["<div style='font-family: Arial, sans-serif; max-width: 760px;'>"]
html.append(f"<h4 style='margin: 0 0 8px 0;'>{title}</h4>")
for label, value in rows:
width = max(3, 100 * float(value) / max_value)
html.append(
"<div style='display:grid; grid-template-columns: 190px 1fr 115px; gap: 8px; "
"align-items:center; margin: 6px 0;'>"
f"<div style='font-size:13px;'>{label}</div>"
"<div style='background:#f6f8fa; border-radius:6px; overflow:hidden; height:22px;'>"
f"<div style='height:22px; width:{width:.1f}%; background:#0969da;'></div></div>"
f"<div style='font-size:13px; font-variant-numeric: tabular-nums;'>{value:.3f} {unit}</div>"
"</div>"
)
html.append(f"<div style='font-size:12px; color:#57606a;'>{'Lower' if lower_is_better else 'Higher'} is better.</div></div>")
display(HTML("".join(html)))
def show_table(headers, rows, title=None, aligns=None):
"""Render rows as an HTML table; `aligns` is an optional per-column list of "left"/"right"/"center"."""
aligns = aligns or ["left"] * len(headers)
html = ["<div style='font-family: system-ui; max-width: 980px;'>"]
if title:
html.append(f"<h4 style='margin-bottom: 8px;'>{title}</h4>")
html.append("<table style='border-collapse: collapse; width: 100%; font-size: 13px;'>")
html.append("<thead><tr>")
for h, a in zip(headers, aligns):
html.append(f"<th style='text-align:{a}; border-bottom:1px solid #d0d7de; padding:6px;'>{h}</th>")
html.append("</tr></thead><tbody>")
for row in rows:
html.append("<tr>")
for cell, a in zip(row, aligns):
html.append(f"<td style='text-align:{a}; border-bottom:1px solid #edf0f2; padding:6px; vertical-align:top; white-space:nowrap;'>{cell}</td>")
html.append("</tr>")
html.append("</tbody></table></div>")
display(HTML("".join(html)))
def show_file_list(root, limit=12):
"""Print files under `root` with their sizes (KB); truncates after `limit` entries."""
root = pathlib.Path(root)
files = [p for p in sorted(root.rglob("*")) if p.is_file()]
for p in files[:limit]:
print(f"{p.stat().st_size / 1024:8.1f} KB {p.relative_to(root)}")
if len(files) > limit:
print(f"... {len(files) - limit} more files")
devices = jax.devices()
gpu_devices = [d for d in devices if d.platform == "gpu"]
print(f"JAX version: {jax.__version__}")
print(f"Default backend: {jax.default_backend()}")
print(f"Devices: {devices}")
print(f"xprof: {XPROF_BIN}")
print(f"tensorboard: {TENSORBOARD_BIN or 'not found, XProf standalone is enough'}")
print(f"nsys: {NSYS_BIN}")
assert gpu_devices, f"This lesson assumes a GPU backend. Available devices: {devices}"
print(f"GPU devices: {gpu_devices}")
Эта ячейка выводит разрешенные пути для xprof и nsys , а также список устройств JAX, после чего проверяет, что по крайней мере одно из них является графическим процессором. Если поиск xprof не удается, повторно выполните команду pip install, указанную выше, и перезапустите ядро.
3. Создайте рабочую нагрузку для профилирования.
Вам нужно что-то, что стоило бы профилировать, но достаточно небольшое, чтобы повторить запуск за считанные секунды. В ячейке ниже определен двухслойный MLP, функция потерь среднеквадратичной ошибки, градиенты из jax.value_and_grad и простое обновление SGD, все это объединено в один JIT-компилированный train_step .
Последние несколько строк так же важны, как и сама модель. Вызов функции разогрева компилирует train_step один раз, поэтому каждая последующая ячейка измерения времени измеряет выполнение, а не случайно настройку.
BATCH = 256
IN_DIM = 1024
HIDDEN = 1024
OUT_DIM = 256
LR = 1e-3
def init_params(key):
"""Initialize the two-layer MLP weights this lesson profiles."""
k1, k2 = jax.random.split(key)
return {
"w1": jax.random.normal(k1, (IN_DIM, HIDDEN), dtype=jnp.float32) * 0.02,
"w2": jax.random.normal(k2, (HIDDEN, OUT_DIM), dtype=jnp.float32) * 0.02,
}
def make_batch(key, batch_size=BATCH):
"""Generate a random (x, target) batch with the standard input/output dims."""
kx, ky = jax.random.split(key)
x = jax.random.normal(kx, (batch_size, IN_DIM), dtype=jnp.float32)
y = jax.random.normal(ky, (batch_size, OUT_DIM), dtype=jnp.float32)
return x, y
def loss_fn(params, batch):
"""Forward pass plus MSE loss"""
x, target = batch
hidden = jax.nn.gelu(x @ params["w1"])
pred = hidden @ params["w2"]
return jnp.mean((pred - target) ** 2)
# This is the function we will profile.
@jax.jit
def train_step(params, batch):
"""One compiled SGD step"""
loss, grads = jax.value_and_grad(loss_fn)(params, batch)
params = jax.tree.map(lambda p, g: p - LR * g, params, grads)
return params, loss
key = jax.random.key(0)
params = init_params(key)
batch = make_batch(jax.random.fold_in(key, 1))
# Warm up once so later timing is on execution
params, loss = train_step(params, batch)
jax.block_until_ready((params, loss))
print(f"x shape/device: {batch[0].shape} on {batch[0].device}")
print(f"target shape/device: {batch[1].shape} on {batch[1].device}")
print(f"w1 shape/device: {params['w1'].shape} on {params['w1'].device}")
print(f"warmup loss: {float(loss):.4f}")
Каждая строка вывода должна содержать имя устройства CUDA: входы в точке (256, 1024) , целевые устройства в точке (256, 256) и w1 в точке (1024, 1024) , за которыми следует одно значение потери при прогреве.
4. Честная изоляция компиляции от выполнения и времени.
JIT-компилированная функция имеет два совершенно разных режима:
- Первый запрос на получение новой входной сигнатуры отслеживает и компилирует код, а затем выполняет его.
- Последующие вызовы с теми же формами и типами данных повторно используют скомпилированный исполняемый файл.
Изменение формы пакета приводит к созданию новой сигнатуры, поэтому JAX приходится компилировать заново. Это основная причина многочисленных зависаний в процессе обучения.
Измерьте затраты на компиляцию относительно кэшированного выполнения.
Здесь вы видите один и тот же train_step трижды: один раз с очищенным кэшем, один раз с прогретым кэшем и один раз с пакетом данных другой формы.
jax.clear_caches()
compile_params = init_params(jax.random.key(101))
compile_batch = make_batch(jax.random.key(102), batch_size=BATCH)
# Trace + compile + execute.
t0 = time.perf_counter()
compile_params, compile_loss = train_step(compile_params, compile_batch)
jax.block_until_ready((compile_params, compile_loss))
first_ms = (time.perf_counter() - t0) * 1000
# Execute using the cached compiled executable.
t0 = time.perf_counter()
compile_params, compile_loss = train_step(compile_params, compile_batch)
jax.block_until_ready((compile_params, compile_loss))
cached_ms = (time.perf_counter() - t0) * 1000
# Different batch shape which triggers another compile.
smaller_batch = make_batch(jax.random.key(103), batch_size=BATCH // 2)
t0 = time.perf_counter()
_, shape_loss = train_step(compile_params, smaller_batch)
shape_change_ms = (time.perf_counter() - t0) * 1000
print(f"First call, same shape (compile + execute): {first_ms:8.2f} ms")
print(f"Cached call, same shape (execute only): {cached_ms:8.2f} ms")
print(f"New batch shape (compile + execute): {shape_change_ms:8.2f} ms")
show_bars(
[
("first call", first_ms),
("cached call", cached_ms),
("new shape", shape_change_ms),
],
title="Compilation Cost vs. Cached Execution",
unit="ms",
lower_is_better=True,
)
# Re-warm the original shape.
params, loss = train_step(params, batch)
jax.block_until_ready((params, loss))
В результате вы получите три числа и столбчатую диаграмму. Два из трех столбцов — первый вызов и новая форма — включают в себя компиляцию, поэтому оба должны быть значительно больше, чем кэшированный вызов. Только средний столбец отражает реальную стоимость шага обучения в стационарном режиме.
Заблокируйте время, прежде чем остановить отсчет.
JAX отправляет задачи на графический процессор асинхронно, поэтому Python может завершить работу до того, как графический процессор закончит. Без block_until_ready() обычно измеряется время, затраченное Python на добавление задачи в очередь, а не время, затраченное графическим процессором на её выполнение.
Здесь вы видите один и тот же этап обучения дважды: один раз с ошибкой, и один раз с блокировкой результата.
t0 = time.perf_counter()
params_dispatch, loss_dispatch = train_step(params, batch)
dispatch_ms = (time.perf_counter() - t0) * 1000
t0 = time.perf_counter()
params_ready, loss_ready = train_step(params, batch)
jax.block_until_ready((params_ready, loss_ready))
ready_ms = (time.perf_counter() - t0) * 1000
print(f"Dispatch-only timing: {dispatch_ms:.3f} ms")
print(f"Blocked timing: {ready_ms:.3f} ms")
show_bars(
[("dispatch only", dispatch_ms), ("block_until_ready", ready_ms)],
title="Timing the Same JAX Step",
unit="ms",
lower_is_better=True,
)
Показатель скорости обработки только диспетчерских вызовов должен быть короче, чем второй. Меньшее значение не означает более быструю программу, это неизмеримый показатель.
5. Исправить шаблон «слишком много мелких операций».
Ещё одна распространённая проблема производительности графического процессора — запуск множества небольших операций из Python. Каждая небольшая операция JAX влечет за собой накладные расходы на диспетчеризацию в Python и может привести к созданию небольшого ядра графического процессора, которое завершится до того, как устройство будет должным образом занято.
jax.jit полезен, поскольку позволяет XLA видеть всю цепочку и объединять или планировать ее как единый скомпилированный блок. Две приведенные ниже функции выполняют одни и те же вычисления. Первая отправляет операции из Python, а вторая компилирует цепочку.
SMALL_OP_STEPS = 50
small_x = jnp.ones((4096,), dtype=jnp.float32)
def many_small_ops(x):
"""Un-jitted chain of small operations."""
# Each loop iteration dispatches JAX operations from Python.
y = x
for _ in range(SMALL_OP_STEPS):
y = jnp.sin(y) + 0.01 * y
return y
@jax.jit
def compiled_chain(x):
"""Same operations under `jax.jit` so XLA can fuse."""
# JAX traces the whole chain, and XLA can optimize it as one compiled computation.
y = x
for _ in range(SMALL_OP_STEPS):
y = jnp.sin(y) + 0.01 * y
return y
_ = compiled_chain(small_x).block_until_ready()
# Time the non-jitted chain.
t0 = time.perf_counter()
_ = many_small_ops(small_x).block_until_ready()
small_ops_ms = (time.perf_counter() - t0) * 1000
# Time the compiled chain.
t0 = time.perf_counter()
_ = compiled_chain(small_x).block_until_ready()
compiled_chain_ms = (time.perf_counter() - t0) * 1000
print(f"Many small Python-dispatched ops: {small_ops_ms:8.3f} ms")
print(f"Compiled chain with jax.jit: {compiled_chain_ms:8.3f} ms")
show_bars(
[("many small ops", small_ops_ms), ("compiled chain", compiled_chain_ms)],
title="Too Many Small Operations",
unit="ms",
lower_is_better=True,
)
Скомпилированная цепочка должна быть короче. Массив здесь содержит всего 4096 элементов, и цикл выполняется 50 раз, поэтому в некомпилированной версии накладные расходы Python на диспетчеризацию увеличиваются в 50 раз при очень небольших арифметических операциях каждый раз. Это тот шаблон, который нужно распознать в собственном коде: цикл Python по небольшим операциям JAX, который вместо этого можно было бы выполнить с помощью одной скомпилированной функции.
6. Запись трассировки профилировщика JAX.
Показатели времени указывают на медленную работу. Трассировка же показывает , что происходило во времени : когда был активен Python, когда компилировался XLA, когда GPU выполнял ядра и где были пробелы.
Исходные трассировки содержат множество низкоуровневых названий операций, поэтому перед записью трассировки необходимо аннотировать рабочую нагрузку. Разметка каждого шага обучения позволит вам использовать удобочитаемые ориентиры для навигации. JAX предлагает три типа аннотаций:
-
StepTraceAnnotationобозначает повторяющиеся шаги, например, итерации обучения. -
TraceAnnotationприсваивает имя области внутри этапа, например, подготовки пакета или обновления оптимизатора. -
annotate_functionуказывает имя функции Python в трассировке.
Эта ячейка отслеживает шесть аннотированных этапов обучения и сохраняет их во временный каталог:
trace_dir = pathlib.Path(tempfile.mkdtemp(prefix="jax-trace-"))
trace_params = params
trace_key = key
with jax.profiler.trace(str(trace_dir), create_perfetto_link=False):
for step_num in range(6):
with jax.profiler.StepTraceAnnotation("train_step", step_num=step_num):
with jax.profiler.TraceAnnotation("make_batch"):
trace_key = jax.random.fold_in(trace_key, step_num + 10)
trace_batch = make_batch(trace_key)
with jax.profiler.TraceAnnotation("sgd_update"):
trace_params, trace_loss = train_step(trace_params, trace_batch)
jax.block_until_ready((trace_params, trace_loss))
print(f"Trace directory: {trace_dir}")
show_file_list(trace_dir)
Вы должны увидеть путь к каталогу трассировки, за которым следует краткий список файлов трассировки, записанных в него, с указанием размера каждого из них. Именно на этот путь к каталогу вы будете указывать программе просмотра на следующем шаге, поэтому не забывайте оставлять вывод ячейки видимым.
7. Откройте трассировку в XProf.
XProf — это пользовательский интерфейс профилировщика, лежащий в основе формата трассировки JAX. Он считывает каталог трассировки напрямую и предоставляет веб-интерфейс с просмотрщиком трассировки, просмотрщиком памяти, просмотрщиком графиков и статистикой операций.
Запустите его внутри капсулы:
xprof_port = 6007
xprof_proc = subprocess.Popen(
[XPROF_BIN, "--port", str(xprof_port), str(trace_dir)],
stdout=subprocess.DEVNULL,
stderr=subprocess.DEVNULL,
)
time.sleep(3)
xprof_url = f"http://localhost:{xprof_port}"
print(f"XProf is running at {xprof_url} (PID {xprof_proc.pid})")
print(f"Stop it later with: kill {xprof_proc.pid}")
display(Javascript(f'window.open("{xprof_url}", "_blank");'))
display(HTML(f'<a href="{xprof_url}" target="_blank" rel="noopener">Open XProf</a>'))
В ячейке отображается URL-адрес, по которому работает XProf, идентификатор процесса и команда kill , которая его останавливает. Запишите этот PID — он понадобится вам на этапе очистки .
Подключитесь к зрителю со своего ноутбука.
По умолчанию ничто внутри Pod-а недоступно с вашего компьютера. kubectl port-forward открывает туннель из Cloud Shell к порту внутри Pod-а, а веб-предварительная версия Cloud Shell публикует этот туннель в вашем браузере.
Откройте новую вкладку Cloud Shell — команда будет выполняться в фоновом режиме, пока вы её не остановите, а ваш блокнот тем временем продолжит работать в Pod:
kubectl port-forward pod/jax-jupyter 8080:6007
Затем нажмите кнопку «Предварительный просмотр веб-страницы — Предварительный просмотр на порту 8080» на панели инструментов Cloud Shell, и вы сможете получить доступ к пользовательскому интерфейсу.
Прочитайте трассировку
После загрузки вкладки «Предварительный просмотр веб-страниц» XProf:
- Выберите маршрут из выпадающего списка «Маршруты» .
- Откройте Инструменты - trace_viewer .
- Найдите диапазоны
train_step— это именаStepTraceAnnotationиз предыдущего шага. - Обратите внимание на фрагменты компиляции, разрывы в работе графического процессора, копии и активность ядра.
Контрольный список для чтения следов
В таблице ниже показано, как визуальный паттерн на диаграмме запускает следующее действие:
Визуальная подсказка | Что это, вероятно, означает | Что попробовать дальше? |
Длительный процесс компиляции перед первым шагом. | Обычная JIT-компиляция первого вызова | Перед измерением необходимо немного разогреться. |
Компиляция фрагментов между шагами | Рекомпиляция | Проверьте возможность изменения форм, типов данных или статических аргументов из практического задания 2. |
В рядах графических процессоров имеются белые промежутки. | Хост не подает данные на графический процессор. | Обратите внимание на загрузку данных, вывод на экран, |
Много крошечных зерен | Запуск избыточных ресурсов или слишком малых скомпилированных областей | Функция JIT (точно в срок) — более масштабная; пакетная работа. |
Активность Memcpy между шагами | Передача данных между хост-устройством | Храните метрики на устройстве; реже ведите журналы; используйте функцию защиты от передачи данных. |
Высокий пик памяти | В качестве основных факторов могут преобладать активации или временные буферы. | Откройте |
Стоит упомянуть еще три точки входа, хотя в этом практическом занятии они не используются: start_trace() и stop_trace() для программных областей трассировки, которые не помещаются в блок with jax.profiler.trace(...) , start_server() плюс python -m jax.collect_profile для профилирования длительно выполняющихся задач, и Perfetto export для открытия трассировок в пользовательском интерфейсе Perfetto. Руководство по профилированию JAX охватывает все три.
8. Диагностика передач данных на хосте и размера пакета.
Два из перечисленных выше пунктов достаточно распространены, чтобы их стоило воспроизвести: случайные передачи данных с хоста и слишком малый размер пакета, не позволяющий заполнить графический процессор.
Передача данных между хост-устройством
Передача значения JAX обратно в Python внутри цикла приводит к синхронизации. Распространенные примеры: float(loss) , .item() , np.asarray(...) и вывод массивов на экран.
Этот код сравнивает два стиля логирования. Плохой вариант преобразует значение функции потерь в число float Python на каждом шаге. Лучший же хранит значения функции потерь в виде массивов JAX и синхронизирует их один раз в конце.
N = 30
# Convert the loss to a Python float every step.
t0 = time.perf_counter()
bad_params = params
bad_losses = []
for _ in range(N):
bad_params, bad_loss = train_step(bad_params, batch)
bad_losses.append(float(bad_loss))
bad_ms = (time.perf_counter() - t0) * 1000 / N
# Keep metrics as JAX arrays and synchronize once.
t0 = time.perf_counter()
good_params = params
good_losses = []
for _ in range(N):
good_params, good_loss = train_step(good_params, batch)
good_losses.append(good_loss)
jax.block_until_ready((good_params, good_losses))
good_ms = (time.perf_counter() - t0) * 1000 / N
print(f"float(loss) every step: {bad_ms:.3f} ms / step")
print(f"deferred sync: {good_ms:.3f} ms / step")
show_bars(
[("float(loss) every step", bad_ms), ("deferred sync", good_ms)],
title="Cost of Pulling Metrics to Python",
unit="ms/step",
lower_is_better=True,
)
# Transfer guard can catch accidental transfers.
try:
_, guard_loss = train_step(params, batch)
guard_loss.block_until_ready()
with jax.transfer_guard("disallow"):
_ = float(guard_loss)
except RuntimeError as e:
print("\nTransfer guard caught an implicit transfer:")
print(str(e).splitlines()[0])
Стоимость за шаг в версии с float(loss) должна быть больше, чем значение в двух столбцах, а в конце ячейки должна быть выведена первая строка RuntimeError , сгенерированной механизмом передачи.
Размер партии и пропускная способность
Небольшие пакеты данных часто не обеспечивают достаточной параллельной работы для графического процессора. Пропускная способность обычно улучшается по мере роста пакета, а затем стабилизируется, когда графический процессор перегружен или объем памяти становится пределом.
Размер пакета также влияет на объем памяти. Приведенный ниже пример включает простую оценку буферов пакетной формы в этом прямом проходе: входные данные, скрытые активационные данные и выходные данные. В реальном обучении используется больше данных, поскольку градиенты и временные буферы также учитываются.
@jax.jit
def forward_only(params, x):
"""Forward pass only."""
hidden = jax.nn.gelu(x @ params["w1"])
return hidden @ params["w2"]
batch_results = []
bytes_per_float32 = np.dtype(np.float32).itemsize
for batch_size in (1, 8, 32, 128, 256, 512, 1024):
xb = jax.random.normal(jax.random.key(batch_size), (batch_size, IN_DIM), dtype=jnp.float32)
_ = forward_only(params, xb).block_until_ready()
reps = 80 if batch_size <= 128 else 30
# Async dispatch lets JAX queue all `reps` calls without blocking. We block once at the end and divide by reps, so this measures *amortized* time per call when the executor stays busy.
t0 = time.perf_counter()
for _ in range(reps):
yb = forward_only(params, xb)
yb.block_until_ready()
ms = (time.perf_counter() - t0) * 1000 / reps
examples_per_sec = batch_size / (ms / 1000)
# Simple memory estimate for batch-shaped forward buffers.
estimated_forward_mib = batch_size * (IN_DIM + HIDDEN + OUT_DIM) * bytes_per_float32 / 2**20
batch_results.append((batch_size, ms, examples_per_sec, estimated_forward_mib))
show_table(
["Batch", "ms/call (avg)", "examples/sec", "estimated batch buffers"],
[
(batch_size, f"{ms:.3f}", f"{examples_per_sec:,.0f}", f"{estimated_mib:.1f} MiB")
for batch_size, ms, examples_per_sec, estimated_mib in batch_results
],
title="Batch Size, Throughput, and Memory",
aligns=["right", "right", "right", "right"],
)
show_bars(
[(f"batch {batch_size}", examples_per_sec) for batch_size, _, examples_per_sec, _ in batch_results],
title="Throughput by Batch Size",
unit="ex/s",
lower_is_better=False,
)
Вы получаете таблицу из семи строк и гистограмму. Обратите внимание на столбец examples/sec , а не ms/call : пропускная способность должна резко возрастать на протяжении первых нескольких размеров пакетов, а затем стабилизироваться по мере насыщения графического процессора. Столбец ms/call растет на протяжении всего времени, что и ожидается.
9. Проверьте уровень нехватки памяти графического процессора.
JAX обычно предварительно выделяет определенный процент памяти графического процессора при первом использовании. Это делается намеренно, поскольку снижает накладные расходы на выделение памяти и фрагментацию. Это также означает, что nvidia-smi может выглядеть почти заполненным, даже если ваша модель очень маленькая, что заставляет многих искать утечку памяти, которой не существует.
Используйте memory_stats для быстрого просмотра данных на уровне процессов, а затем инструменты XProf для анализа памяти для более глубокого изучения.
def gib(value):
"""Convert a byte count to gibibytes."""
return value / 2**30
print(f"{'device':<26} {'limit':>12} {'in use':>12} {'peak':>12}")
print("-" * 66)
for device in gpu_devices:
stats = device.memory_stats()
if not stats:
print(f"{str(device):<26} memory_stats unavailable")
continue
limit = stats.get("bytes_limit")
in_use = stats.get("bytes_in_use")
peak = stats.get("peak_bytes_in_use")
limit_s = f"{gib(limit):.2f} GiB" if limit is not None else "n/a"
in_use_s = f"{gib(in_use):.2f} GiB" if in_use is not None else "n/a"
peak_s = f"{gib(peak):.2f} GiB" if peak is not None else "n/a"
print(f"{str(device):<26} {limit_s:>12} {in_use_s:>12} {peak_s:>12}")
На каждый видимый графический процессор должна приходиться одна строка. Столбец limit отображает резервирование памяти распределителем, а не общий объем памяти карты, и in use этого небольшого MLP он должен составлять лишь небольшую его часть.
Настройки памяти необходимо задать до импорта JAX . В ноутбуке это означает установку параметров перед запуском ядра, а затем перезапуск ядра.
Переменная | Пример | Использовать при |
| | Вы используете один и тот же графический процессор и хотите, чтобы JAX резервировал меньше памяти. |
| | Вы хотите получать ресурсы по запросу, принимая на себя больший риск фрагментации. |
| | Вы отлаживаете память и хотите освободить её; это слишком медленно для обычного обучения. |
В этой ячейке показано, какие из этих параметров установлены в текущем ядре:
for name in (
"XLA_PYTHON_CLIENT_MEM_FRACTION",
"XLA_PYTHON_CLIENT_PREALLOCATE",
"XLA_PYTHON_CLIENT_ALLOCATOR",
):
print(f"{name}={os.environ.get(name, '<unset>')}")
print("\nExample for a shared GPU, set before launching Python/Jupyter:")
print("export XLA_PYTHON_CLIENT_MEM_FRACTION=0.50")
Все три, скорее всего, будут напечатаны. Это означает, что у вас установлены параметры по умолчанию: предварительное выделение ресурсов включено, доля 75%.
10. Создание временной шкалы CUDA с помощью Nsight Systems.
XProf — это подходящий первый профилировщик для JAX. Nsight Systems — это второй вариант, когда вам нужна временная шкала CUDA: потоки, вызовы API CUDA, ядра, копирование памяти, вызовы cuBLAS и cuDNN, и, в конечном итоге, NCCL.
Рабочий процесс состоит из четырех частей:
- Оформите задачу в виде короткого скрипта.
- Добавьте диапазоны NVTX вокруг интересующих вас этапов.
- Запустите скрипт с
nsys profile. - Откройте файл
.nsys-repс помощью графического интерфейса Nsight Systems.
Напишите скрипт для захвата данных.
Профилирование ноутбука напрямую — непростая задача, поэтому приведенный ниже код вместо этого создает небольшой автономный скрипт. Он выполняет один цикл прогрева, а затем использует диапазоны NVTX для обозначения области, заслуживающей захвата данных.
Вам не нужно читать каждую строку. Большая часть текста повторяет ту же самую небольшую модель, что и раньше, чтобы nsys мог профилировать новый процесс Python.
nsight_dir = pathlib.Path(tempfile.mkdtemp(prefix="jax-nsight-"))
script_path = nsight_dir / "nsight_train_step.py"
report_base = nsight_dir / "jax_train_step"
report_path = pathlib.Path(f"{report_base}.nsys-rep")
script_source = f"""
import nvtx
import jax
import jax.numpy as jnp
BATCH = {BATCH}
IN_DIM = {IN_DIM}
HIDDEN = {HIDDEN}
OUT_DIM = {OUT_DIM}
LR = {LR}
def init_params(key):
k1, k2 = jax.random.split(key)
return {{ "{{" }}
"w1": jax.random.normal(k1, (IN_DIM, HIDDEN), dtype=jnp.float32) * 0.02,
"w2": jax.random.normal(k2, (HIDDEN, OUT_DIM), dtype=jnp.float32) * 0.02,
{{ "}}" }}
def make_batch(key):
kx, ky = jax.random.split(key)
return (
jax.random.normal(kx, (BATCH, IN_DIM), dtype=jnp.float32),
jax.random.normal(ky, (BATCH, OUT_DIM), dtype=jnp.float32),
)
def loss_fn(params, batch):
x, target = batch
hidden = jax.nn.gelu(x @ params["w1"])
pred = hidden @ params["w2"]
return jnp.mean((pred - target) ** 2)
@jax.jit
def train_step(params, batch):
loss, grads = jax.value_and_grad(loss_fn)(params, batch)
params = jax.tree.map(lambda p, g: p - LR * g, params, grads)
return params, loss
key = jax.random.key(0)
params = init_params(key)
batch = make_batch(key)
# Warm up before the region we care about.
params, loss = train_step(params, batch)
jax.block_until_ready((params, loss))
with nvtx.annotate("profile_region", domain="jax_course"):
for step in range(8):
with nvtx.annotate(f"train_step_{{ "{{" }}step{{ "}}" }}", domain="jax_course"):
key = jax.random.fold_in(key, step)
batch = make_batch(key)
params, loss = train_step(params, batch)
jax.block_until_ready((params, loss))
"""
script_path.write_text(textwrap.dedent(script_source))
print(f"Wrote script: {script_path}")
print(f"Report path: {report_path}")
Ячейка выводит на экран два пути, которые она будет использовать. Пока ничего не выполнялось.
Run Night Systems
Приведенный ниже вызов запускает сбор данных только в начале диапазона NVTX profile_region@jax_course и останавливается в конце этого диапазона. Это позволяет сосредоточить внимание на выполнении отчета после предварительной подготовки, а не скрывать его под запуском процесса и JIT-компиляцией.
import base64
nsys_cmd = [
NSYS_BIN,
"profile",
"--trace=cuda,nvtx,osrt,cudnn,cublas",
"--capture-range=nvtx",
"--capture-range-end=stop",
"--nvtx-capture=profile_region@jax_course",
"--force-overwrite=true",
f"--output={report_base}",
"-e",
"NSYS_NVTX_PROFILER_REGISTER_ONLY=0",
sys.executable,
str(script_path),
]
print("Running Nsight Systems:")
print(" ".join(nsys_cmd))
nsys_result = subprocess.run(nsys_cmd, capture_output=True, text=True, check=False)
if nsys_result.returncode != 0:
print("nsys profile failed")
print(nsys_result.stdout[-2000:])
print(nsys_result.stderr[-4000:])
raise RuntimeError("Nsight Systems capture failed. This lesson requires nsys to run successfully.")
print(f"Nsight report: {report_path} ({report_path.stat().st_size / 1024:.1f} KB)")
b64 = base64.b64encode(report_path.read_bytes()).decode()
download_html = (
f'<a href="data:application/octet-stream;base64,{b64}" '
f'download="{report_path.name}" '
f'style="font-weight:600;">Download {report_path.name}</a>'
)
display(HTML(download_html))
Процесс захвата данных занимает некоторое время, поскольку запускается новый процесс Python, который импортирует JAX и заново компилирует train_step . После завершения вы получаете путь к отчету, его размер и ссылку для скачивания . Эта ссылка представляет собой URL-адрес данных, поэтому он работает прямо из блокнота в вашем браузере — переадресация портов не требуется.
Ознакомьтесь с хронологией Nsight.
Чтобы открыть отчет в графическом интерфейсе пользователя:
- Установите Nsight Systems на свой локальный компьютер с сайта developer.nvidia.com/nsight-systems . Это бесплатное программное обеспечение, работающее на Linux, macOS и Windows.
- Загрузите сгенерированный отчет
.nsys-repиспользуя ссылку выше. - Запустите
nsys-uiи откройте файл с помощью меню «Файл» — «Открыть» или перетащите его в окно.
Начните с этих строк:
- NVTX : найти
profile_regionиtrain_step_*. - Строки CUDA GPU : проверьте покрытие ядра и белые промежутки, где белый промежуток означает время простоя графического процессора.
- API CUDA : найдите вызовы
cudaLaunchKernel,cudaMemcpy*и синхронизации. - cuBLAS / cuDNN : вызовы библиотек matmul, convolution и attention отображаются здесь при трассировке.
- Среда выполнения ОС (
osrt) : хост ожидает, блокирует, переходит в спящий режим и выполняет другие блокирующие действия.
Для получения метрик использования графического процессора повторно запустите сбор данных с параметром --gpu-metrics-devices=cuda-visible если ваша среда разрешает сбор метрик графического процессора.
Кратко изложите содержание отчета.
Основной интерфейс Nsight — это графический интерфейс пользователя, но nsys stats предоставляет полезную текстовую сводку, которую можно прочитать. Этот код выводит таблицы ядра, операций с памятью и API CUDA.
stats_cmd = [
NSYS_BIN,
"stats",
"--quiet",
"--report",
"cuda_gpu_kern_sum,cuda_gpu_mem_time_sum,cuda_api_sum",
"--format",
"csv",
"--timeunit",
"msec",
str(report_path),
]
stats_result = subprocess.run(stats_cmd, capture_output=True, text=True, check=False)
if stats_result.returncode != 0:
print("nsys stats failed")
print(stats_result.stderr[-3000:])
else:
sections = [s for s in stats_result.stdout.strip().split("\n\n") if s.strip()]
for idx, section in enumerate(sections[:3], start=1):
reader = csv.DictReader(io.StringIO(section))
rows = list(reader)
if not rows:
continue
headers = reader.fieldnames or []
print(f"\nReport section {idx}: {headers}")
compact_rows = []
for row in rows[:8]:
name = row.get("Name") or row.get("Operation") or row.get("Name:Demangled") or ""
time_pct = row.get("Time (%)", "")
total = row.get("Total Time (ms)") or row.get("Total Time (msec)") or row.get("Total Time") or row.get("Total Time (ns)") or ""
count = row.get("Instances") or row.get("Count") or row.get("Num Calls") or row.get("Calls") or ""
compact_rows.append((time_pct, total, count, name[:90]))
show_table(["Time (%)", "Total time (ms)", "Count", "Name / operation"], compact_rows, title=f"Nsight stats section {idx}", aligns=["right", "right", "right", "left"])
Вы должны получить до трех таблиц, в каждой из которых будут отображаться восемь строк с наибольшим временем выполнения. Сначала ознакомьтесь со сводкой ядра: она ранжирует ядра GPU по общему времени, что показывает, где устройство фактически его потратило. Сравните это с таблицей операций с памятью — если копии конкурируют с ядрами за время, вы снова сталкиваетесь с проблемой передачи данных между хостом и устройством, о которой говорилось ранее.
11. Уборка
Плата за использование графического процессора взимается независимо от того, запущены ли на нем какие-либо приложения, поэтому не пропускайте этот шаг.
Сначала остановите зрителей, которых вы запустили внутри Pod: выполните команду kill команды, которые вывели ячейки XProf и TensorBoard, и нажмите Ctrl+C в любой вкладке Cloud Shell, где все еще запущена kubectl port-forward .
Удалите рабочую нагрузку Jupyter, включая балансировщик нагрузки и постоянный том:
kubectl delete -f deploy/jupyter.yaml
Уничтожьте кластер, пул узлов, VPC и учетную запись службы:
cd terraform
terraform destroy
При появлении запроса введите yes , затем подтвердите, что ничего не осталось:
gcloud container clusters list
gcloud compute instances list
Оба поля должны быть пустыми для этого проекта. Если вы создали проект только для этой серии, вы можете удалить весь проект из консоли Cloud .
12. Поздравляем!
Вы провели профилирование этапа обучения JAX от начала до конца, начиная с измерения time.perf_counter и заканчивая временной шкалой CUDA.
Что вы узнали
- Как корректно измерить время выполнения JAX с помощью
block_until_ready()и почему измерение только времени выполнения бессмысленно. - Как отделить время компиляции первого вызова от времени выполнения в кэше, и как новая форма входных данных приводит к перекомпиляции.
- Как получить трассировку с помощью
jax.profiler.traceи пометить её с помощьюStepTraceAnnotationиTraceAnnotation - Как открыть этот трассировочный файл в XProf, а также в TensorBoard в качестве альтернативного интерфейса, используя переадресацию портов Cloud Shell.
- Как распознать распространённые признаки замедления работы: перекомпиляция, слишком много мелких операций, передача данных между хостом и устройством, неэффективные размеры пакетов и нехватка памяти.
- Как проверить память графического процессора с помощью
memory_stats()и какXLA_PYTHON_CLIENT_MEM_FRACTIONизменяет резервирование памяти JAX перед запуском. - Как захватить и прочитать временную шкалу CUDA с помощью Nsight Systems, используя диапазоны NVTX, и как суммировать её с помощью
nsys stats
Следующие шаги
- Практическое занятие 4: Обучение модели на графическом процессоре с использованием JAX, Optax и Fashion-MNIST. Вы пройдете реальный цикл обучения на реальных данных, уже освоив методы профилирования, изученные в этом практическом занятии.
- Повторно запустите захват Nsight, добавив параметр
--gpu-metrics-devices=cuda-visibleвnsys_cmd, и сравните счетчики использования графического процессора с покрытием ядра, которое вы видели на временной шкале. - Перед перезапуском ядра установите значение
XLA_PYTHON_CLIENT_MEM_FRACTION=0.50, затем повторно запустите ячейкуmemory_stats()и проверьте, как изменилось значение в столбцеlimit.