Как превратить привычный NumPy в компилируемый код для GPU

22 авг 2026
36,193
3,741
328
3 недели
logo

Если вы хоть раз пытались ускорить тяжелые вычисления на Python, то наверняка упирались в ограничения стандартного стека. NumPy работает быстро благодаря C-бэкенду, но не умеет работать с видеокартами из коробки. PyTorch и TensorFlow закрывают эту проблему, но тащат за собой громоздкие абстракции, классы слоев и собственную семантику работы с графом вычислений.

В 2018 году инженеры из Google выложили в открытый доступ библиотеку JAX. Идея под капотом простая: дать разработчику привычный синтаксис NumPy, но добавить к нему автоматическое дифференцирование и компилятор XLA. На выходе получается чистый функциональный Python, который на лету превращается в оптимизированный машинный код для видеокарт или TPU.

Что такое JAX на самом деле

Многие воспринимают JAX как очередной ML-фреймворк, но это не совсем так. Сами авторы в репозитории прямо пишут, что это система композируемых функциональных преобразований для массивов.

Вместо создания сложных объектных моделей JAX предлагает работать с чистыми функциями. Вы пишете обычный код на jax.numpy, а затем применяете к нему функции-трансформеры. Никаких скрытых глобальных состояний или мутаций данных на месте. Если вам нужно посчитать градиент, скомпилировать кусок программы или раскидать вычисления по батчам, вы просто оборачиваете функцию в соответствующий декоратор.

Посмотрим на четыре главных преобразования, вокруг которых строится вся библиотека.

Реклама

Автоматическое дифференцирование через grad

Функция jax.grad берет вашу функцию и возвращает новую, которая вычисляет градиент исходной:

import jax
import jax.numpy as jnp

def tanh(x):
    y = jnp.exp(-2.0 * x)
    return (1.0 - y) / (1.0 + y)

grad_tanh = jax.grad(tanh)
print(grad_tanh(1.0))  # 0.4199743

Градиенты можно брать любого порядка, просто вызывая jax.grad вложенно. При этом алгоритм спокойно проходит через стандартные питоновские ветвления if/else, циклы и рекурсию.

Компиляция через jit

Обычный Python выполняет каждую операцию над массивами последовательно, с накладными расходами на вызовы функций и выделение промежуточной памяти. Декоратор jax.jit отправляет тело функции в компилятор XLA:

def calculate(x):
    return x * x + x * 2.0

x = jnp.ones((5000, 5000))

# Обычный запуск против скомпилированного
fast_calculate = jax.jit(calculate)

Компилятор объединяет (fuses) элементарные операции в одно вычислительное ядро. В итоге данные не гоняются туда-сюда между кэшем и памятью GPU, а производительность вырастает в разы.

Векторизация кода через vmap

Каждый, кто писал алгоритмы машинного обучения, тратил часы на подгонку размерностей тензоров под batch-размер. jax.vmap решает эту головную боль: вы пишете логику для одного элемента или вектора, а JAX сам векторизует операцию:

def l1_distance(x, y):
    return jnp.sum(jnp.abs(x - y))

# Превращаем функцию для векторов в функцию для матриц
def pairwise_distances(xs):
    return jax.vmap(jax.vmap(l1_distance, (0, None)), (None, 0))(xs, xs)

xs = jax.random.normal(jax.random.key(0), (100, 3))
matrix = pairwise_distances(xs)  # форма (100, 100)

Вместо медленного цикла на Python векторизатор пробрасывает цикл внутрь низкоуровневых операций, превращая матрично-векторные умножения в полноценные матричные.

Параллелизм и шардирование данных

Когда модель перестает влезать в память одного ускорителя, JAX предлагает декларативный подход к параллелизму. Вы задаете сетку устройств (mesh) и правила разбиения массивов (partition spec), а компилятор сам распределяет вычисления и настраивает обмен данными между карточками:

from jax.sharding import set_mesh, AxisType, PartitionSpec as P

# Создаем сетку из 8 ускорителей
mesh = jax.make_mesh((8,), ('data',), axis_types=(AxisType.Explicit,))
set_mesh(mesh)

# Шардируем входные данные
inputs, targets = jax.device_put((inputs, targets), P('data'))

# Обычная функция градиента теперь выполняется параллельно
grad_fn = jax.jit(jax.grad(loss_fn))
grads = grad_fn(params, (inputs, targets))

Подводные камни и специфика

У JAX есть обратная сторона, к которой приходится привыкать.

Во-первых, функциональная парадигма требует отсутствия побочных эффектов. Нельзя просто так взять и изменить элемент массива по индексу (arr[0] = 5), потому что массивы в JAX неизменяемы (immutable). Для этого используется метод arr.at[0].set(5).

Во-вторых, генерация случайных чисел требует явного проброса ключей состояния (jax.random.key), так как глобальный сид сломал бы воспроизводимость при параллельной компиляции.

В-третьих, отладка JIT-кода может быть непривычной: при первом вызове функция проходит этап трассировки (tracing), и обычные питоновские принты внутри нее сработают только один раз.

Установка и платформы

Библиотека официально поддерживает Linux и macOS, а также запускается на Windows через подсистему WSL2.

Для работы на обычном процессоре:

pip install -U jax

Для сборки с поддержкой NVIDIA CUDA:

pip install -U "jax[cuda13]"

Также есть поддержка ускорителей Google TPU и AMD ROCm.

Стоит ли пробовать

JAX отлично подходит для исследовательских задач, физического моделирования, нестандартных оптимизаций и научных вычислений, где PyTorch кажется слишком громоздким, а чистого NumPy не хватает по скорости. Если ваш проект упирается в производительность математических операций или требует вычисления сложных производных, взглянуть на jax-ml/jax точно стоит.

🍪 Мы используем файлы cookie и сервис аналитики Яндекс.Метрика, чтобы сайт работал лучше. Продолжая пользоваться devtrends.ru, вы соглашаетесь с обработкой данных согласно Политике конфиденциальности.