Фаза 03 · урок 12

Введение в JAX

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

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

Курс
AI Engineering from Scratch
Фаза
Основы глубокого обучения
Чтение
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(). Вместо этого:

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 — промпт для выбора подходящей конфигурации оптимизатора 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 — Отладка нейронных сетей