JAX : accélère ton calcul numérique et calcule tes gradients en une ligne

De quoi avez-vous besoin

Version de Python

3.x

Packages

  • {"nom":"jax","version":"0.10.2"}
  • {"nom":"jaxlib","version":"0.10.2"}

Difficulté

Intermédiaire

Tu connais NumPy ? Alors tu connais déjà la moitié de JAX. Ce framework, développé par Google, reprend la même syntaxe que NumPy, mais compile ton code à la volée, le fait tourner sur GPU ou TPU sans rien réécrire, et calcule les gradients automatiquement. C'est la brique cachée derrière une partie de l'IA moderne : DeepMind s'en sert pour AlphaFold et Gemini. On va voir comment t'en servir en cinq étapes, avec du code que tu peux exécuter chez toi.


Installer JAX

Rien de sorcier : un environnement virtuel et un pip install. La version CPU suffit pour démarrer, pas besoin de carte graphique.

bash
python3 -m venv jaxenv
source jaxenv/bin/activate
pip install jax

1. jax.numpy : du NumPy qui court sur GPU

Premier réflexe : jax.numpy s'utilise exactement comme NumPy. Tu importes, tu calcules, c'est tout. La différence est invisible à l'œil nu, mais sous le capot, chaque opération est tracée pour être compilée ou dérivée plus tard.

python
import jax.numpy as jnp

x = jnp.arange(5)
print(x)               # [0 1 2 3 4]
print(jnp.sum(x ** 2)) # 30

2. jax.jit : compile ton code à la volée

Le vrai moteur, c'est jax.jit. Il prend ta fonction Python et la compile en XLA, un code machine optimisé, à la première exécution. Les appels suivants sont nettement plus rapides, surtout sur les gros tableaux.

python
import time
import jax
import jax.numpy as jnp

def travail(x):
    return jnp.sum(jnp.sin(x) ** 2 + jnp.cos(x) ** 3 + jnp.tanh(x) ** 2)

x = jnp.arange(10_000_000, dtype=jnp.float32)

travail_compile = jax.jit(travail)
travail_compile(x).block_until_ready()   # la compilation a lieu ici

t0 = time.perf_counter()
resultat = travail_compile(x).block_until_ready()
print(f"Compilé : {time.perf_counter() - t0:.4f} s -> {float(resultat):.2f}")

Sur une machine lambda, ce calcul passe d'environ 0,3 seconde en NumPy pur à 0,06 seconde une fois compilé. Branche un GPU et l'écart explose : XLA répartit le travail sur des milliers de cœurs sans que tu touches à ton code.

3. jax.grad : les gradients gratuits

Voici la raison pour laquelle l'IA adore JAX : la dérivation automatique. Tu écris une fonction de perte, et jax.grad t'en donne le gradient sans écrire une seule formule mathématique. C'est le moteur de l'apprentissage par descente de gradient.

python
import jax
import jax.numpy as jnp

def perte(w, X, y):
    return jnp.mean((X @ w - y) ** 2)

X = jnp.array([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]])
y = jnp.array([3.0, 7.0, 11.0])
w = jnp.array([0.5, 0.5])

gradient = jax.grad(perte)(w, X, y)
print(gradient)   # [-26.33 -33.33]

4. jax.vmap : vectorise sans boucle

Quand tu veux appliquer une fonction à un lot de données, jax.vmap la vectorise pour toi. Fini les boucles for : tu écris la fonction pour un seul élément, vmap la déroule sur tout le lot en parallèle.

python
import jax
import jax.numpy as jnp

def normalise(v):
    return v / jnp.linalg.norm(v)

mat = jnp.array([[3.0, 4.0], [0.0, 5.0], [6.0, 8.0]])
print(jax.vmap(normalise)(mat))
# [[0.6 0.8]
#  [0.  1. ]
#  [0.6 0.8]]

5. Une régression linéaire en dix lignes

Assemble tout : un jeu de données synthétique, une perte, un gradient compilé, et une boucle de descente. En dix lignes, tu entraînes un vrai modèle qui retrouve les bons poids tout seul.

python
import jax
import jax.numpy as jnp

key = jax.random.PRNGKey(42)
key, sk = jax.random.split(key)
X = jax.random.normal(sk, (200, 3))
w_vrai = jnp.array([2.0, -1.0, 0.5])
y = X @ w_vrai + 0.1 * jax.random.normal(key, (200,))

def mse(w):
    return jnp.mean((X @ w - y) ** 2)

g = jax.jit(jax.grad(mse))
w = jnp.zeros(3)
for _ in range(300):
    w = w - 0.05 * g(w)

print(w)   # [ 2.004 -1.000  0.487]  (proche de w_vrai)

Pourquoi JAX change la donne

JAX n'est pas qu'un NumPy plus rapide. C'est un langage fonctionnel déguisé : les fonctions sont pures, les tableaux immuables, et c'est cette discipline qui permet la compilation XLA et la dérivation automatique. C'est pour ça que les labos l'adoptent. Et tu n'as pas besoin d'un GPU pour commencer : tout tourne déjà en CPU. Le jour où tu veux passer sur GPU, tu réinstalles jax avec le support CUDA, sans toucher à ton code.

Conclusion

Tu as maintenant les quatre réflexes qui font 90 % de JAX : jnp pour calculer, jit pour accélérer, grad pour dériver, vmap pour vectoriser. À partir de là, PyTorch et les réseaux de neurones n'ont plus rien de magique : sous le capot, c'est exactement ce moteur qui tourne. Copie le code, casse-le, observe les temps de compilation tomber. C'est en faisant tourner la machine qu'on comprend comment elle pense.