Bibliothèque / SDK
jax-ml/jax avatar
jax-ml/jax

jax : ce que le dépôt permet réellement de faire

Transformations composables de programmes Python+NumPy : différenciation, vectorisation, JIT vers GPU/TPU, et plus encore.

36 304 étoiles3 779 forksPythonApache-2.0

En bref

De quoi s’agit-il ?
Composable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more.
À qui s’adresse-t-il ?
jax convient aux lecteurs dont le besoin correspond aux fonctions décrites dans le README et qui peuvent fournir son environnement requis. Il convient moins à ceux qui attendent une garantie absente du dépôt.
Puis-je l’utiliser commercialement ?
Oui. Apache-2.0 est une licence permissive : vous pouvez utiliser, modifier et vendre un logiciel qui en dépend, à condition de conserver les mentions de droit d’auteur et de licence.
Est-il encore maintenu ?
Oui. Le dépôt a reçu de nouveaux commits au cours des dernières 24 heures.
En quel langage est-il écrit ?
Principalement Python, d’après les statistiques de langage de GitHub.

Ces réponses reposent sur les données GitHub du projet (dernière synchronisation le 15 septembre 2026) et sur notre analyse. Elles ne constituent pas un avis juridique.

ANALYSE OPEN SOURCE APPROFONDIE

Une bibliothèque pour transformer des fonctions numériques

JAX est une bibliothèque Python pour le calcul sur tableaux orienté accélérateurs et la transformation de programmes, conçue pour le calcul numérique haute performance et l'apprentissage automatique à grande échelle. L'affirmation centrale du README est que JAX peut différencier automatiquement des fonctions Python et NumPy natives, y compris à travers les boucles, les branchements, la récursion et les fermetures. La différenciation en mode inverse est exposée via jax.grad, le mode direct est également pris en charge, et les deux peuvent être composés dans n'importe quel ordre. En dessous, XLA compile et exécute les programmes NumPy sur les TPU, GPU et autres accélérateurs matériels. Le README précise aussi que la compilation et la différenciation automatique peuvent être composées.

Trois transformations au cœur du système : grad, jit et vmap

Le README présente trois transformations comme le cœur du système. jax.grad calcule les gradients en mode inverse ; les exemples montrent une fonction tanh différenciée une fois puis trois fois, ainsi qu'une fonction avec une branche if/else. jax.jit compile les fonctions de bout en bout avec XLA, utilisable comme décorateur ou comme fonction d'ordre supérieur, et contraint le type de flux de contrôle Python qu'une fonction peut utiliser. jax.vmap applique une fonction le long des axes d'un tableau en repoussant la boucle vers les opérations primitives de la fonction, transformant des multiplications matrice-vecteur en multiplications matrice-matrice et évitant d'avoir à transporter les dimensions de lot dans le code. Pour aller plus loin, le README renvoie à l'Autodiff Cookbook et à la documentation de référence.

Composer grad, jit et vmap sur une même fonction de perte

L'exemple principal du README définit une fonction de prédiction qui boucle sur des paires de paramètres, une perte d'erreur quadratique, puis construit grad_loss = jax.jit(jax.grad(loss)) et perex_grads = jax.jit(jax.vmap(grad_loss, in_axes=(None, 0, 0))) pour obtenir rapidement des gradients par exemple. Le même motif réapparaît dans la section sur le passage à l'échelle : une perte traitée par jit et grad est évaluée sur des paramètres et des données explicitement shardés. La sortie de l'exemple imprime des types comme f32[512@data,512], montrant que le tenseur de paramètres porte une annotation de sharding sur l'axe data.

Trois modes de passage à l'échelle, de l'automatique au manuel

Pour passer à l'échelle sur des milliers d'appareils, le README documente trois approches. La parallélisation automatique pilotée par le compilateur permet de programmer comme sur une seule machine globale, le compilateur choisissant comment sharder les données et partitionner le calcul, avec quelques contraintes fournies par l'utilisateur. Le sharding explicite avec partitionnement automatique conserve la vue globale mais rend les shardings de données visibles dans les types JAX, inspectables avec jax.typeof. La programmation manuelle par appareil passe à une vue par appareil et autorise des collectifs explicites. Le tableau du README résume les trois modes selon la vue, le sharding explicite et les collectifs explicites.

Plates-formes prises en charge et commandes d'installation documentées

La matrice de plates-formes du README liste un support CPU sur Linux x86_64, Linux aarch64, Mac aarch64, Windows x86_64 et Windows WSL2 x86_64. Le GPU NVIDIA est pris en charge sur les deux variantes Linux, expérimental sous Windows WSL2 et indisponible ailleurs. Le TPU Google est limité à Linux x86_64. Le GPU AMD est pris en charge sur Linux x86_64 et expérimental sous Windows WSL2. Le GPU Apple est expérimental sur Mac aarch64, et le GPU Intel expérimental sur Linux x86_64. Les commandes d'installation indiquées dans le README sont pip install -U jax pour le CPU, pip install -U "jax[cuda13]" pour le GPU NVIDIA, pip install -U "jax[tpu]" pour le TPU Google et pip install -U "jax[rocm7-local]" pour le GPU AMD sous Linux. Pour le GPU Intel, il faut suivre les instructions d'Intel. Pour la compilation depuis les sources, Docker, d'autres versions de CUDA et une construction conda communautaire, le README renvoie à la documentation.

Statut du projet, citation et licence Apache-2.0

Le README déclare clairement que JAX est un projet de recherche, pas un produit officiel de Google, et prévient qu'il a des angles vifs, renvoyant au notebook Gotchas et au suivi des problèmes. L'entrée de citation liste les auteurs par ordre alphabétique, utilise la version 0.3.13 de jax/version.py et donne 2018 comme année de la publication open source. Une version naissante, ne prenant en charge que la différenciation automatique et la compilation XLA, a été décrite dans un article de SysML 2018 ; le README indique qu'un article plus complet et plus à jour est en préparation. Le dépôt est sous Apache-2.0, qui accorde une licence de droit d'auteur perpétuelle, mondiale, non exclusive, gratuite, sans redevance et irrévocable, ainsi qu'une licence de brevet correspondante qui prend fin si le licencié intente une action en contrefaçon de brevet. L'extrait de licence ne traite ni de garantie, ni de support, ni de posture de sécurité. Les métadonnées du dépôt montrent 36 103 étoiles, 3 716 forks et 2 554 problèmes ouverts au moment de la préparation de cet article.

Contrôler jax dans son environnement

Installez JAX selon la commande du README avec le backend correspondant, exécutez un calcul jax.numpy puis jax.jit et vérifiez le device choisi. Comparez ensuite grad ou vmap sur un exemple minimal et inspectez les erreurs de compilation ou de mémoire. Cette séquence porte sur jax, ses fichiers et ses sorties documentées ; elle ne remplace pas les informations absentes du README.

Conclusion éditoriale

jax convient aux lecteurs dont le besoin correspond aux fonctions décrites dans le README et qui peuvent fournir son environnement requis. Il convient moins à ceux qui attendent une garantie absente du dépôt. Avant décision, exécutez ce contrôle propre au projet : Installez JAX selon la commande du README avec le backend correspondant, exécutez un calcul jax.numpy puis jax.jit et vérifiez le device choisi. Comparez ensuite grad ou vmap sur un exemple minimal et inspectez les erreurs de compilation ou de mémoire.

Sources officielles

  1. Official documentation
  2. Official README
  3. Project repository
  4. Release notes
Notes de la communauté

Notes de la communauté