Фаза 03 · урок 12

Введение в JAX

Цель урока: Вы знаете, как строить нейронные сети в PyTorch. Вы определяете nn.Module , вызываете .backward() , делаете шаг оптимизатора. Это работает. Миллионы людей этим пользуются.

Текущий релиз AlexBred.com: первые 100 уроков русскоязычной программы.

Курс
AI Engineering
Фаза
Основы глубокого обучения
Чтение
18 мин.
Проверено
Содержание урока
  1. Цели обучения
  2. Проблема
  3. Концепция
  4. Философия JAX
  5. jax.numpy: знакомая поверхность
  6. jax.grad: функциональное автоматическое дифференцирование
  7. jit: компиляция в XLA
  8. vmap: автоматическая векторизация
  9. pmap: параллелизм данных между устройствами
  10. Pytrees: универсальная структура данных
  11. Функциональный и объектно-ориентированный подходы
  12. Экосистема JAX
  13. Когда использовать JAX, а когда PyTorch
  14. Случайные числа в JAX
  15. Соберите сами
  16. Шаг 1: настройка и данные
  17. Шаг 2: инициализация параметров
  18. Шаг 3: прямой проход
  19. Шаг 4: JIT-скомпилированный шаг обучения
  20. Шаг 5: цикл обучения
  21. Используйте это
  22. Flax: стандарт Google
  23. Equinox: альтернатива в стиле Python
  24. Optax: компонуемые оптимизаторы
  25. Внедрите это
  26. Упражнения
  27. Ключевые термины
  28. Дополнительное чтение

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(). Вместо этого:

PyTorchJAX
Класс 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

ФакторJAXPyTorch
Поддержка TPUПервоклассная (Google создала оба)Поддерживается сообществом (torch_xla)
Поддержка GPUХорошая (CUDA через XLA)Лучшая в классе (нативная CUDA)
ОтладкаСложная (трассировка + компиляция)Простая (нетерпеливое, построчное выполнение)
ЭкосистемаОриентирована на исследования (Flax, Equinox)Огромная (HuggingFace, torchvision и т. д.)
НаймНишевая (Google/DeepMind/Anthropic)Массовая (повсеместно)
Масштабное обучениеПревосходное (XLA, pmap, mesh)Хорошее (FSDP, DeepSpeed)
Скорость прототипированияМедленнее (функциональные накладные расходы)Быстрее (изменяйте и запускайте)
Производственный инференсTensorFlow Serving, Vertex AITorchServe, 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 — промпт для выбора подходящей конфигурации оптимизатора JAX
  • outputs/skill-jax-patterns.md — навык, описывающий функциональные паттерны JAX

Упражнения

  1. Добавьте dropout в MLP. В JAX для dropout нужен PRNG-ключ: передавайте ключ через прямой проход и разделяйте его для каждого слоя dropout. Сравните тестовую точность с dropout и без него.

  2. Используйте jax.vmap, чтобы вычислить градиенты для каждого примера в пакете из 32 изображений MNIST. Вычислите норму градиента для каждого примера. У каких примеров градиенты наибольшие и почему?

  3. Замените ручную функцию прямого прохода универсальной mlp_forward(params, x), работающей с любым числом слоёв. Используйте jax.tree.leaves, чтобы автоматически определить глубину.

  4. Сравните по времени шаг обучения с @jax.jit и без него. Измерьте 100 шагов каждого варианта. Насколько велико ускорение на вашем оборудовании? Каковы накладные расходы на компиляцию при первом вызове?

  5. Реализуйте отсечение градиента, скомпоновав 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 — Отладка нейронных сетей