1. Введение

В практическом занятии «Запуск вашей первой программы JAX на графических процессорах NVIDIA с помощью GKE» вы обернули функцию в jax.jit на графическом процессоре и увидели, что первый вызов занимает гораздо больше времени, чем все последующие. Это не случайность: JAX отслеживает вашу функцию Python с помощью абстрактных заполнителей, передает записанную программу в XLA и кэширует скомпилированный исполняемый файл, который он запускает на графическом процессоре. В этом практическом занятии вы откроете этот процесс, узнаете, что добавляет запись в кэш компиляции, и исправите две вещи, которые отнимают у пользователей JAX больше всего времени, такие как случайные перекомпиляции и управление потоком выполнения Python на основе отслеживаемых значений.
Что вы будете делать
- Отслеживание трассировки можно осуществить, поместив оператор
printиз Python внутрь скомпилированной функции. - Сопоставьте стоимость компиляции со стоимостью выполнения в кэше на графическом процессоре.
- Определите, что относится к ключу кэша компиляции и что запускает перекомпиляцию.
- Замените управляющий поток Python для отслеживаемых значений на
jnp.whereиjax.lax.cond - Для обеспечения стабильности форм используйте
jax.lax.scan, отступы и маскирование, а такжеstatic_argnums - Просмотрите результаты трассировки JAX с помощью
jax.make_jaxpr
Что вам понадобится
- Проект Google Cloud с включенной оплатой и кредитами для семинара или резервированием, покрывающим использование графического процессора.
- Квота на использование как минимум двух видеокарт NVIDIA L4 в выбранном вами регионе ( как проверить квоту на видеокарты )
- Завершение практического занятия 1: Запустите свою первую программу на JAX на графических процессорах NVIDIA с помощью GKE или аналогичной среды JAX для графических процессоров.
Примерное время выполнения: 50 минут .
2. Прежде чем начать
Выберите свой проект
В консоли Google Cloud выберите или создайте проект с включенной функцией выставления счетов.
Открытая облачная оболочка
Чтобы запустить сеанс Cloud Shell , нажмите кнопку «Активировать Cloud Shell» (значок терминала в правом верхнем углу консоли), а затем укажите в качестве исполнителя свой проект:
gcloud config set project <YOUR_PROJECT_ID>
Этот практический пример выполняется в той же среде, что и практический пример 1: Запуск вашей первой программы JAX на графических процессорах NVIDIA с помощью GKE . Если ваш кластер GKE и под JupyterLab все еще работают, перейдите к разделу «Настройка и проверка графического процессора» . В противном случае, подготовьте среду сейчас.
Подготовка среды для работы с графическим процессором.
Выполните следующие команды в 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 занимает около 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 . Каждый блок кода из этого практического занятия помещается в ячейку этого блокнота.
Настройте и проверьте графический процессор.
Импортируйте JAX, NumPy и несколько вспомогательных программ из стандартной библиотеки, затем убедитесь, что используете графический процессор (GPU).
import time
from functools import partial
import jax
import jax.numpy as jnp
import numpy as np
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}")
assert gpu_devices, f"This lab assumes a GPU backend. Available devices: {devices}"
print(f"GPU devices: {gpu_devices}")
В списке устройств по умолчанию должен отображаться gpu , а также как минимум одно CudaDevice . Для этого практического задания требуется только один графический процессор, поэтому допустимо, если узел предоставляет доступ к нескольким графическим процессорам.
3. Просмотрите трассировку вашей функции с помощью jax.jit.
При вызове обычной, не скомпилированной с помощью JIT-компилятора функции JAX каждая операция выполняется через Python и передается на графический процессор по мере ее выполнения. jax.jit меняет это. Вместо запуска вашей функции с реальными массивами, он отслеживает ее выполнение: JAX вызывает ее один раз с абстрактными заполнителями, которые содержат только форму и тип данных, и записывает каждую операцию JAX, которую вы выполняете с этими заполнителями, в промежуточное представление, называемое jaxpr .
JAX преобразует jaxpr в StableHLO, передает преобразованную программу в XLA , и XLA компилирует оптимизированный исполняемый файл для целевого устройства. XLA может объединять операции, но скомпилированная функция все еще может преобразовываться в несколько ядер GPU. С этого момента вызов функции переходит непосредственно к кэшированному исполняемому файлу.
Таким образом, каждый JIT-вызов состоит из трех фаз:
Фаза | Что это делает | Что происходит |
След | Запишите результаты вычислений. | Python запускается один раз, а JAX записывает каждую операцию в абстрактные заполнители в файл |
Компиляция | Преобразовать в исполняемый файл для графического процессора. | JAX понижает уровень |
Выполнять | Используйте повторно кэшированный исполняемый файл. | При каждом последующем вызове с совпадающими формами и типами данных трассировка и компиляция пропускаются, и запускается кэшированная программа. |
Трассировка (Trace) объясняет, почему оператор print в Python внутри JIT-функции срабатывает только при первом вызове. Трассировка и компиляция вместе объясняют медленную работу первого вызова. Выполнение (Execute) объясняет быструю работу каждого последующего вызова.
Посмотрите, как работает трассировка.
Убедитесь сами, что тело функции выполняется только один раз для каждой входной сигнатуры. Поместите внутрь функции print на уровне Python: он выполняется во время трассировки, но не является частью скомпилированной программы для графического процессора, поэтому последующие вызовы с той же формой и типом данных ничего не выведут.
@jax.jit
def f(x):
"""Jitted demo function that prints during tracing so we can see exactly when JAX retraces."""
# This print runs during tracing only not on every GPU execution.
print(f" tracing with shape={x.shape} dtype={x.dtype}")
return x ** 2 + 1
print("Call 1 (new shape):")
_ = f(jnp.arange(4, dtype=jnp.float32)).block_until_ready()
print("Call 2 (same shape):")
_ = f(jnp.arange(4, dtype=jnp.float32)).block_until_ready()
print("Call 3 (new shape):")
_ = f(jnp.arange(5, dtype=jnp.float32)).block_until_ready()
Вы должны увидеть результат, похожий на следующий:
Call 1 (new shape): tracing with shape=(4,) dtype=float32 Call 2 (same shape): Call 3 (new shape): tracing with shape=(5,) dtype=float32
Вывод срабатывает при первом вызове, когда JAX впервые видит форму (4,) с типом данных float32 , и при третьем вызове, когда он впервые видит форму (5,) . При втором вызове JAX обнаруживает существующий скомпилированный исполняемый файл и пропускает как трассировку, так и компиляцию.
4. Проведите оценку компиляции на основе кэшированного выполнения.
Новая стоимость подписи реальна, и именно отсюда берутся сообщения о медленной работе JAX. Измерьте, какая часть первого вызова приходится на компиляцию, а какая — на выполнение.
Приведённая ниже функция объединяет 20 нелинейных зависимостей, в результате чего компиляция оказывается заметно дороже выполнения.
def heavy(x):
"""20 chained nonlinearities so the first-call compilation is visibly more expensive than the cached execution."""
y = x
for _ in range(20):
y = jnp.sin(y) * jnp.cos(y) + jnp.tanh(y)
return y
heavy_jit = jax.jit(heavy)
x = jnp.arange(1_000_000, dtype=jnp.float32)
# Empty in-process cache so a re-run shows the first-call compile cost again.
jax.clear_caches()
t0 = time.perf_counter()
_ = heavy_jit(x).block_until_ready()
first_ms = (time.perf_counter() - t0) * 1000
t0 = time.perf_counter()
for _ in range(20):
_ = heavy_jit(x).block_until_ready()
cached_ms = (time.perf_counter() - t0) * 1000 / 20
print(f"First call (compile + execute): {first_ms:8.2f} ms")
print(f"Cached call (execute only): {cached_ms:8.2f} ms")
print(f"Compilation cost (approx): {first_ms - cached_ms:8.2f} ms")
Вы должны увидеть время первого вызова, которое значительно превышает время вызова из кэша. Разница между ними примерно равна времени, затраченному XLA на компиляцию.
Для крошечных функций этот промежуток составляет несколько десятков миллисекунд; для полного этапа обучения трансформера он легко может составлять несколько секунд. Хорошая новость в том, что вы платите один раз за каждую комбинацию формы и типа данных , а не один раз за каждый вызов. Остальная часть этого практического занятия посвящена тому, как не платить чаще, чем это необходимо.
5. Выясните, что вызывает перекомпиляцию.
JAX формирует кэш компиляции на основе структурной сигнатуры входных данных: их формы, типов данных и любых аргументов, помеченных как статические. Если сигнатура совпадает с той, которую JAX уже видел ранее, запускается кэшированный исполняемый файл. Если что-либо изменяется, JAX выполняет трассировку и компиляцию заново.
Три фактора обычно приводят к перекомпиляции:
Ключевая часть кэша | Какие изменения | Эффект |
Форма | Разная форма | |
D-тип | Различные типы данных | |
Статический аргумент | Различные статические значения | Значение любого аргумента |
Значения обычных массивов на входе не имеют значения. Два массива float32 (32, 128) с совершенно разным содержимым попадают в один и тот же скомпилированный исполняемый файл.
Посмотрите, как происходит перекомпиляция. Приведённый ниже цикл вызывает одну JIT-скомпилированную функцию с пятью массивами, три из которых имеют форму, ранее не встречавшуюся в JAX.
@jax.jit
def f(x):
"""Simple jitted scalar function used to demonstrate one compile per new input shape (a new dtype would trigger the same recompile)."""
return jnp.sum(x ** 2)
# clear JAX's in-process compilation cache.
jax.clear_caches()
# Feed in a few different shapes and measure each call.
shapes = [(100,), (200,), (100,), (200,), (300,)]
for s in shapes:
x = jnp.ones(s, dtype=jnp.float32)
t0 = time.perf_counter()
_ = f(x).block_until_ready()
dt = (time.perf_counter() - t0) * 1000
print(f"shape={str(s):8s} {dt:7.2f} ms")
Вы должны увидеть три медленных вызова, по одному для каждой новой фигуры, и два быстрых вызова для повторяющихся (100,) и (200,) .
В реальных рабочих нагрузках это происходит случайно постоянно: последовательности переменной длины, последний пакет в эпохе, неровный результат токенизации. Решение почти во всех случаях — не допускать изменения формы .
6. Замените управляющий поток Python в отслеживаемых значениях.
Во время трассировки входные данные вашей функции не являются конкретными массивами. Это абстрактные значения с известной формой и типом данных. Любая конструкция Python, требующая численного сравнения этого содержимого ( if , while , bool(x) , int(x) ), нарушает трассировку.
Вот как это выглядит. Эта функция ReLU изначально неправильная:
@jax.jit
def relu_bad(x):
"""ReLU using a Python `if` on a traced value with JIT errors out at trace time."""
if x > 0:
return x
return jnp.zeros_like(x)
try:
print(relu_bad(jnp.array(1.0)))
except Exception as e:
print(f"{type(e).__name__}: {str(e).splitlines()[0]}")
Вы должны увидеть ошибку TracerBoolConversionError . Сообщение указывает на условие if : JAX не может определить, какую ветвь сохранить, если значение является абстрактным.
Представьте свой выбор в виде данных с помощью jnp.where
Решение состоит в том, чтобы выразить выбор как данные , а не как управление потоком выполнения Python. Для небольшого поэлементного выбора, такого как ReLU, jnp.where — наиболее удобный инструмент. Обе ветви всегда выполняются, и предикат указывает JAX, какую из них использовать в каждой позиции.
@jax.jit
def relu(x):
"""ReLU using `jnp.where` with both branches are computed so tracing works."""
return jnp.where(x > 0, x, 0.0)
print(relu(jnp.array([-1.0, -0.5, 0.0, 0.5, 1.0])))
На этот раз ошибки нет. Два отрицательных значения и ноль возвращаются как 0. , а 0.5 и 1.0 проходят без изменений.
Выберите реальную ветку с помощью jax.lax.cond
Для ветвей, вычисляющих совершенно разные вещи, где запуск обеих функций был бы неэффективным, используйте jax.lax.cond . Обе функции ветвления отслеживаются, но во время выполнения lax.cond представляет собой условную операцию XLA, поэтому обычно выполняется только выбранная ветвь. Одно замечание: в vmap cond может быть преобразована в операцию, подобную select, а не в реальную ветвь.
@jax.jit
def soft_or_sharp(x, sharp):
"""Switch between hard ReLU and softplus inside the compiled graph via `lax.cond`, controlled by a traced bool."""
# `sharp` is a scalar bool and lax.cond compiles to a real if-then-else
return jax.lax.cond(
sharp,
lambda x: jnp.where(x > 0, x, 0.0),
lambda x: jax.nn.softplus(x),
x,
)
x = jnp.array([-1.0, 0.5, 2.0])
print(f"sharp=True: {soft_or_sharp(x, jnp.array(True))}")
print(f"sharp=False: {soft_or_sharp(x, jnp.array(False))}")
Вы должны увидеть два разных массива: строка sharp=True ограничивает отрицательное значение нулем, а строка sharp=False возвращает небольшие положительные значения softplus повсюду.
7. Сохраняйте длинные петли компактными с помощью lax.scan.
Для циклов по отслеживаемым данным используйте структурированные примитивы управления потоком выполнения, такие как jax.lax.while_loop , jax.lax.fori_loop и jax.lax.scan .
На самом деле, цикл for Python со статическим ограничением допустим внутри jit , но JAX разворачивает цикл во время трассировки. Это означает, что 200 итераций цикла превращаются примерно в 200 повторяющихся блоков в скомпилированной программе. lax.scan сохраняет цикл как примитив, подобный циклу, что обычно значительно ускоряет компиляцию длинных циклов фиксированной длины.
# python_for_loop compile time scales with NUM_STEPS while scan_loop compile time stays roughly constant. Try NUM_STEPS = 2000 to see the gap widen.
NUM_STEPS = 200
@jax.jit
def python_for_loop(x):
"""Python `for` loop inside jit."""
y = x
for _ in range(NUM_STEPS):
y = jnp.sin(y) + 0.01 * y
return y
@jax.jit
def scan_loop(x):
"""Same logic expressed with `lax.scan`."""
def body(y, _):
y = jnp.sin(y) + 0.01 * y
return y, None
y, _ = jax.lax.scan(body, x, xs=None, length=NUM_STEPS)
return y
x = jnp.ones((1024,), dtype=jnp.float32)
jax.clear_caches()
t0 = time.perf_counter()
_ = python_for_loop(x).block_until_ready()
python_for_ms = (time.perf_counter() - t0) * 1000
jax.clear_caches()
t0 = time.perf_counter()
_ = scan_loop(x).block_until_ready()
scan_ms = (time.perf_counter() - t0) * 1000
print(f"Python for loop first call: {python_for_ms:8.2f} ms")
print(f"lax.scan first call: {scan_ms:8.2f} ms")
Оба числа включают компиляцию, и обе функции вычисляют одну и ту же рекуррентную формулу. Развернутый цикл Python должен компилировать гораздо большую программу, поэтому его первый вызов происходит медленнее, чем первый.
Важное различие заключается в структуре на этапе компиляции: цикл Python разворачивается во время трассировки, тогда как lax.scan преобразует его в примитив цикла. Для коротких циклов часто подходит цикл for Python. Для длинных дифференцируемых циклов lax.scan обычно является лучшим вариантом по умолчанию.
8. Стабилизация форм с помощью отступов и статических аргументов.
Реальные рабочие нагрузки различаются по форме. Например, последний пакет в эпохе меньше по размеру, последовательности имеют разную длину. Каждая из этих ситуаций запускает новую компиляцию, если форма данных передается в JAX. Стандартное решение — дополнить входные данные до фиксированной формы и замаскировать неиспользуемые позиции .
MAX_LEN = 16
@jax.jit
def masked_mean(x, mask):
"""Mean of `x` ignoring positions where `mask==0`."""
# Always called with shape (MAX_LEN,) - no recompile when actual length varies
return jnp.sum(x * mask) / jnp.maximum(jnp.sum(mask), 1.0)
def pad(seq):
"""Right-pad a variable-length list of floats to `MAX_LEN` and return the padded array plus a 0/1 mask."""
actual_len = len(seq)
if actual_len > MAX_LEN:
raise ValueError(f"sequence length {actual_len} exceeds MAX_LEN={MAX_LEN}")
pad_len = MAX_LEN - actual_len
x = jnp.concatenate([
jnp.asarray(seq, dtype=jnp.float32),
jnp.zeros(pad_len, dtype=jnp.float32),
])
mask = jnp.concatenate([
jnp.ones(actual_len, dtype=jnp.float32),
jnp.zeros(pad_len, dtype=jnp.float32),
])
return x, mask
# Several different sequence lengths, but a single compiled function handles them all
for seq in [[1.0, 2.0, 3.0], [10.0] * 8, [5.0, -2.0]]:
x, mask = pad(seq)
print(f"len={len(seq):2d} mean={masked_mean(x, mask):.3f}")
Для каждой последовательности должна отображаться одна строка, в каждой из которых указывается только среднее значение действительных чисел — отступы не приближают среднее значение к нулю.
Все три вызова обращаются к одному и тому же скомпилированному исполняемому файлу, поскольку форма на устройстве всегда равна (MAX_LEN,) . Меняется только маска. Эта закономерность проявляется на всех уровнях, от простого среднего значения из 16 элементов до масок внимания с отступами при крупномасштабном обучении трансформеров.
Если вам нужна перекомпиляция: static_argnums
Добавление параметров (padding) позволяет избежать перекомпиляции, которую вы не запрашивали, а static_argnums позволяет запросить её намеренно.
Иногда параметр действительно является константой на стороне Python, например, количество слоев, флаг точности или размер ядра, и вы хотите, чтобы JAX встроил его значение в скомпилированную программу. Пометьте такие аргументы с помощью static_argnums или static_argnames для именованных аргументов. JAX хеширует значения этих аргументов в ключ кэша, поэтому для каждого уникального значения создается свой собственный скомпилированный исполняемый файл.
@partial(jax.jit, static_argnums=0)
def power(n: int, x):
"""Repeated squaring with `n` is static so JAX unrolls the loop and compiles a fresh program per value of `n`."""
# `n` is a Python int and JAX bakes it into the trace and unrolls the loop
y = x
for _ in range(n):
y = y * y
return y
jax.clear_caches()
for n in (2, 3, 2): # n=2 reuses the cache the second time
t0 = time.perf_counter()
_ = power(n, jnp.arange(4, dtype=jnp.float32)).block_until_ready()
print(f"n={n}: {(time.perf_counter() - t0) * 1000:7.2f} ms")
Вы должны увидеть, как первые два вызова компилируются, а третий, повторяющий n=2 , возвращается быстро.
Каждое новое значение n запускает компиляцию, но для фиксированных конфигураций это именно то, что нужно: цикл полностью разворачивается, и XLA может видеть каждую операцию. Компромисс очевиден: не следует помещать постоянно изменяющееся значение в static_argnums , иначе вам придётся перекомпилировать при каждом вызове.
9. Просмотрите трассировку с помощью jax.make_jaxpr
Если что-то компилируется не так, как вы ожидаете, jax.make_jaxpr позволяет увидеть трассировку до того, как XLA это изменит. jaxpr — это промежуточное представление компилятора на уровне JAX: типизированное функциональное представление того, что JAX вывел перед преобразованием в StableHLO, а затем в XLA.
Это не окончательная оптимизированная версия кода для графического процессора, но она очень полезна для понимания того, что отслеживал JAX.
def f(x):
return jnp.tanh(x) * jnp.sin(x) + jnp.log1p(x * x)
print(jax.make_jaxpr(f)(jnp.arange(4, dtype=jnp.float32)))
Вы должны увидеть небольшую типизированную программу: по одному примитиву на строку — tanh , sin , умножения, log1p и, наконец, сложение — каждый из которых аннотирован типом массива, который он создает.
Если вы подозреваете, что JAX перекомпилируется из-за неожиданного изменения формы или типа данных, сравнение двух JAXPR, полученных при «быстром» и «медленном» вызове, обычно позволяет точно определить виновника. Тот же приём работает и для выяснения того, почему такие преобразования, как grad или vmap создают больше работы, чем вы предполагали.
10. Уборка
Удалите рабочую нагрузку Jupyter, включая балансировщик нагрузки и постоянный том:
kubectl delete -f deploy/jupyter.yaml
Уничтожьте кластер, пул узлов, VPC и учетную запись службы:
cd terraform
terraform destroy
При появлении запроса введите yes , затем подтвердите, что ничего не осталось:
gcloud container clusters list
gcloud compute instances list
Оба поля должны быть пустыми для этого проекта. Если вы создали проект только для этой серии, вы можете удалить весь проект из консоли Cloud .
11. Поздравляем!
Вы создали ментальную модель, лежащую в основе jax.jit : JAX отслеживает вашу функцию Python, преобразует отслеженные вычисления во входные данные компилятора и кэширует скомпилированный исполняемый файл для сопоставления с входными сигнатурами.
Что вы узнали
- Как отслеживать выполнение функции Python с помощью
jax.jitи почему побочные эффекты Python, такие какprintвыполняются во время трассировки, а не при каждом выполнении. - Как отличить время компиляции от времени выполнения в кэше с помощью простого измерения времени выполнения функции
block_until_ready() - На что основаны ключи кэша компиляции: входная структура PyTree, формы, типы данных и значения статических аргументов.
- Как избежать ошибок управления потоком выполнения, связанных с отслеживанием значений, заменив ветвление в Python на
jnp.whereилиjax.lax.cond, и почемуjnp.whereможет приводить к появлению значений NaN в градиенте. - Как
lax.scanпомогает компактно обрабатывать длинный цикл фиксированной длины, вместо того чтобы разворачивать его в скомпилированную программу. - Как стабилизировать форму входных данных с помощью отступов и маскирования, и как встроить константы со стороны Python в трассировку с помощью
static_argnumsесли вы намеренно хотите создать отдельный исполняемый файл. - Как проверить трассировку на уровне JAX с помощью
jax.make_jaxprесли функция компилируется не так, как вы ожидаете.
Следующие шаги
- В мастер-классе Codelab 3: Профилирование и отладка JAX на GPU с использованием XProf и Nsight Systems вы научитесь наблюдать за трассировкой, компиляцией и выполнением ядра в реальном профиле.
- Увеличьте значение
NUM_STEPSс 200 до 2000 и повторно выполните сравнение циклов, чтобы увеличить разрыв междуlax.scanи развернутым циклом Python. - Поместите постоянно изменяющееся значение в
static_argnums: вызовитеpowerсn, взятым из счетчика, который увеличивается при каждом вызове, и наблюдайте, как каждый вызов перекомпилируется.