Comment transformer du code NumPy familier en code GPU compilable
Si vous avez déjà essayé d'accélérer des calculs lourds en Python, vous avez probablement atteint les limites de la pile standard. NumPy est rapide grâce à son backend en C, mais ne peut pas fonctionner avec les GPU out-of-the-box. PyTorch et TensorFlow résolvent ce problème, mais apportent avec eux des abstractions volumineuses, des classes de couches et leur propre sémantique de graphe de calcul.
En 2018, les ingénieurs de Google ont rendu open-source la bibliothèque JAX. L'idée sous le capot est simple : donner aux développeurs une syntaxe NumPy familière, mais ajouter la différentiation automatique et le compilateur XLA par-dessus. Le résultat est du Python fonctionnel propre qui se compile à la volée en code machine optimisé pour les GPU ou les TPU.
Ce que JAX est réellement
Beaucoup de gens considèrent JAX comme un simple framework ML supplémentaire, mais ce n'est pas tout à fait exact. Les auteurs eux-mêmes indiquent dans le dépôt qu'il s'agit d'un système de transformations fonctionnelles composables pour les tableaux.
Au lieu de construire des modèles d'objets complexes, JAX encourage le travail avec des fonctions pures. Vous écrivez du code Python classique, puis appliquez des fonctions de transformation dessus. Pas d'états globaux cachés ni de mutations de données sur place. Si vous avez besoin de calculer un gradient, de compiler un bout de programme ou de distribuer des calculs sur des batches, vous encapsulez simplement la fonction dans le décorateur approprié.
Examinons les quatre transformations principales autour desquelles toute la bibliothèque est construite.
Différenciation automatique via grad
La fonction jax.grad prend votre fonction et retourne une nouvelle qui calcule le gradient de l'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
Vous pouvez calculer des gradients de n'importe quel ordre en imbriquant simplement les appels à jax.grad. L'algorithme gère les conditionnels Python standard if/else, les boucles et la récursion sans problème.
Compilation via jit
Le Python classique exécute chaque opération sur tableaux séquentiellement, avec des frais généraux liés aux appels de fonctions et à l'allocation mémoire intermédiaire. Le décorateur jax.jit envoie le corps de la fonction au compilateur XLA :
def calculate(x):
return x * x + x * 2.0
x = jnp.ones((5000, 5000))
# Обычный запуск против скомпилированного
fast_calculate = jax.jit(calculate)
Le compilateur fusionne les opérations élémentaires en un seul noyau de calcul. Ainsi, les données ne font pas d'allers-retours entre le cache et la mémoire GPU, et les performances augmentent de plusieurs ordres de grandeur.
Vectorisation du code via vmap
Tout ceux qui ont écrit des algorithmes de machine learning ont passé des heures à ajuster les dimensions des tenseurs pour correspondre à la taille du batch. jax.vmap résout ce casse-tête : vous écrivez la logique pour un seul élément ou vecteur, et JAX lui-même vectorise l'opération :
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)
Au lieu d'une lente boucle Python, le vectoriseur pousse la boucle à l'intérieur des opérations de bas niveau, transformant les multiplications matrice-vecteur en multiplications matrices complètes.
Parallélisme et partitionnement des données
Quand un modèle ne tient plus dans la mémoire d'un seul accélérateur, JAX offre une approche déclarative du parallélisme. Vous définissez un maillage de dispositifs et les règles de partitionnement des tableaux (partition spec), et le compilateur lui-même distribue les calculs et configure l'échange de données entre les cartes :
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))
Avertissements et spécificités
JAX a un revers auquel il faut s'habituer.
Premièrement, le paradigme fonctionnel interdit les effets secondaires. Vous ne pouvez pas simplement changer un élément de tableau par index (arr[0] = 5), car les tableaux dans JAX sont immuables. Pour cela, utilisez la méthode arr.at[0].set(5).
Deuxièmement, la génération de nombres aléatoires nécessite le passage explicite de clés d'état (jax.random.key), car une graine globale romprait la reproductibilité lors de la compilation parallèle.
Troisièmement, le débogage du code JIT peut être déroutant : lors du premier appel, la fonction passe par une phase de traçage, et les instructions print Python classiques à l'intérieur ne s'exécuteront qu'une seule fois.
Installation et plateformes
La bibliothèque supporte officiellement Linux et macOS, et fonctionne également sur Windows via le sous-système WSL2.
Pour une exécution sur CPU classique :
pip install -U jax
Pour une compilation avec support NVIDIA CUDA :
pip install -U "jax[cuda13]"
Il y a également un support pour les accélérateurs Google TPU et AMD ROCm.
Cela vaut-il la peine d'essayer
JAX est excellent pour les tâches de recherche, la modélisation physique, les optimisations non standard et le calcul scientifique là où PyTorch semble trop volumineux et NumPy classique n'est pas assez rapide. Si votre projet atteint des limites de performance sur les opérations mathématiques ou nécessite de calculer des dérivées complexes, jetter un œil à jax-ml/jax vaut définitivement le coup.
Projets similaires