Come Trasformare il Familiar NumPy in Codice GPU Compilabile
Se hai mai provato ad accelerare calcoli pesanti in Python, probabilmente hai incontrato i limiti dello stack standard. NumPy è veloce grazie al suo backend in C, ma non può lavorare con le GPU out of the box. PyTorch e TensorFlow risolvono questo problema, ma portano con sé astrazioni ingombranti, classi di layer e la loro propria semantica del grafo computazionale.
Nel 2018, gli ingegneri di Google hanno reso open-source la libreria JAX. L'idea alla base è semplice: dare agli sviluppatori la familiare sintassi NumPy, ma aggiungendo la differenziazione automatica e il compilatore XLA. Il risultato è codice Python funzionale pulito che viene compilato on-the-fly in codice macchina ottimizzato per GPU o TPU.
Cos'è JAX in Realtà
Molti pensano a JAX come a un semplice framework per ML, ma non è del tutto accurato. Gli stessi autori affermano nel repository che si tratta di un sistema di trasformazioni funzionali componibili per array.
invece di costruire modelli a oggetti complessi, JAX incoraggia a lavorare con funzioni pure. Scrivi codice Python normale, e poi applichi funzioni trasformatrici. Nessuno stato globale nascosto o mutazione dei dati in-place. Se hai bisogno di calcolare un gradiente, compilare una parte del programma, o distribuire i calcoli tra batch, ti basta racchiudere la funzione nel decorator appropriato.
Vediamo le quattro trasformazioni principali su cui è costruita l'intera libreria.
Differenziazione Automatica tramite grad
La funzione jax.grad prende la tua funzione e ne restituisce una nuova che calcola il gradiente dell'originale:
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
Puoi calcolare gradienti di qualsiasi ordine semplicemente annidando le chiamate a jax.grad. L'algoritmo gestisce senza problemi i costrutti condizionali standard di Python if/else, i cicli e la ricorsione.
Compilazione tramite jit
Il Python normale esegue ogni operazione su array sequenzialmente, con overhead dalle chiamate di funzione e dall'allocazione di memoria intermedia. Il decorator jax.jit invia il corpo della funzione al compilatore XLA:
def calculate(x):
return x * x + x * 2.0
x = jnp.ones((5000, 5000))
# Обычный запуск против скомпилированного
fast_calculate = jax.jit(calculate)
Il compilatore fonde le operazioni elementari in un singolo kernel computazionale. Di conseguenza, i dati non rimbalzano avanti e indietro tra cache e memoria GPU, e le prestazioni aumentano di diversi ordini di grandezza.
Vettorizzazione del Codice tramite vmap
Chiunque abbia scritto algoritmi di machine learning ha passato ore a regolare le dimensioni dei tensori per adattarle alla dimensione del batch. jax.vmap risolve questo mal di testa: scrivi la logica per un singolo elemento o vettore, e JAX stesso vettorizza l'operazione:
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)
Invece di un lento ciclo Python, il vettorizzatore spinge il ciclo all'interno delle operazioni di basso livello, trasformando le moltiplicazioni matrice-vettore in moltiplicazioni matrice-matrice complete.
Parallelismo e Partizionamento dei Dati
Quando un modello non rientra più nella memoria di un singolo acceleratore, JAX offre un approccio dichiarativo al parallelismo. Definisci una mesh di dispositivi e le regole di partizionamento dell'array (partition spec), e il compilatore stesso distribuisce i calcoli e configura lo scambio di dati tra le schede:
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))
Caveat e Specificità
JAX ha un rovescio della medaglia a cui bisogna abituarsi.
Prima di tutto, il paradigma funzionale richiede l'assenza di effetti collaterali. Non puoi semplicemente cambiare un elemento dell'array per indice (arr[0] = 5), perché gli array in JAX sono immutabili. Per questo, usa il metodo arr.at[0].set(5).
In secondo luogo, la generazione di numeri casuali richiede il passaggio esplicito delle chiavi di stato (jax.random.key), perché un seed globale romperebbe la riproducibilità durante la compilazione parallela.
In terzo luogo, il debugging del codice JIT può essere poco familiare: alla prima chiamata, la funzione attraversa una fase di tracing, e le istruzioni print Python regolari al suo interno verranno eseguite solo una volta.
Installazione e Piattaforme
La libreria supporta ufficialmente Linux e macOS, e funziona anche su Windows tramite il sottosistema WSL2.
Per l'esecuzione su una CPU normale:
pip install -U jax
Per la compilazione con supporto NVIDIA CUDA:
pip install -U "jax[cuda13]"
C'è anche supporto per gli acceleratori Google TPU e AMD ROCm.
Vale la Pena Provarlo
JAX è ottimo per task di ricerca, modellazione fisica, ottimizzazioni non standard e calcolo scientifico dove PyTorch risulta troppo ingombrante e NumPy normale non è abbastanza veloce. Se il tuo progetto sta raggiungendo i limiti di prestazioni sulle operazioni matematiche o richiede il calcolo di derivate complesse, dare un'occhiata a jax-ml/jax vale sicuramente la pena.
Progetti correlati