Como Transformar o Familiar NumPy em Código GPU Compilável
Se você já tentou acelerar cálculos pesados em Python, provavelmente esbarrou nas limitações da stack padrão. O NumPy é rápido graças ao seu backend em C, mas não funciona com GPUs nativamente. O PyTorch e o TensorFlow resolvem esse problema, mas trazem junto abstrações volumosas, classes de camadas e sua própria semântica de grafo computacional.
Em 2018, engenheiros do Google open-sourceram a biblioteca JAX. A ideia por trás é simples: dar aos desenvolvedores a sintaxe familiar do NumPy, mas adicionar diferenciação automática e o compilador XLA por cima. O resultado é Python funcional limpo que é compilado on-the-fly em código de máquina otimizado para GPUs ou TPUs.
O Que o JAX Realmente É
Muitas pessoas pensam no JAX como apenas mais um framework de ML, mas isso não é bem preciso. Os próprios autores afirmam no repositório que é um sistema de transformações funcionais composáveis para arrays.
Em vez de construir modelos de objetos complexos, o JAX incentiva o trabalho com funções puras. Você escreve código Python comum e então aplica funções transformadoras nele. Sem estados globais ocultos ou mutações de dados in-place. Se você precisa calcular um gradiente, compilar um pedaço do programa ou distribuir cálculos entre batches, basta envolver a função no decorator apropriado.
Vamos ver as quatro principais transformações nas quais toda a biblioteca é construída.
Diferenciação Automática via grad
A função jax.grad pega sua função e retorna uma nova que calcula o gradiente da original:
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
Você pode calcular gradientes de qualquer ordem simplesmente aninhando chamadas de jax.grad. O algoritmo lida com condicionais padrão do Python if/else, loops e recursão sem problemas.
Compilação via jit
Python regular executa cada operação de array sequencialmente, com overhead de chamadas de função e alocação de memória intermediária. O decorator jax.jit envia o corpo da função para o compilador XLA:
def calculate(x):
return x * x + x * 2.0
x = jnp.ones((5000, 5000))
# Обычный запуск против скомпилированного
fast_calculate = jax.jit(calculate)
O compilador funde operações elementares em um único kernel computacional. Como resultado, os dados não vão e voltam entre cache e memória GPU, e a performance aumenta em ordens de magnitude.
Vetorização de Código via vmap
Todo mundo que já escreveu algoritmos de machine learning passou horas ajustando dimensões de tensores para corresponder ao tamanho do batch. jax.vmap resolve essa dor de cabeça: você escreve a lógica para um único elemento ou vetor, e o próprio JAX vetoriza a operação:
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)
Em vez de um loop Python lento, o vetorizador empurra o loop para dentro de operações de baixo nível, transformando multiplicações matriz-vetor em multiplicações matriz-matriz completas.
Paralelismo e Sharding de Dados
Quando um modelo não cabe mais na memória de um único acelerador, o JAX oferece uma abordagem declarativa para paralelismo. Você define um device mesh e regras de particionamento de arrays (partition spec), e o próprio compilador distribui os cálculos e configura a troca de dados entre os cards:
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))
ressalvas e Especificidades
O JAX tem um lado negativo que você precisa se acostumar.
Primeiro, o paradigma funcional não permite efeitos colaterais. Você não pode simplesmente mudar um elemento do array por índice (arr[0] = 5), porque arrays no JAX são imutáveis. Para isso, use o método arr.at[0].set(5).
Segundo, a geração de números aleatórios requer passagem explícita de chaves de estado (jax.random.key), porque uma seed global quebraria a reprodutibilidade durante a compilação paralela.
Terceiro, debugar código JIT pode ser unfamiliar: na primeira chamada, a função passa por uma fase de tracing, e statements print Python regulares dentro dela só serão executados uma vez.
Instalação e Plataformas
A biblioteca suporta oficialmente Linux e macOS, e também roda no Windows via subsistema WSL2.
Para rodar em uma CPU comum:
pip install -U jax
Para construir com suporte a NVIDIA CUDA:
pip install -U "jax[cuda13]"
Também há suporte para aceleradores Google TPU e AMD ROCm.
Vale a Pena Experimentar
O JAX é ótimo para tarefas de pesquisa, modelagem física, otimizações não convencionais e computação científica onde o PyTorch parece muito volumoso e o NumPy comum não é rápido o suficiente. Se seu projeto está batendo nos limites de performance em operações matemáticas ou requer o cálculo de derivadas complexas, definitivamente vale a pena dar uma olhada em jax-ml/jax.
Projetos relacionados