Как превратить привычный NumPy в компилируемый код для GPU
Если вы хоть раз пытались ускорить тяжелые вычисления на 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 точно стоит.
