Фаза 03 · урок 12
Введение в JAX
Цель урока: Вы знаете, как строить нейронные сети в PyTorch. Вы определяете nn.Module , вызываете .backward() , делаете шаг оптимизатора. Это работает. Миллионы людей этим пользуются.
Текущий релиз AlexBred.com: первые 100 уроков русскоязычной программы.
Содержание урока
- Цели обучения
- Проблема
- Концепция
- Философия JAX
- jax.numpy: знакомая поверхность
- jax.grad: функциональное автоматическое дифференцирование
- jit: компиляция в XLA
- vmap: автоматическая векторизация
- pmap: параллелизм данных между устройствами
- Pytrees: универсальная структура данных
- Функциональный и объектно-ориентированный подходы
- Экосистема JAX
- Когда использовать JAX, а когда PyTorch
- Случайные числа в JAX
- Соберите сами
- Шаг 1: настройка и данные
- Шаг 2: инициализация параметров
- Шаг 3: прямой проход
- Шаг 4: JIT-скомпилированный шаг обучения
- Шаг 5: цикл обучения
- Используйте это
- Flax: стандарт Google
- Equinox: альтернатива в стиле Python
- Optax: компонуемые оптимизаторы
- Внедрите это
- Упражнения
- Ключевые термины
- Дополнительное чтение
PyTorch изменяет тензоры. TensorFlow строит графы. JAX компилирует чистые функции. Последнее меняет то, как вы мыслите о глубоком обучении.
Тип: Сборка Языки: Python Предварительные требования: Фаза 03, уроки 01–10; базовый NumPy Время: ~90 минут
Цели обучения
- Писать код нейронных сетей в стиле чистых функций, используя функциональный API JAX (
jax.numpy,jax.grad,jax.jit,jax.vmap) - Объяснять ключевое отличие между нетерпеливым выполнением с изменением состояния в PyTorch и моделью функциональной компиляции JAX
- Применять JIT-компиляцию и векторизацию
vmap, чтобы ускорять циклы обучения по сравнению с наивным Python - Обучить простую сеть в JAX и сопоставить явное управление состоянием с объектно-ориентированным подходом PyTorch
Проблема
Вы знаете, как строить нейронные сети в PyTorch. Вы определяете nn.Module, вызываете .backward(), делаете шаг оптимизатора. Это работает. Миллионы людей этим пользуются.
Но в ДНК PyTorch заложено ограничение: он нетерпеливо трассирует операции по одной в Python. Каждое tensor + tensor — отдельный запуск ядра. Каждый шаг обучения заново интерпретирует один и тот же Python-код. Это вполне подходит, пока вам не нужно обучать модель с 540 миллиардами параметров на 2 048 TPU. Тогда накладные расходы вас уничтожают.
Google DeepMind обучает Gemini на JAX. Anthropic обучала Claude на JAX. Это не небольшие задачи — это крупнейшие запуски обучения нейронных сетей на Земле. Они выбрали JAX, потому что он рассматривает ваш цикл обучения как компилируемую программу, а не как последовательность вызовов Python.
JAX — это NumPy с тремя суперсилами: автоматическим дифференцированием, JIT-компиляцией в XLA и автоматической векторизацией. Вы пишете функцию, которая обрабатывает один пример. JAX даёт вам функцию, которая обрабатывает пакет, вычисляет градиенты, компилируется в машинный код и выполняется на нескольких устройствах. И всё это без изменения исходной функции.
Концепция
Философия JAX
JAX — функциональный фреймворк. Никаких классов, изменяемого состояния или метода .backward(). Вместо этого:
| PyTorch | JAX |
|---|---|
Класс nn.Module с состоянием |
Чистая функция: f(params, x) -> y |
loss.backward() |
jax.grad(loss_fn)(params, x, y) |
| Нетерпеливое выполнение | JIT-компиляция через XLA |
Ручной цикл for x in batch: |
Автовекторизация jax.vmap(f) |
DataParallel / FSDP |
Автопараллелизм jax.pmap(f) |
Изменяемый model.parameters() |
Неизменяемое pytree из массивов |
Это не вопрос стилевых предпочтений. Это ограничение компилятора. Для JIT-компиляции требуются чистые функции: одни и те же входы всегда дают одинаковые выходы, побочные эффекты отсутствуют. Именно это ограничение делает возможным ускорение в 100 раз.
jax.numpy: знакомая поверхность
JAX реализует API NumPy на ускорителях:
import jax.numpy as jnp
a = jnp.array([1.0, 2.0, 3.0])
b = jnp.array([4.0, 5.0, 6.0])
c = jnp.dot(a, b)
Те же имена функций. Те же правила broadcasting. Та же семантика срезов. Но массивы находятся на GPU/TPU, а каждая операция доступна для трассировки компилятором.
Одно критически важное отличие: массивы JAX неизменяемы. Нельзя написать a[0] = 5. Вместо этого: a = a.at[0].set(5). Первую неделю это кажется неудобным, а затем становится понятно: неизменяемость делает композицию преобразований наподобие grad, jit и vmap возможной.
jax.grad: функциональное автоматическое дифференцирование
PyTorch прикрепляет градиенты к тензорам (.grad). JAX прикрепляет градиенты к функциям.
import jax
def f(x):
return x ** 2
df = jax.grad(f)
df(3.0)
jax.grad принимает функцию и возвращает новую функцию, вычисляющую градиент. Никакого вызова .backward(). Никакого графа вычислений, хранимого в тензорах. Градиент — просто ещё одна функция, которую можно вызывать, компоновать или JIT-компилировать.
Эта композиция произвольна:
d2f = jax.grad(jax.grad(f))
d2f(3.0)
Вторые производные. Третьи производные. Якобианы. Гессианы. Всё — композицией grad. PyTorch тоже умеет это (torch.autograd.functional.hessian), но там это добавлено поверх основной модели. В JAX это фундамент.
Ограничение: grad работает только с чистыми функциями. Внутри не должно быть операторов print (они выполняются при трассировке, а не при исполнении). Нельзя изменять внешнее состояние. Нельзя генерировать случайные числа без явного управления ключами.
jit: компиляция в XLA
@jax.jit
def train_step(params, x, y):
loss = loss_fn(params, x, y)
return loss
fast_step = jax.jit(train_step)
При первом вызове JAX трассирует функцию — записывает происходящие операции, не исполняя их. Затем он передаёт трассу XLA (Accelerated Linear Algebra), компилятору Google для TPU и GPU. XLA объединяет операции, устраняет избыточные копирования памяти и генерирует оптимизированный машинный код.
Последующие вызовы полностью обходят Python. Скомпилированный код выполняется на ускорителе со скоростью C++.
Когда JIT помогает:
- Шаги обучения (одинаковое вычисление повторяется тысячи раз)
- Инференс (та же модель, другие входы)
- Любая функция, вызываемая больше одного раза с входами схожей формы
Когда JIT мешает:
- Функции с потоком управления Python, зависящим от значений (
if x > 0, гдеx— трассируемый массив) - Однократные вычисления (накладные расходы на компиляцию превышают время выполнения)
- Отладка (трассировка скрывает фактическое исполнение)
Ограничение потока управления реально. jax.lax.cond заменяет if/else. jax.lax.scan заменяет циклы for. Это не необязательно — это цена компиляции.
vmap: автоматическая векторизация
Вы пишете функцию, обрабатывающую один пример:
def predict(params, x):
return jnp.dot(params['w'], x) + params['b']
vmap поднимает её до обработки пакета:
batch_predict = jax.vmap(predict, in_axes=(None, 0))
in_axes=(None, 0) означает: не группировать params в пакет (они общие), группировать по оси 0 массива x. Никакого ручного цикла for. Никакого изменения формы. Никакого протягивания размерности пакета. JAX сам определяет пакетную размерность и векторизует всё вычисление.
Это не синтаксический сахар. vmap генерирует слитый векторизованный код, выполняющийся в 10–100 раз быстрее Python-цикла. И он компонуется с jit и grad:
per_example_grads = jax.vmap(jax.grad(loss_fn), in_axes=(None, 0, 0))
Градиенты для каждого примера. Одна строка. В PyTorch это почти невозможно без ухищрений.
pmap: параллелизм данных между устройствами
parallel_step = jax.pmap(train_step, axis_name='devices')
pmap реплицирует функцию на все доступные устройства (GPU/TPU) и разбивает пакет. Внутри функции jax.lax.pmean и jax.lax.psum синхронизируют градиенты между устройствами.
Google обучает Gemini на тысячах чипов TPU v5e, используя pmap (и его преемника shard_map). Модель программирования: напишите версию для одного устройства, оберните в pmap — и готово.
Pytrees: универсальная структура данных
JAX работает с «pytrees» — вложенными комбинациями списков, кортежей, словарей и массивов. Параметры вашей модели — это pytree:
params = {
'layer1': {'w': jnp.zeros((784, 256)), 'b': jnp.zeros(256)},
'layer2': {'w': jnp.zeros((256, 128)), 'b': jnp.zeros(128)},
'layer3': {'w': jnp.zeros((128, 10)), 'b': jnp.zeros(10)},
}
Каждое преобразование JAX — grad, jit, vmap — умеет обходить pytrees. jax.tree.map(f, tree) применяет f к каждому листу. Так оптимизаторы обновляют все параметры разом:
params = jax.tree.map(lambda p, g: p - lr * g, params, grads)
Никакого метода .parameters(). Никакой регистрации параметров. Структура дерева и есть модель.
Функциональный и объектно-ориентированный подходы
PyTorch хранит состояние внутри объектов:
class Model(nn.Module):
def __init__(self):
self.linear = nn.Linear(784, 10)
def forward(self, x):
return self.linear(x)
JAX использует чистые функции с явным состоянием:
def predict(params, x):
return jnp.dot(x, params['w']) + params['b']
Параметры передаются явно. Ничего не хранится. Ничего не изменяется. Это делает каждую функцию тестируемой, компонуемой и компилируемой. Это также означает, что параметрами нужно управлять самостоятельно или использовать библиотеку вроде Flax или Equinox.
Экосистема JAX
JAX предоставляет примитивы. Библиотеки обеспечивают удобство:
| Библиотека | Роль | Стиль |
|---|---|---|
| Flax (Google) | Слои нейронных сетей | nn.Module с явным состоянием |
| Equinox (Patrick Kidger) | Слои нейронных сетей | На основе pytree, в стиле Python |
| Optax (DeepMind) | Оптимизаторы и расписания LR | Компонуемые преобразования градиентов |
| Orbax (Google) | Контрольные точки | Сохранение/восстановление pytrees |
| CLU (Google) | Метрики и логирование | Утилиты цикла обучения |
Optax — стандартная библиотека оптимизаторов. Она отделяет преобразование градиента (Adam, SGD, clipping) от обновления параметров, поэтому их легко компоновать:
optimizer = optax.chain(
optax.clip_by_global_norm(1.0),
optax.adam(learning_rate=1e-3),
)
Когда использовать JAX, а когда PyTorch
| Фактор | JAX | PyTorch |
|---|---|---|
| Поддержка TPU | Первоклассная (Google создала оба) | Поддерживается сообществом (torch_xla) |
| Поддержка GPU | Хорошая (CUDA через XLA) | Лучшая в классе (нативная CUDA) |
| Отладка | Сложная (трассировка + компиляция) | Простая (нетерпеливое, построчное выполнение) |
| Экосистема | Ориентирована на исследования (Flax, Equinox) | Огромная (HuggingFace, torchvision и т. д.) |
| Найм | Нишевая (Google/DeepMind/Anthropic) | Массовая (повсеместно) |
| Масштабное обучение | Превосходное (XLA, pmap, mesh) | Хорошее (FSDP, DeepSpeed) |
| Скорость прототипирования | Медленнее (функциональные накладные расходы) | Быстрее (изменяйте и запускайте) |
| Производственный инференс | TensorFlow Serving, Vertex AI | TorchServe, Triton, ONNX |
| Кто использует | DeepMind (Gemini), Anthropic (Claude) | Meta (Llama), OpenAI (GPT), Stability AI |
Честный ответ: используйте PyTorch, если у вас нет конкретной причины использовать JAX. Такими причинами могут быть доступ к TPU, потребность в градиентах для каждого примера, обучение на множестве устройств в огромном масштабе или работа в Google/DeepMind/Anthropic.
Случайные числа в JAX
У JAX нет глобального состояния генератора случайных чисел. Для каждой случайной операции нужен явный PRNG-ключ:
key = jax.random.PRNGKey(42)
key1, key2 = jax.random.split(key)
w = jax.random.normal(key1, shape=(784, 256))
Сначала это раздражает. Но это гарантирует воспроизводимость на разных устройствах и при разных компиляциях — свойство, которое torch.manual_seed в многGPU-сценариях гарантировать не может.
batchnorm-effect
Соберите сами
Шаг 1: настройка и данные
Мы обучим трёхслойный MLP на MNIST с использованием JAX и Optax. 784 входа, два скрытых слоя на 256 и 128 нейронов, 10 выходных классов.
import jax
import jax.numpy as jnp
from jax import random
import optax
def get_mnist_data():
from sklearn.datasets import fetch_openml
mnist = fetch_openml('mnist_784', version=1, as_frame=False, parser='auto')
X = mnist.data.astype('float32') / 255.0
y = mnist.target.astype('int')
X_train, X_test = X[:60000], X[60000:]
y_train, y_test = y[:60000], y[60000:]
return X_train, y_train, X_test, y_test
Шаг 2: инициализация параметров
Никакого класса. Только функция, возвращающая pytree:
def init_params(key):
k1, k2, k3 = random.split(key, 3)
scale1 = jnp.sqrt(2.0 / 784)
scale2 = jnp.sqrt(2.0 / 256)
scale3 = jnp.sqrt(2.0 / 128)
params = {
'layer1': {
'w': scale1 * random.normal(k1, (784, 256)),
'b': jnp.zeros(256),
},
'layer2': {
'w': scale2 * random.normal(k2, (256, 128)),
'b': jnp.zeros(128),
},
'layer3': {
'w': scale3 * random.normal(k3, (128, 10)),
'b': jnp.zeros(10),
},
}
return params
Инициализация He выполнена вручную. Три PRNG-ключа разделены из одного seed. Каждый вес — неизменяемый массив во вложенном словаре.
Шаг 3: прямой проход
def forward(params, x):
x = jnp.dot(x, params['layer1']['w']) + params['layer1']['b']
x = jax.nn.relu(x)
x = jnp.dot(x, params['layer2']['w']) + params['layer2']['b']
x = jax.nn.relu(x)
x = jnp.dot(x, params['layer3']['w']) + params['layer3']['b']
return x
def loss_fn(params, x, y):
logits = forward(params, x)
one_hot = jax.nn.one_hot(y, 10)
return -jnp.mean(jnp.sum(jax.nn.log_softmax(logits) * one_hot, axis=-1))
Чистые функции. Параметры на входе, предсказание на выходе. Никакого self, никакого сохранённого состояния. loss_fn вычисляет кросс-энтропию с нуля: softmax, логарифм, отрицательное среднее.
Шаг 4: JIT-скомпилированный шаг обучения
@jax.jit
def train_step(params, opt_state, x, y):
loss, grads = jax.value_and_grad(loss_fn)(params, x, y)
updates, opt_state = optimizer.update(grads, opt_state, params)
params = optax.apply_updates(params, updates)
return params, opt_state, loss
@jax.jit
def accuracy(params, x, y):
logits = forward(params, x)
preds = jnp.argmax(logits, axis=-1)
return jnp.mean(preds == y)
jax.value_and_grad возвращает и значение функции потерь, и градиенты за один проход. Декоратор @jax.jit компилирует обе функции в XLA. После первого вызова каждый шаг обучения выполняется без обращения к Python.
Шаг 5: цикл обучения
optimizer = optax.adam(learning_rate=1e-3)
X_train, y_train, X_test, y_test = get_mnist_data()
X_train, X_test = jnp.array(X_train), jnp.array(X_test)
y_train, y_test = jnp.array(y_train), jnp.array(y_test)
key = random.PRNGKey(0)
params = init_params(key)
opt_state = optimizer.init(params)
batch_size = 128
n_epochs = 10
for epoch in range(n_epochs):
key, subkey = random.split(key)
perm = random.permutation(subkey, len(X_train))
X_shuffled = X_train[perm]
y_shuffled = y_train[perm]
epoch_loss = 0.0
n_batches = len(X_train) // batch_size
for i in range(n_batches):
start = i * batch_size
xb = X_shuffled[start:start + batch_size]
yb = y_shuffled[start:start + batch_size]
params, opt_state, loss = train_step(params, opt_state, xb, yb)
epoch_loss += loss
train_acc = accuracy(params, X_train[:5000], y_train[:5000])
test_acc = accuracy(params, X_test, y_test)
print(f"Epoch {epoch + 1:2d} | Loss: {epoch_loss / n_batches:.4f} | "
f"Train Acc: {train_acc:.4f} | Test Acc: {test_acc:.4f}")
10 эпох. Примерно 97 % точности на тестовой выборке. Первая эпоха медленная (JIT-компиляция). Эпохи 2–10 быстрые.
Обратите внимание, чего здесь нет: ни .zero_grad(), ни .backward(), ни .step(). Всё обновление — один составной вызов функции. Градиенты вычисляются, преобразуются Adam и применяются к параметрам — всё внутри train_step.
Используйте это
Flax: стандарт Google
Flax — наиболее распространённая библиотека нейронных сетей JAX. Она возвращает nn.Module, но с явным управлением состоянием:
import flax.linen as nn
class MLP(nn.Module):
@nn.compact
def __call__(self, x):
x = nn.Dense(256)(x)
x = nn.relu(x)
x = nn.Dense(128)(x)
x = nn.relu(x)
x = nn.Dense(10)(x)
return x
model = MLP()
params = model.init(jax.random.PRNGKey(0), jnp.ones((1, 784)))
logits = model.apply(params, x_batch)
Та же структура, что и в PyTorch, но params отделены от модели. model.init() создаёт параметры. model.apply(params, x) выполняет прямой проход. У объекта модели нет состояния.
Equinox: альтернатива в стиле Python
Equinox (автор Patrick Kidger) представляет модели как pytrees:
import equinox as eqx
model = eqx.nn.MLP(
in_size=784, out_size=10, width_size=256, depth=2,
activation=jax.nn.relu, key=jax.random.PRNGKey(0)
)
logits = model(x)
Сама модель — pytree. Вызов .apply() не нужен. Параметры — просто листья модели. Это ближе к способу мышления JAX.
Optax: компонуемые оптимизаторы
Optax отделяет преобразование градиента от обновления:
schedule = optax.warmup_cosine_decay_schedule(
init_value=0.0, peak_value=1e-3,
warmup_steps=1000, decay_steps=50000
)
optimizer = optax.chain(
optax.clip_by_global_norm(1.0),
optax.adamw(learning_rate=schedule, weight_decay=0.01),
)
Отсечение градиента, разогрев скорости обучения, затухание весов — всё скомпоновано как цепочка преобразований. Каждое преобразование получает градиенты, изменяет их и передаёт следующему. Никакого монолитного класса оптимизатора.
Внедрите это
Установка:
pip install jax jaxlib optax flax
Для поддержки GPU:
pip install jax[cuda12]
Для TPU (Google Cloud):
pip install jax[tpu] -f https://storage.googleapis.com/jax-releases/libtpu_releases.html
Подводные камни производительности:
- Первый вызов JIT медленный (компиляция). Выполните прогрев до бенчмаркинга.
- Избегайте циклов Python по массивам JAX внутри JIT. Используйте
jax.lax.scanилиjax.lax.fori_loop. jax.debug.print()работает внутри JIT. Обычныйprint()— нет.- Профилируйте с
jax.profilerили TensorBoard. Компиляция XLA может скрывать узкие места. - По умолчанию JAX предварительно выделяет 75 % памяти GPU. Чтобы отключить это, задайте
XLA_PYTHON_CLIENT_PREALLOCATE=false.
Контрольные точки:
import orbax.checkpoint as ocp
checkpointer = ocp.PyTreeCheckpointer()
checkpointer.save('/tmp/model', params)
restored = checkpointer.restore('/tmp/model')
Этот урок создаёт:
outputs/prompt-jax-optimizer.md— промпт для выбора подходящей конфигурации оптимизатора JAXoutputs/skill-jax-patterns.md— навык, описывающий функциональные паттерны JAX
Упражнения
-
Добавьте dropout в MLP. В JAX для dropout нужен PRNG-ключ: передавайте ключ через прямой проход и разделяйте его для каждого слоя dropout. Сравните тестовую точность с dropout и без него.
-
Используйте
jax.vmap, чтобы вычислить градиенты для каждого примера в пакете из 32 изображений MNIST. Вычислите норму градиента для каждого примера. У каких примеров градиенты наибольшие и почему? -
Замените ручную функцию прямого прохода универсальной
mlp_forward(params, x), работающей с любым числом слоёв. Используйтеjax.tree.leaves, чтобы автоматически определить глубину. -
Сравните по времени шаг обучения с
@jax.jitи без него. Измерьте 100 шагов каждого варианта. Насколько велико ускорение на вашем оборудовании? Каковы накладные расходы на компиляцию при первом вызове? -
Реализуйте отсечение градиента, скомпоновав
optax.chain(optax.clip_by_global_norm(1.0), optax.adam(1e-3)). Обучите модель с отсечением и без него. Постройте график нормы градиента в процессе обучения, чтобы увидеть эффект.
Ключевые термины
| Термин | Как обычно говорят | Что это действительно означает |
|---|---|---|
| XLA | «То, что делает JAX быстрым» | Accelerated Linear Algebra — компилятор, объединяющий операции и генерирующий оптимизированные ядра GPU/TPU из графа вычислений |
| JIT | «Компиляция точно в срок» | JAX трассирует функцию при первом вызове, компилирует в XLA, затем запускает скомпилированную версию при последующих вызовах |
| Чистая функция | «Без побочных эффектов» | Функция, чей выход зависит только от входов: без глобального состояния, изменений и случайности без явных ключей |
| vmap | «Автоматическое пакетирование» | Преобразует функцию, обрабатывающую один пример, в функцию для пакета без переписывания |
| pmap | «Автопараллелизм» | Реплицирует функцию на нескольких устройствах и разбивает входной пакет |
| Pytree | «Вложенный словарь массивов» | Любая вложенная структура списков, кортежей, словарей и массивов, которую JAX может обходить и преобразовывать |
| Трассировка | «Запись вычисления» | JAX выполняет функцию с абстрактными значениями, чтобы построить граф вычислений, не вычисляя реальных результатов |
| Функциональное автодифференцирование | «grad функции» | Вычисление производных путём преобразования функций, а не присоединения хранилища градиентов к тензорам |
| Optax | «Библиотека оптимизаторов JAX» | Компонуемая библиотека преобразований градиентов: Adam, SGD, отсечение, расписания — которые объединяются в цепочку |
| Flax | «nn.Module для JAX» | Библиотека нейронных сетей Google для JAX, добавляющая абстракции слоёв при сохранении явного состояния |
Дополнительное чтение
- Документация JAX: https://jax.readthedocs.io/ — официальная документация с отличными руководствами по grad, jit и vmap
- «JAX: composable transformations of Python+NumPy programs» (Bradbury et al., 2018) — оригинальная статья, объясняющая философию проектирования
- Документация Flax: https://flax.readthedocs.io/ — библиотека нейронных сетей Google для JAX
- Patrick Kidger, «Equinox: neural networks in JAX via callable PyTrees and filtered transformations» (2021) — альтернатива Flax в стиле Python
- DeepMind, «Optax: composable gradient transformation and optimisation» — стандартная библиотека оптимизаторов
- «You Don’t Know JAX» (Colin Raffel, 2020) — практическое руководство по ловушкам и паттернам JAX от одного из авторов T5
Источник: Introduction to JAX 03.11 — Введение в PyTorch · Фаза 3 — Основы глубокого обучения · Полный каталог · 03.13 — Отладка нейронных сетей