À propos du projet

JAX fournit un système de transformations de fonctions composables pour les programmes Python et NumPy. Il est conçu pour mettre à l'échelle le calcul numérique sur des accélérateurs matériels, notamment les GPU et TPU, en utilisant XLA (Accelerated Linear Algebra) pour la compilation. Les capacités clés incluent : - Différenciation automatique : via `jax.grad`, il prend en charge la différenciation en mode inverse et en mode direct à travers les boucles Python, les branches, la récursion et les fermetures, permettant des dérivées d'ordre supérieur. - Compilation Just-In-Time (JIT) : via `jax.jit`, il compile des fonctions pures via XLA pour améliorer les performances d'exécution sur les accélérateurs. - Auto-vectorisation : via `jax.vmap`, il mappe des fonctions le long des axes de tableaux en poussant les boucles vers les opérations primitives, éliminant ainsi le besoin d'une gestion manuelle des dimensions de batch. - Mise à l'échelle et parallélisation : prend en charge divers modes de mise à l'échelle, notamment la parallélisation automatique basée sur le compilateur, le sharding explicite et la programmation manuelle par appareil avec des collectifs explicites. JAX est compatible avec plusieurs plateformes, dont Linux, macOS et Windows, avec un support pour CPU, NVIDIA GPU, Google TPU, AMD GPU, et un support expérimental pour les GPU Apple et Intel.