Cómo Convertir NumPy Familiar en Código GPU Compilable
Si alguna vez has intentado acelerar cálculos pesados en Python, probablemente te has topado con las limitaciones del stack estándar. NumPy es rápido gracias a su backend en C, pero no puede trabajar con GPUs de forma nativa. PyTorch y TensorFlow resuelven este problema, pero traen consigo abstracciones complejas, clases de capas y su propia semántica de grafos de computación.
En 2018, ingenieros de Google liberaron el código de la biblioteca JAX. La idea detrás es simple: dar a los desarrolladores la sintaxis familiar de NumPy, pero agregar diferenciación automática y el compilador XLA encima. El resultado es Python funcional limpio que se compila sobre la marcha en código máquina optimizado para GPUs o TPUs.
Qué es realmente JAX
Mucha gente piensa en JAX como solo otro framework de ML, pero eso no es del todo preciso. Los propios autores declaran en el repositorio que es un sistema de transformaciones funcionales componibles para arrays.
En lugar de construir modelos de objetos complejos, JAX fomenta trabajar con funciones puras. Escribes código Python normal, y luego aplicas funciones transformadoras sobre él. No hay estados globales ocultos ni mutaciones de datos en el lugar. Si necesitas calcular un gradiente, compilar una parte del programa, o distribuir cálculos entre batches, simplemente envuelves la función en el decorador apropiado.
Veamos las cuatro transformaciones principales sobre las que se construye toda la biblioteca.
Diferenciación Automática mediante grad
La función jax.grad toma tu función y devuelve una nueva que calcula el gradiente de la 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
Puedes calcular gradientes de cualquier orden simplemente anidando llamadas a jax.grad. El algoritmo maneja condicionales estándar de Python if/else, bucles y recursión sin problemas.
Compilación mediante jit
Python normal ejecuta cada operación de array secuencialmente, con sobrecarga de llamadas a funciones y asignación de memoria intermedia. El decorador jax.jit envía el cuerpo de la función al compilador XLA:
def calculate(x):
return x * x + x * 2.0
x = jnp.ones((5000, 5000))
# Обычный запуск против скомпилированного
fast_calculate = jax.jit(calculate)
El compilador fusiona operaciones elementales en un único kernel de computación. Como resultado, los datos no van de un lado a otro entre la caché y la memoria GPU, y el rendimiento aumenta órdenes de magnitud.
Vectorización de Código mediante vmap
Todo el que ha escrito algoritmos de aprendizaje automático ha pasado horas ajustando las dimensiones de los tensores para que coincidan con el tamaño del batch. jax.vmap resuelve este dolor de cabeza: escribes la lógica para un solo elemento o vector, y JAX mismo vectoriza la operación:
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)
En lugar de un lento bucle Python, el vectorizador empuja el bucle dentro de operaciones de bajo nivel, convirtiendo multiplicaciones matriz-vector en multiplicaciones de matrices completas.
Paralelismo y Particionado de Datos
Cuando un modelo ya no cabe en la memoria de un solo acelerador, JAX ofrece un enfoque declarativo del paralelismo. Defines una malla de dispositivos y reglas de partición de arrays (partition spec), y el compilador mismo distribuye los cálculos y configura el intercambio de datos entre tarjetas:
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))
Advertencias y Particularidades
JAX tiene una cara opuesta a la que tienes que acostumbrarte.
Primero, el paradigma funcional requiere que no haya efectos secundarios. No puedes simplemente cambiar un elemento de un array por índice (arr[0] = 5), porque los arrays en JAX son inmutables. Para esto, usa el método arr.at[0].set(5).
Segundo, la generación de números aleatorios requiere pasar explícitamente claves de estado (jax.random.key), porque una semilla global rompería la reproducibilidad durante la compilación paralela.
Tercero, depurar código JIT puede resultar unfamiliar: en la primera llamada, la función pasa por una fase de tracing, y las sentencias print normales de Python dentro de ella solo se ejecutarán una vez.
Instalación y Plataformas
La biblioteca soporta oficialmente Linux y macOS, y también funciona en Windows mediante el subsistema WSL2.
Para ejecutar en una CPU normal:
pip install -U jax
Para construir con soporte para NVIDIA CUDA:
pip install -U "jax[cuda13]"
También hay soporte para aceleradores Google TPU y AMD ROCm.
¿Vale la Pena Probarlo
JAX es excelente para tareas de investigación, modelado de física, optimizaciones no estándar y computación científica donde PyTorch se siente demasiado pesado y NumPy plano no es lo suficientemente rápido. Si tu proyecto está alcanzando límites de rendimiento en operaciones matemáticas o requiere calcular derivadas complejas, definitivamente vale la pena echarle un vistazo a jax-ml/jax.
Proyectos relacionados