JAX
Gratuit
JAX est un framework de programmation différenciable lancé par Google. Il fournit l'API NumPy et des capacités de compilation différentielle automatique XLA et d'accélération matérielle, devenant ainsi une infrastructure importante pour la recherche de pointe en matière de ML.
JAX
Paramètres de base et statistiques de JAX
JAX a emprunté une voie unique parmi les frameworks d'apprentissage profond traditionnels : il ne s'appelle pas une « bibliothèque de réseaux neuronaux », mais un « framework de calcul numérique différenciable ». C'est sur cette conception sous-jacente qu'une grande partie des recherches principales de DeepMind (AlphaFold, améliorations de l'infrastructure partielle Gemini AlphaGo) sont basées sur JAX. Contrairement à PyTorch et TensorFlow, JAX ne fournit pas d'API de réseau neuronal de haut niveau. Au lieu de cela, il fournit un ensemble de transformations de fonctions composables qui permettent aux développeurs d'exprimer des calculs dans un style purement fonctionnel, puis de les compiler dans des noyaux GPU/TPU efficaces via le compilateur XLA.
| Projets | JAX | PyTorch | TensorFlow |
|---|---|---|---|
| Positionnement officiel | Cadre de programmation différenciable haute performance | Cadre de recherche sur l'apprentissage profond | Plateforme ML de bout en bout |
| Paradigme de programmation | Fonctionnel (fonction pure + convertisseur) | Impératif (impatient par défaut) | Hybride déclaratif + impératif |
| Différenciation automatique | grad (mode inverse)/jacfwd (mode avant) | autograd (mode inversé) | GradientTape (mode inversé) |
| Mécanisme de compilation | XLA (décorateur jit) | TorchDynamo/Inducteur | XLA (tf.fonction) |
| Stratégie parallèle | pmap/pjit/shard_map | DDP/FSDP | Stratégie miroir/FSDP |
| Assistance matérielle | GPU NVIDIA, GPU AMD, Google TPU | GPU NVIDIA, GPU AMD, Apple MPS | GPU NVIDIA, GPU AMD, TPU |
| Bibliothèque de réseaux neuronaux | Lin/Haiku (tiers) | Torche intégrée.nn | tf.keras intégré |
| Licence Open Source | Apache2.0 | BSD | Apache2.0 |
| Étoiles GitHub | 33 000+ | 87 000+ | 188 000+ |
| Première version | 2018-12 | 2016-09 | 2015-11 |
| Principaux utilisateurs | Recherche de pointe en ML (DeepMind, etc.) | Universités + Industrie | Déploiement de production au niveau de l'entreprise |
Différence fondamentale : la conception fonctionnelle de JAX est la différence fondamentale par rapport à PyTorch/TensorFlow : il n'a pas les concepts d'« objets de modèle » et de « cycles de formation », mais utilise une combinaison de fonctions pures et de fonctions de conversion (jit, grad, vmap, pmap) pour exprimer les calculs. Cette conception confère à JAX des avantages uniques dans les scénarios de formation parallèle à grande échelle et de calcul de recherche scientifique personnalisé, mais elle entraîne également une courbe d'apprentissage plus abrupte.
Reconnaissance des utilisateurs et du marché de JAX
Adoption institutionnelle de recherche : JAX a une pénétration extrêmement élevée parmi les principales institutions de recherche en ML. DeepMind utilise JAX comme cadre de recherche principal depuis 2020. Les réalisations marquantes telles que AlphaFold 2/3, le modèle Chinchilla de la série Gemini et Gopher sont toutes implémentées sur la base de JAX ou de sa bibliothèque de couche supérieure. L'infrastructure expérimentale à grande échelle de Google Brain (maintenant Google DeepMind) utilise également JAX comme moteur informatique sous-jacent.
Communauté Open Source : le référentiel principal JAX sur GitHub a reçu plus de 33 000 étoiles et le nombre de forks a dépassé 3 100. Il existe plus de 200 projets écologiques construits autour de JAX, couvrant les bibliothèques de réseaux neuronaux (Flax, Haiku), les optimiseurs (Optax), l'apprentissage par renforcement (RLax, Acme), les réseaux de neurones graphiques (Jraph), l'inférence bayésienne (NumPyro, TensorFlow Probability pour JAX) et d'autres directions.
Applications d'entreprise : outre Google, NVIDIA (optimisant en profondeur les performances JAX via CUDA et cuDNN), Hugging Face (Transformers prend en charge le backend JAX/Flax), Cohere, Anthropic et d'autres sociétés utilisent également JAX pour certains travaux de formation ou d'inférence. Hugging Face dispose déjà de milliers de modèles pré-entraînés prenant en charge JAX/Flax dans sa bibliothèque de modèles.
Analyse comparative de l'industrie : dans les principaux documents de conférence tels que NeurIPS, ICML et ICLR, la proportion d'utilisation de JAX passera de moins de 5 % en 2020 à environ 35 à 40 % en 2025, et est devenue une infrastructure importante pour la méthodologie de recherche. La proportion de JAX utilisé comme outil pédagogique dans les cours universitaires augmente également d’année en année.
L'avantage financier de JAX : une infrastructure informatique haute performance sans frais de licence
La structure des coûts de JAX doit être évaluée indépendamment à partir de deux dimensions : le framework lui-même et le matériel en cours d'exécution :
Côté C/développeur individuel :
- Frais de cadre : JAX est entièrement open source, protocole Apache 2.0, sans frais de licence et peut être utilisé sans condition à des fins commerciales.
- Coût matériel : les particuliers peuvent exécuter JAX gratuitement sur leur propre GPU (série NVIDIA GeForce, série AMD Radeon). Pour les expériences à petite échelle qui ne nécessitent pas de GPU, l’exécution d’un processeur pur est également gratuite. L'accès TPU est facturé à l'heure via Google Cloud TPU, mais Google propose un quota TPU gratuit limité (comme le projet TRC).
Couche d'appel développeur/API :
- JAX lui-même ne fournit pas de services API cloud ; les développeurs n’ont rien à payer pour le framework lui-même.
- Les coûts d'infrastructure de formation dépendent de la plateforme de cloud computing choisie. Prenons l'exemple de Google Cloud :
- Instance GPU (par exemple A100 80G) : environ 3,50 $ à 5,00 $/heure
- Pod TPU v5p (tranchage multi-puces) : environ 30 $ à 100 $+/heure, selon la configuration
- AWS et Azure prennent également en charge la formation GPU JAX et sont facturés en fonction de la tarification de leur instance GPU respective.
Déploiements d'entreprise/privés :
- Zéro coût de framework : pas de frais de licence d'entreprise, pas de limite d'utilisateurs, pas de limite d'appels API.
- Coûts cachés :
-Acquisition de talents : les ingénieurs ML familiers avec la programmation fonctionnelle JAX bénéficient d'une prime salariale supérieure à celle des développeurs PyTorch, ce qui rend le recrutement plus difficile.
- Coût de migration : la migration de PyTorch/TensorFlow vers JAX nécessite une réécriture du pipeline de formation et du processus de traitement des données, et il peut y avoir une période de transformation initiale de 2 à 6 mois.
- Coût d'exploitation et de maintenance : la formation JAX à grande échelle nécessite le déploiement de Google Cloud TPU ou d'un cluster GPU auto-construit, et la complexité d'exploitation et de maintenance est proportionnelle à l'échelle.
- Avantages cachés : les optimisations de compilation XLA et de gestion de la mémoire de JAX peuvent réduire la consommation de ressources informatiques de 15 à 30 % dans le cadre d'une formation à grande échelle (par rapport aux implémentations équivalentes de PyTorch), ce qui peut compenser les coûts de migration à long terme.
| Dimension du coût | JAX | PyTorch | TensorFlow |
|---|---|---|---|
| Frais de licence-cadre | 0 $ | 0 $ | 0 $ |
| Modèle de licence d'entreprise | Aucun (Apache 2.0) | Aucun (BSD) | Aucun (Apache 2.0) |
| Seuil minimum de fonctionnement | Le processeur est suffisant (gratuit) | Le processeur est suffisant (gratuit) | Le processeur est suffisant (gratuit) |
| Coûts typiques de formation GPU | Facturation par instance GPU cloud | Facturation par instance GPU cloud | Facturation par instance GPU cloud |
| Coût d'utilisation du TPU | Nécessite Google Cloud (30 $+/h) | Ne prend pas directement en charge TPU | Nécessite Google Cloud (même prix) |
| Difficulté à acquérir des talents | Élevé (moins de développeurs) | Faible (grande communauté) | Moyen |
| Coût de la migration | Élevé (changement de paradigme) | — | Moyen (Keras existe déjà) |
| Efficacité des ressources de formation à grande échelle | Excellent (compilation et optimisation XLA) | Bon (Dynamo continue de s'améliorer) | Bon (compilation et optimisation XLA) |
Principales fonctions de JAX
- Différenciation automatique (
grad) : dérivé de toute fonction Python, prend en charge le mode inverse (le plus couramment utilisé) et le mode avant (jacfwd). Il peut être imbriqué pour calculer des dérivées d'ordre supérieur (telles que les matrices de Hesse), ce qui constitue une capacité essentielle pour les problèmes de calcul scientifique et d'optimisation.value_and_gradpeut renvoyer des valeurs de fonction et des dégradés en même temps, réduisant ainsi les calculs répétés. - Compilation juste à temps (
jit) : Compilez les fonctions Python dans des noyaux GPU/TPU efficaces via XLA. Le premier appel déclenche la compilation (~ 5 à 60 secondes, selon la complexité de la fonction), et les appels suivants exécutent directement le code hautes performances compilé. Les fonctions compilées s'exécutent souvent à des vitesses proches de celles du CUDA manuscrit, atteignant des accélérations de 50 à 100 fois supérieures à celles du Python pur sur les opérations matricielles intensives. - Vectorisation automatique (
vmap) : mappez automatiquement la logique de traitement par lots aux fonctions, éliminant ainsi le besoin d'écrire manuellement des boucles par lots. Par exemple, l'application de « vmap » à une fonction d'inférence à échantillon unique obtient automatiquement des capacités d'inférence par lots. Sous le capot,vmapfusionnera la dimension du lot dans l'opération de vectorisation existante, et les performances seront bien meilleures que celles de la boucle for manuelle. - Parallélisme multi-appareils (
pmap/pjit/shard_map) :pmapcopie automatiquement les calculs sur plusieurs appareils et effectue le parallélisme des données ; « pjit » (Partitioned JIT) partitionne automatiquement le graphique de calcul en tableaux de périphériques via des spécifications de partitionnement ;shard_map(JAX 0.4.16+) fournit un modèle de programmation SPMD explicite adapté aux stratégies de partitionnement personnalisées. Les trois couvrent tous les scénarios, du simple parallélisme de données au parallélisme de modèles complexes. - Pallas Kernel Language : DSL de noyau GPU personnalisé introduit dans JAX 0.4.20+, permettant aux noyaux GPU de bas niveau d'être écrits en Python (similaire à CUDA mais avec une syntaxe plus simple) et compilés et exécutés via XLA. Convient aux opérateurs personnalisés ayant des exigences de performances extrêmes, telles que les implémentations personnalisées de Flash Attention.
- Génération de nombres aléatoires (
jax.random) : système de nombres aléatoires fonctionnel - chaque fonction aléatoire reçoit et renvoie explicitement une valeur de clé PRNG, évitant ainsi l'état global implicite. Cette conception garantit la reproductibilité et
Naturellement thread-safe dans le calcul parallèle.
- Algèbre linéaire et API compatible NumPy (
jax.numpy/jax.lax/jax.scipy) :jax.numpyfournit une interface presque identique à NumPy et peut être accéléré de manière transparente sur GPU/TPU.jax.laxfournit des primitives d'algèbre linéaire de bas niveau etjax.scipycouvre les fonctions de calcul scientifique courantes.
Evolution du modèle et de la version de JAX
JAX a été open source par Google en décembre 2018 et a subi une évolution complète d'un cadre expérimental à une infrastructure de production.
Version principale
| Version | Dates | Changements clés |
|---|---|---|
| 0.1.0 | ~2019-02 | Première version publique, fournissant des convertisseurs de base grad, jit, vmap, pmap |
| 0.2.0 | ~2020-06 | Stabilisez l'API NumPy et introduisez l'interface complète jax.numpy ; DeepMind commence à adopter pleinement |
| 0.3.0 | ~2022-03 | Ajout de la compilation de fragments pjit pour prendre en charge la formation multi-machines et multi-TPU ; améliorations significatives des performances |
| 0.4.0 | ~2023-01 | Jalon de stabilité de l'API ; introduction du SPMD explicite shard_map ; AMD GPU prend en charge la version expérimentale |
| 0.4.16 | ~2024-06 | shard_map stable ; Version bêta du langage du noyau Pallas |
| 0.4.20 | ~2024-10 | Pallas officiellement libéré ; Améliorations de l'infrastructure de débogage (jax.debug) |
| 0.4.30 | ~2025-06 | Amélioration de la prise en charge du GPU AMD ROCm ; optimisation du cache de compilation ; nouvel aperçu du backend MLIR |
| 0.4.35 | ~2025-12 | Prise en charge du niveau de production des GPU AMD ; optimisation de la communication multi-nœuds ; amélioration de la lisibilité des messages d'erreur |
| 0.5.0 | ~2026-05 | Les performances de compilation XLA continuent de s'améliorer ; Extension du noyau Pallas ; Nettoyage des API |
Interprétation des points forts de la version
Série 0.2.x (2020-2021) : Une période critique pour que JAX établisse le positionnement trinitaire « NumPy + différenciation automatique + XLA ». Au cours de cette période, DeepMind a achevé la migration de sa pile de recherche principale de TensorFlow vers JAX, vérifiant ainsi la faisabilité de JAX dans la recherche sur le ML à grande échelle.
Série 0.3.x (2022-2023) : l'introduction de pjit fait de JAX l'un des rares frameworks à prendre en charge la « compilation de partitions en un clic » : les développeurs n'ont qu'à décrire l'intention de distribution des tenseurs sur chaque appareil (PartitionSpec), et pjit génère automatiquement un plan d'exécution multi-appareils. Au cours de la même période, des bibliothèques de formation à grande échelle telles que EasyLM, T5X et PaLM ont été créées sur la base de JAX.
Série 0.4.x (2023-2025) : L'écosystème JAX accélère sa maturité. Le langage du noyau Pallas comble le vide des opérateurs GPU personnalisés ; shard_map modifie le modèle de programmation SPMD d'implicite à explicite, abaissant ainsi le seuil de partitionnement personnalisé pour la formation à grande échelle ; Le GPU AMD prend en charge le passage de l'expérimentation à la production.
0.5.0 (2026-05) : en tant que première version de la gamme 0.5, elle poursuit la stratégie de stabilité de 0.4.x, en se concentrant sur l'optimisation de la surcharge de compilation XLA et de l'expérience de développement du noyau Pallas. Il n’y a pas encore de date officielle précise.
Avantages techniques de JAX
Conception fonctionnelle : déterminisme + composabilité
La conception « fonction pure » de JAX constitue la différence fondamentale par rapport à PyTorch/TensorFlow. Chaque fonction JAX ne contient aucun état interne et toutes les entrées et sorties sont transmises explicitement via des paramètres. Cela signifie : le même ensemble de paramètres et d'entrées produit toujours le même résultat (déterminisme) et les fonctions peuvent être librement combinées sans effets secondaires (composabilité). Cette conception est particulièrement importante dans le calcul parallèle : sans avoir à se soucier des conditions de concurrence dans un état partagé, pmap/pjit peut distribuer en toute sécurité des fonctions à des périphériques arbitraires.
Mécanisme → Effet : L'architecture combinée de fonctions pures + convertisseurs permet à grad, jit, vmap et pmap d'être imbriqués et composés arbitrairement (comme jit(grad(vmap(fn)))). Chaque couche de transformation se concentre uniquement sur la sémantique informatique d'une dimension et n'interfère pas avec les autres dimensions. C'est le principal avantage de JAX en termes d'expressivité - torch.vmap et torch.compile de PyTorch sont des capacités de "rattrapage" ultérieures, et leur composabilité et leur stabilité ne sont pas aussi bonnes que la conception native de JAX.
Compilation XLA : compilez une fois et exécutez sur tous les appareils
XLA (Accelerated Linear Algebra) est le compilateur sous-jacent de JAX, qui compile des graphiques de calcul au niveau des fonctions Python en code exécutable optimisé pour le matériel cible. Par rapport au mode d'exécution rapide de PyTorch (chaque opération est planifiée indépendamment), la compilation XLA permet d'améliorer les performances grâce aux mécanismes suivants :
- Opération Fusion : Fusion de petites opérations continues (telles que
add → relu → matmul → softmax) en un seul noyau GPU, réduisant ainsi les allers-retours en mémoire et les frais de lancement du noyau. Dans la formation Transformer, la fusion réduit généralement le nombre d'appels au noyau de 30 à 50 %. - Optimisation de la mémoire vidéo : XLA analyse le cycle de vie des tenseurs pendant la phase de compilation et insère automatiquement des stratégies de réutilisation et de suppression de tampon. Par rapport à la gestion manuelle, elle peut réduire l'utilisation maximale de la mémoire vidéo de 10 à 20 %.
- Indépendant du périphérique : le même code JAX peut s'exécuter sur le CPU, le GPU NVIDIA, le GPU AMD, le Google TPU sans modification, et XLA s'adapte automatiquement au matériel cible au moment de la compilation.
Formation à grande échelle : extension transparente d'une seule carte à dix mille cartes
L'abstraction parallèle de JAX (pmap → pjit → shard_map) forme un chemin d'expansion progressif depuis une seule machine vers un pod TPU à grande échelle :
- pmap (parallélisme des données) : copiez le modèle sur N appareils, chaque appareil traite différents micro-lots et synchronise les dégradés via all-reduce. Convient aux scénarios multi-cartes sur une seule machine avec le coût de configuration le plus bas.
- pjit (parallélisme de modèle + parallélisme de données) : en décrivant la distribution des tenseurs par appareil via
PartitionSpec, le compilateur génère automatiquement des graphiques de calcul et des plans de communication entre appareils. Convient aux formations à moyenne et grande échelle où les paramètres du modèle dépassent la mémoire d'un seul appareil. - shard_map (SPMD explicite) : introduit dans la version 0.4.16+, permettant aux développeurs d'écrire directement des fonctions qui s'exécutent sur des données fragmentées, et le compilateur gère automatiquement la communication entre fragments. Convient aux stratégies de partitionnement personnalisées (telles que le parallélisme séquentiel, le parallélisme expert).
Effet : DeepMind a utilisé JAX + pjit pour entraîner un modèle GShard-MoE avec 500 milliards de paramètres sur 6 144 puces TPU v4, obtenant ainsi une efficacité de mise à l'échelle quasi linéaire. Cette capacité de parallélisme à grande échelle ne peut être obtenue que par la combinaison JAX + TPU dans les frameworks traditionnels actuels.
Limite d'adaptation (scénarios applicables et inapplicables)
Les scénarios dans lesquels JAX est le meilleur :
- Entraînements distribués à grande échelle (niveau 100 calories à 10 000 calories), notamment entraînements sur clusters TPU
- Calculs scientifiques (simulations physiques, dynamique moléculaire, modélisation climatique) nécessitant des dérivées d'ordre élevé ou des calculs de gradient personnalisés
- Code expérimental orienté recherche (nécessite des modifications fréquentes de la structure du modèle, de la fonction de perte personnalisée, de l'opérateur expérimental)
- Formation de grands modèles avec des stratégies parallèles de modèles complexes (MoE, parallélisme de séquence, tensor sharding, etc.)
Scénarios pour lesquels JAX n'est pas bon :
- Débuter avec le prototypage et l'enseignement rapides (courbe d'apprentissage beaucoup plus raide que PyTorch)
- Modèles à flux de contrôle dynamique intensifs (tels que Tree-RNN, réseaux de graphes récursifs), bien que
jax.lax.while_loop/condfournisse une prise en charge, l'expression et le débogage sont beaucoup moins pratiques que les graphes dynamiques PyTorch - Pipelines d'inférence de production qui nécessitent une interaction fréquente avec des systèmes externes non Python
- Projets ML occasionnels/non liés à la recherche (la richesse des bibliothèques et des outils de modèles communautaires est bien moindre que celle de PyTorch)
- Vous disposez déjà d'une base de code PyTorch mature et d'une expérience d'équipe, et le coût de la migration est supérieur aux avantages.
Performances et débit
Les performances de JAX obtenues grâce à la compilation XLA sont compétitives par rapport au code optimisé écrit à la main dans les dimensions suivantes :
- TTFT (Time to First Token) : la compilation
jitde JAX prend beaucoup de temps pour la première fois (généralement 5 à 60 secondes) car l'analyse complète du graphique de calcul et la génération du code matériel doivent être terminées. La surcharge des appels ultérieurs, y compris la détection de recompilation après modification des paramètres, est considérablement réduite. En comparaison, le mode impatient de PyTorch n'a aucun délai de compilation et le temps de préchauffage de TorchDynamo est d'environ 10 à 30 secondes. - Débit (débit de formation) : dans les tâches de formation standard de Transformer, le débit de la combinaison JAX + TPU est généralement 20 à 50 % plus élevé que celui de PyTorch avec la même configuration GPU. Dans le contexte du GPU, l'écart de performances entre JAX et PyTorch se réduit, et JAX reste en tête dans certains opérateurs spécifiques bien intégrés. La valeur spécifique dépend de l'architecture du modèle, de la taille du lot et du type de matériel, et il n'existe pas de référence officielle unifiée.
- Contrôle de fréquence TPM/RPM : JAX en tant que framework local n'a pas de contrôle de fréquence d'appel API ; lors de l'utilisation de Google Cloud TPU, il est soumis à des restrictions de quota de ressources cloud (quota horaire de puce TPU) et à des restrictions TPM/RPM non liées à l'API.
Comment utiliser JAX
Installation
JAX fournit des packages d'installation pip pour différents backends matériels :
# Version CPU (universelle, aucun GPU requis)
pip installer jax jaxlib
# Version GPU NVIDIA (CUDA 12)
pip installer jax[cuda12]
# Version du GPU AMD (ROCm)
pip installer jax[rocm]
# Version TPU (doit s'exécuter dans l'environnement Google Cloud TPU)
pip installer jax[tpu]
Après l'installation, vérifiez la situation : python -c "import jax; print(jax.devices())", qui devrait afficher une liste des périphériques matériels actuellement disponibles.
Exemples de code API de base
Exemple de différenciation automatique :
importer jax
importer jax.numpy en tant que jnp
déf f(x):
retourner jnp.sin(x) * jnp.exp(-x**2)
# Dérivée première
df = jax.grad(f)
print(df(1.0)) # df/dx à x=1.0
#Dérivée seconde (imbrication des diplômes)
d2f = jax.grad(jax.grad(f))
print(d2f(1.0)) # d²f/dx² à x=1.0
# Renvoie à la fois la valeur de la fonction et le dégradé
val_grad = jax.value_and_grad(f)
print(val_grad(1.0)) # (f(1.0), df(1.0))
Exemple de compilation juste à temps :
importer jax
importer jax.numpy en tant que jnp
# Compiler une fonction de multiplication matricielle
@jax.jit
def matmul_fast(A,B):
retourner jnp.dot(A, B)
# Le premier appel déclenche la compilation XLA (prend un peu plus de temps)
A = jnp.ones((4096, 4096))
B = jnp.ones((4096, 4096))
C = matmul_fast(A, B) # compiler + exécuter
# Les appels suivants exécutent directement le code compilé
C = matmul_fast(A, B) # Exécution uniquement, pas de surcharge de compilation
# Exemple de paramètre statique : spécifiez les paramètres qui n'ont pas besoin d'être suivis dans le graphique de calcul
@jax.jit(static_argnums=(2,))
def conv_with_padding(x, w, padding_mode) :
retourner jnp.convolve(x, w, mode=padding_mode)
Exemple d'autovectorisation :
importer jax
importer jax.numpy en tant que jnp
#Fonction d'inférence d'échantillon unique
def predict_single(params, x) :
retourner jnp.dot(params, x)
# Inférence automatique par lots
batch_predict = jax.vmap(predict_single, in_axes=(Aucun, 0))
# in_axes=(Aucun, 0) signifie que les paramètres ne sont pas divisés (partagés), x est divisé le long de la 0ème dimension
paramètres = jnp.ones((256, 64))
batch_x = jnp.ones((32, 64)) # 32 échantillons
résultats = batch_predict(params, batch_x) # forme : (32, 256)
Exemple de parallélisme multi-appareils :
importer jax
importer jax.numpy en tant que jnp
# Parallélisme des données : pmap copie les fonctions sur tous les appareils
def train_step (params, batch):
perte = calculate_loss (params, lot)
diplômés = jax.grad(compute_loss)(params, batch)
perte de retour, jax.pmean(grads, axis_name='devices')
#num_devices appareils traitent chacun une partie du lot
paramètres = jnp.ones((1024, 512))
batch = jnp.ones((64, 512)) # Sera automatiquement divisé en chaque appareil
perte, diplômes = jax.pmap(train_step, axis_name='devices')(params, batch)
Description du paramètre clé :
jax.jit(fun, static_argnums=(), donate_argnums=()):static_argnumsspécifie les index de paramètres à ne pas tracer dans le graphique de calcul (s'applique aux paramètres de forme/configuration) ;donate_argnumsdéclare que le tampon d'entrée peut être écrasé pour économiser la mémoire vidéo.jax.grad(fun, argnums=0, has_aux=False):argnumsspécifie quels paramètres sont différenciés ; lorsquehas_aux=True, la fonction renvoie(sortie primaire, données auxiliaires)et grad ne différencie que la sortie principale.jax.vmap(fun, in_axes=0, out_axes=0):in_axes/out_axesspécifie quelles dimensions des tenseurs d'entrée/sortie correspondent aux dimensions du lot.jax.pmap(fun, axis_name, devices=None):axis_nameest un identifiant nommé utilisé pour les opérations de communication collective telles quepmean/all_gather; « appareils » peut spécifier un sous-ensemble d'appareils participants.jax.lax.with_sharding_constraint(x, sharding): spécifiez explicitement la stratégie de partitionnement du tenseur dans pjit.
Outils de développement et débogage
- jax.debug : 0.4.20+ fournit des outils de point d'arrêt et d'impression pour afficher les valeurs intermédiaires compilées.
- jax.make_jaxpr : convertissez les fonctions en représentation interne JAX (Jaxpr) pour analyser les structures de graphiques informatiques.
- jax.profiler : un outil d'analyse des performances intégré à TensorBoard qui peut afficher la consommation de temps du noyau et l'allocation de mémoire vidéo.
- Orbax : la bibliothèque de points de contrôle JAX officielle de Google, prend en charge la sauvegarde asynchrone et les points de contrôle fragmentés SPMD.
Prix des produits pour JAX
JAX lui-même est entièrement open source et gratuit, et son coût total se compose de deux parties : le coût d'utilisation du framework et le coût de fonctionnement du matériel.
Coût d'utilisation du framework :
| Projet | Tarifs | Descriptif |
|---|---|---|
| Cadre JAX | 0 $ | Protocole open source Apache 2.0, utilisation commerciale illimitée |
| Lin / Haïku / Optax | 0 $ | La bibliothèque de niveau supérieur est également open source et gratuite |
| Licence d'entreprise | 0 $ | Aucun accord d'entreprise supplémentaire ni frais de licence requis |
| Assistance technique | Assistance technique communautaire gratuite / payante Google Cloud | Plan d'assistance officiel non rémunéré ; Les clients Google Cloud peuvent bénéficier d'une assistance relative au TPU |
Frais de fonctionnement du matériel :
| Type de matériel | Comment obtenir | Prix de référence |
|---|---|---|
| Processeur | Propre serveur ou toute instance de CPU cloud | Inclus dans les ressources informatiques existantes |
| GPU NVIDIA (personnel) | Propre GPU | Investissement matériel ponctuel (300 $ à 3 000 $) |
| GPU NVIDIA (nuage) | Instance GPU Google Cloud/AWS/Azure | 0,50 $ à 5,00 $/heure (variant de T4/A100/H100) |
| GPU AMD (nuage) | Instance Google Cloud A3 / auto-construite | Semblable au GPU cloud NVIDIA |
| Google Cloud TPU v5e | Google Cloud à la demande/préempté | ~1,50 $ à 4,00 $/heure (puce unique) |
| Google Cloud TPU v5p | Google Cloud à la demande/préempté | ~12,00 $ à 30,00 $+/heure (puce unique) |
| Pod TPU (tranchage multi-puces) | Pré-occupation de Google Cloud | Devis commercial requis, généralement 100 $+/heure |
Quota gratuit : Google propose le projet TPU Research Cloud (TRC), qui offre un quota d'accès gratuit limité au TPU aux chercheurs universitaires. Les nouveaux utilisateurs de Google Cloud peuvent bénéficier d'un crédit d'essai de 300 $ pour tester les instances TPU/GPU.
Suggestion payante :
- Recherche personnelle : utiliser votre propre quota de GPU ou de TPU gratuit TRC est le meilleur moyen, avec un coût pratiquement nul.
- Équipes petites et moyennes : utilisez des instances cloud NVIDIA GPU (A100 80G, ~ 4 $/heure), budget mensuel de 1 000 $ à 5 000 $.
- Équipes de formation à grande échelle : nécessité d'évaluer le rapport coût/performance des clusters TPU par rapport aux clusters GPU. TPU Pod est plus efficace dans les scénarios parallèles à grande échelle (plus de 256 puces), mais le coût de configuration initial est plus élevé et il est lié à Google Cloud. Il est recommandé d'effectuer une comparaison pilote à petite échelle pendant 2 à 4 semaines avant de prendre une décision.
Scénarios d'application JAX
- Recherche de pointe en ML et récurrence des articles : NeurIPS/ICML/ICLR Environ 35 % des articles en 2024-2025 impliquent des implémentations JAX, des variantes de Transformer aux modèles de diffusion en passant par les algorithmes d'apprentissage par renforcement. Conseils d'implémentation : lors de la reproduction d'articles JAX, donnez la priorité à la recherche d'implémentations open source basées sur Flax ou Haiku ; Le code JAX pur (qui ne repose pas sur des bibliothèques de haut niveau) est généralement difficile à migrer directement vers l'environnement de production.
- Infrastructure de formation de modèles à grande échelle : la bibliothèque de formation (T5X, EasyLM, pipeline PaLM) construite sur la base de JAX prend en charge la formation de la plupart des 100 B+ modèles de paramètres internes de Google. Conseils de mise en œuvre : avant de commencer la formation de dizaines de milliards de paramètres, l'équipe doit disposer d'au moins 1 à 2 ingénieurs familiers avec la sémantique de partitionnement pjit/shard_map, sinon le cycle de débogage peut durer jusqu'à 2 à 4 semaines.
- Calcul scientifique et simulation physique : Les caractéristiques différenciables de JAX lui confèrent des avantages uniques dans des domaines tels que la dynamique moléculaire (JAX-MD), la modélisation astrophysique (JAX-Cosmo) et la simulation climatique (JAX-Climate). Par rapport aux outils informatiques scientifiques traditionnels (tels que MATLAB et Fortran), JAX offre une différenciation automatique et une accélération GPU/TPU, abaissant ainsi le seuil de développement de modèles scientifiques. Conseils d'implémentation : dans les scénarios de calcul scientifique, le mode 64 bits de JAX (
jax.config.update("jax_enable_x64", True)) doit être utilisé en premier. Le mode 32 bits par défaut peut introduire des erreurs de précision cumulatives. - Plateforme de formation par apprentissage par renforcement : les bibliothèques RL open source de DeepMind (Acme, RLax, Mava) sont toutes construites sur JAX, en utilisant vmap et pmap pour obtenir le parallélisme de contexte et le parallélisme de formation. Conseils de mise en œuvre : la formation RL implique souvent un grand nombre d'interactions contextuelles. Le modèle fonctionnel pur de JAX s'adapte naturellement au cycle « état-action-récompense » de RL. Cependant, vous devez faire attention au gaspillage de calcul causé par les différentes conditions de terminaison de chaque contexte lorsque vmap est parallèle au contexte.
- Développement du noyau GPU/TPU et vérification des prototypes : le langage du noyau Pallas offre un niveau d'abstraction plus élevé que CUDA pour le développement du noyau GPU et convient à la vérification rapide des opérateurs personnalisés (tels que les variantes Flash Attention). Conseil de mise en œuvre : Pallas ne prend actuellement en charge que les GPU NVIDIA et TPU, la prise en charge des GPU AMD n'est pas encore disponible
Écurie; Le développement du noyau au niveau de la production doit encore revenir à CUDA pour un réglage précis.
Groupes applicables de JAX
- Chercheurs de pointe en ML (utilisateurs principaux) : il s'agit du principal groupe cible de JAX. Si vous effectuez des recherches en ML chez DeepMind, Google Brain, un laboratoire d'IA de premier plan ou une université de premier plan, JAX est votre « langue maternelle ». Une maîtrise approfondie de la programmation fonctionnelle JAX et des stratégies de partitionnement pjit/shard_map sont des compétences essentielles pour faire progresser les expériences à grande échelle. Prérequis : Vous devez comprendre le principe de différenciation automatique, les concepts de base de la formation distribuée et avoir de l'expérience dans l'utilisation d'au moins un framework d'apprentissage profond.
- Chercheurs en informatique scientifique et en équations différentielles : chercheurs qui ont besoin de simulation numérique et de résolution d'équations différentielles dans les domaines de la physique, de la chimie, de la biologie, du climat, etc. La combinaison grad/vmap/pmap de JAX peut raccourcir considérablement le cycle des formules mathématiques aux simulations exécutables. Prérequis : Familier avec l'écosystème NumPy/SciPy, pas besoin d'expérience en deep learning pour débuter avec la partie calcul numérique de JAX.
- Ingénieur de formation des grands modèles : L'équipe d'ingénierie responsable de la formation du modèle à l'échelle paramétrique 10B-1T. JAX + TPU est l'une des rares solutions de formation éprouvées au niveau Wanka. Prérequis : une compréhension approfondie du modèle de programmation SPMD, de la topologie de communication (tout réduire/tout rassembler/réduire-diffusion) et des connaissances en matière d'exploitation et de maintenance de Google Cloud TPU sont requises.
- Ingénieur en apprentissage automatique (nécessite une évaluation minutieuse) : Si votre travail quotidien consiste à utiliser des modèles pré-entraînés pour le réglage fin, le déploiement et l'intégration commerciale, JAX n'est pas le meilleur choix - l'écosystème communautaire de PyTorch, les outils de déploiement (TorchServe, ONNX, TensorRT) et l'exhaustivité dépassent de loin JAX. Conditions inappropriées : dans les scénarios où il n'y a pas de besoins de recherche à long terme, l'équipe utilise PyTorch comme pile principale et le cycle de livraison du projet est de 3 mois, il n'est pas recommandé d'introduire JAX.
- Étudiants et débutants (non recommandé en priorité) : la haute abstraction et la conception fonctionnelle de JAX ne sont pas adaptées aux débutants en ML. Il est recommandé d'établir d'abord les concepts de base du deep learning (tenseurs, différenciation automatique, boucles d'entraînement) via PyTorch, puis de l'utiliser lors du calcul haute performance ou de la reproduction de données spécifiques.
Apprenez JAX tout en faisant des recherches. Conditions inappropriées : Pour les apprenants qui débutent dans le deep learning depuis moins de 6 mois, la courbe d'apprentissage de JAX peut entraîner une charge cognitive excessive.
Résumé et Outlook
JAX a un statut déterminant dans le sens technique de la « programmation différenciable » : sa conception fonctionnelle et son haut niveau d'abstraction du matériel sous-jacent le rendent irremplaçable dans la recherche de pointe en ML avec le seuil le plus élevé.
Compétences de base :
- Paradigm Leadership : La conception de convertisseurs fonctionnels + est théoriquement plus adaptée pour exprimer et combiner des calculs complexes que les cadres impératifs. Cet avantage est particulièrement important dans les scénarios distribués et multi-appareils.
- Profondeur d'abstraction matérielle : la combinaison de JAX + XLA fournit un modèle de programmation unifié du CPU au pod TPU, qui peut être écrit une seule fois et exécuté sur différents backends matériels, ce qui est unique parmi les frameworks grand public actuels.
- Vérification d'entraînement à grande échelle : Après plusieurs années de vérification de production à l'échelle de milliers, voire de dizaines de milliers de puces au sein de DeepMind et de Google, la maturité technique de JAX en matière d'entraînement parallèle à grande échelle a été testée en combat réel.
Limites actuelles :
- Courbe d'apprentissage abrupte : des concepts tels que le paradigme fonctionnel, la composition du convertisseur et la sémantique de partitionnement nécessitent un changement de pensée spécialisé. Il faut généralement 1 à 3 mois aux développeurs pour migrer depuis PyTorch.
- Richesse écologique insuffisante : La richesse des bibliothèques de modèles communautaires, des outils tiers, des solutions de déploiement et des ressources didactiques est bien inférieure à celle de PyTorch. À la mi-2026, le nombre de packages liés à JAX sur PyPI représente environ 1/10 de l'écosystème PyTorch.
- Difficultés de débogage : le message d'erreur de la fonction compilée n'est pas assez intuitif et le débogueur Python (pdb) à l'intérieur de
jita une prise en charge limitée. Bien quejax.debugetjax.make_jaxpraméliorent la situation, l'expérience globale de débogage est toujours en retard par rapport au mode impatient de PyTorch. - Risque stratégique Google : le développement principal de JAX est dirigé par Google, avec une influence limitée des contributeurs externes. Il existe une situation parallèle entre les doubles frameworks TensorFlow/JAX au sein de Google, et il existe une incertitude quant à l'orientation à long terme de la feuille de route technique.
Points d'observation de suivi :
- Unification interne de Google : Google DeepMind unifiera-t-il les itinéraires techniques de TensorFlow et de JAX au cours des 2-3 prochaines années, ou clarifiera-t-il le statut de JAX en tant que seul cadre de recherche.
- Taux de croissance écologique : l'écosystème JAX peut-il réduire l'écart avec PyTorch dans les dimensions de la bibliothèque de modèles (proportion de modèles Hugging Face JAX/Flax) et de la chaîne d'outils (débogueur Profiler, plan de déploiement).
- GPU AMD et prise en charge Apple Silicon : La maturité de la prise en charge de JAX pour le matériel non NVIDIA aura un impact direct sur l'expansion de son adoption.
- Structure de gouvernance communautaire : Google établira-t-il un modèle de gouvernance communautaire plus ouvert (tel que la Fondation JAX) pour réduire le risque de dépendance à l'égard d'une seule entreprise.
Évaluation des risques en matière d'approvisionnement et d'adoption :
- Pour les équipes de recherche de pointe (visant à publier les meilleurs articles de conférence et à explorer de nouvelles architectures) : JAX est une compétence fondamentale qui doit être maîtrisée. Il est recommandé d'investir 1 à 2 ingénieurs pour apprendre d'abord et établir des capacités JAX internes dans un délai de 3 à 6 mois.
- Pour l'équipe de formation de grands modèles (modèle de paramètre de formation cible 10B+) : la solution JAX + TPU est toujours en avance sur la solution PyTorch + GPU en termes d'efficacité de mise à l'échelle (en particulier la taille de puce 512+), mais la disponibilité et le coût de Google Cloud TPU doivent être évalués. Il est recommandé de demander d'abord le quota TPU gratuit de Google TRC et d'effectuer une vérification technique pendant 4 à 8 semaines.
- Pour les équipes ML petites à moyennes (réglage/inférence du modèle en dessous de l'objectif 7B) : JAX n'est pas recommandé. PyTorch dispose d'une meilleure chaîne d'outils, d'un meilleur support communautaire et d'un meilleur vivier de talents, et les coûts cachés de l'adoption de JAX (embauche, formation, migration) peuvent dépasser les gains de performances. Si la maturité écologique du JAX s’améliore significativement dans le futur, elle pourra être réévaluée en 2027-2028.
Outils associés :
Visage câlin, replicate
Informations de version
- JAX 0.5.0 :Il n’y a pas encore de date officielle précise. Améliorations continues des performances de compilation XLA et des noyaux Pallas.
- JAX 0.4.35 :Il n’y a pas encore de date officielle précise. Prise en charge améliorée et optimisation des performances pour les GPU AMD.
Avis des utilisateurs