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.
python3 -m venv jaxenv
source jaxenv/bin/activate
pip install jax1. 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.
import jax.numpy as jnp
x = jnp.arange(5)
print(x) # [0 1 2 3 4]
print(jnp.sum(x ** 2)) # 302. 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.
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.
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.
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.
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.






