Créer une IA de production sur des Cloud TPU avec JAX

La pile d'IA JAX étend le cœur numérique JAX avec une collection de bibliothèques composables soutenues par Google, ce qui en fait une plate-forme Open Source de bout en bout robuste pour le machine learning à des échelles extrêmes. À ce titre, la pile d'IA JAX se compose d'un écosystème complet et robuste qui couvre l'ensemble du cycle de vie du ML :

  • Fondation à l'échelle industrielle : la pile d'IA JAX est conçue pour une mise à l'échelle massive, en s'appuyant sur ML Pathways pour orchestrer l'entraînement sur des dizaines de milliers de puces et sur Orbax pour un checkpointing asynchrone résilient et à haut débit, ce qui permet un entraînement de qualité production de modèles de pointe.

  • Boîte à outils complète et prête pour la production : la pile d'IA JAX fournit un ensemble complet de bibliothèques pour l'ensemble du processus de développement : Flax pour la création de modèles flexibles, Optax pour les stratégies d'optimisation composables et Grain pour les pipelines de données déterministes essentiels aux exécutions reproductibles à grande échelle.

  • Performances de pointe et spécialisées : pour maximiser l'utilisation du matériel, la pile d'IA JAX propose des bibliothèques spécialisées, y compris Tokamax pour les noyaux personnalisés de pointe, Qwix pour la quantification non intrusive qui améliore la vitesse d'entraînement et d'inférence, et XProf pour le profilage des performances approfondi et intégré au matériel.

  • Chemin complet vers la production : la pile d'IA JAX permet une transition fluide de la recherche au déploiement. Cela inclut MaxText comme référence évolutive pour l'entraînement des modèles de fondation, Tunix pour l'apprentissage par renforcement (RL) et l'alignement de pointe, ainsi qu'une solution d'inférence unifiée avec l'intégration vLLM TPU et l'environnement d'exécution JAX pour le service.

La philosophie de la pile d'IA JAX repose sur des composants faiblement couplés, chacun d'eux étant spécialisé dans une tâche. Plutôt que d'être un framework de ML monolithique, JAX est lui-même de portée limitée et se concentre sur les opérations de tableaux et les transformations de programmes efficaces. L'écosystème est basé sur ce framework principal pour fournir un large éventail de fonctionnalités, liées à l'entraînement des modèles de ML et à d'autres types de charges de travail telles que le calcul scientifique.

Ce système de composants faiblement couplés vous permet de sélectionner et de combiner des bibliothèques de la manière la plus adaptée à vos besoins. Du point de vue de l'ingénierie logicielle, cette architecture vous permet également de mettre à jour de manière itérative les fonctionnalités qui seraient traditionnellement considérées comme des composants de framework de base (par exemple, les pipelines de données et la création de points de contrôle), sans risquer de déstabiliser le framework de base ni d'être pris dans les cycles de publication. Étant donné que la plupart des fonctionnalités sont implémentées dans des bibliothèques plutôt que dans des modifications apportées à un framework monolithique, cela rend la bibliothèque de nombres de base plus durable et adaptable aux futurs changements du paysage technologique.

Les sections suivantes présentent un aperçu technique de la pile d'IA JAX, de ses principales fonctionnalités, des décisions de conception qui les sous-tendent et de la manière dont elles se combinent pour créer une plate-forme durable pour les charges de travail de ML modernes.

Pile JAX AI et autres composants de l'écosystème

Composant Fonction / Description
Composants et cœur de la pile JAX AI1
JAX Calcul de tableaux et transformation de programmes orientés accélérateur (JIT, grad, vmap, pmap).
Flax Bibliothèque flexible de création de réseaux neuronaux pour une création et une modification intuitives des modèles.
Optax Bibliothèque de transformations composables pour le traitement et l'optimisation des gradients.
Orbax Bibliothèque de point de contrôle distribuée "toute échelle" pour la résilience de l'entraînement à l'échelle héroïque.
Grain Bibliothèque de pipeline de données d'entrée évolutive, déterministe et vérifiable.
Pile d'IA JAX : infrastructure
XLA Compilateur de machine learning Open Source pour les TPU, les processeurs et les GPU.
Pathways Runtime distribué pour orchestrer le calcul sur des dizaines de milliers de puces.
Pile JAX AI : avancé Développement
Pallas Extension JAX permettant d'écrire des noyaux personnalisés de bas niveau et hautes performances implémentés en Python.
Tokamax Une bibliothèque organisée de noyaux personnalisés hautes performances et de pointe (par exemple, Attention).
Qwix Une bibliothèque complète et non intrusive pour la quantification (PTQ, QAT, QLoRA).
Pile JAX AI : application
MaxText / MaxDiffusion Frameworks de référence phares et évolutifs pour l'entraînement des modèles de fondation (par exemple, LLM et diffusion).
Tunix Framework pour l'entraînement et l'alignement post-entraînement de pointe (RLHF, DPO).
vLLM Solution d'inférence LLM hautes performances utilisant l'intégration intégrée du framework vLLM.
XProf Profileur intégré au matériel pour une analyse des performances à l'échelle du système.

1 Inclus dans le package Python jax-ai-stack.

Figure 1 : Composants de la pile et de l'écosystème JAX AI

Pile JAX AI

L'impératif architectural : des performances au-delà des frameworks

Alors que les architectures de modèles convergent (par exemple, sur les Transformers multimodaux Mixture-of-Experts (MoE)), la recherche de performances maximales conduit à l'émergence des Megakernels. Un mégakernel correspond à l'intégralité (ou une grande partie) de la passe avant d'un modèle spécifique, codée manuellement à l'aide d'une API de niveau inférieur comme le SDK CUDA sur les GPU NVIDIA. Cette approche permet d'utiliser au maximum le matériel en chevauchant de manière agressive le calcul, la mémoire et la communication. Des travaux récents de la communauté de recherche ont démontré que cette approche peut générer des gains de débit importants (plus de 22 % dans certains cas) pour l'inférence sur les GPU. Cette tendance ne se limite pas à l'inférence. Des éléments suggèrent que certains efforts d'entraînement à grande échelle ont impliqué un contrôle matériel de bas niveau pour obtenir des gains d'efficacité importants.

Si cette tendance s'accélère, tous les frameworks de haut niveau tels qu'ils existent aujourd'hui risquent de devenir moins pertinents, car l'accès de bas niveau au matériel est ce qui compte en fin de compte pour les performances sur les architectures matures et stables. Cela représente un défi pour toutes les piles ML modernes : comment fournir un contrôle matériel de niveau expert sans sacrifier la productivité et la flexibilité d'un framework de haut niveau ?

Pour que les TPU offrent une voie claire vers ce niveau de performances, l'écosystème doit exposer une couche d'API plus proche du matériel, permettant le développement de ces kernels hautement spécialisés. La pile JAX est conçue pour résoudre ce problème en offrant un continuum d'abstraction (voir la figure 2), des optimisations automatisées de haut niveau du compilateur XLA au contrôle manuel et précis de la bibliothèque de création de noyaux Pallas.

Figure 2 : Continuum d'abstraction JAX

Continuum d'abstraction JAX

Pile JAX AI de base

La pile d'IA JAX de base se compose de cinq bibliothèques clés qui fournissent les bases du développement de modèles :

JAX : une base pour la transformation de programmes composables et hautes performances

JAX est une bibliothèque Python pour le calcul de tableaux et la transformation de programmes orientés accélérateur. Elle est conçue pour le calcul numérique hautes performances et le machine learning à grande échelle. Avec son modèle de programmation fonctionnelle et son API de type NumPy, JAX fournit une base solide pour les bibliothèques de niveau supérieur.

Grâce à sa conception axée sur le compilateur, JAX favorise intrinsèquement l'évolutivité en tirant parti de XLA (voir la section XLA) pour une analyse, une optimisation et un ciblage matériel agressifs et complets. L'accent mis par JAX sur la programmation fonctionnelle (par exemple, les fonctions pures) rend ses transformations de programme de base plus faciles à gérer et, surtout, composables.

Ces transformations de base peuvent être combinées pour obtenir des performances élevées et une mise à l'échelle des charges de travail en fonction de la taille du modèle, de la taille du cluster et des types de matériel :

  • jit : compilation à la volée des fonctions Python en exécutables XLA optimisés et fusionnés.
  • grad : différenciation automatique, compatible avec les modes forward et reverse, ainsi qu'avec les dérivées d'ordre supérieur.
  • vmap : vectorisation automatique, permettant le traitement par lot et le parallélisme des données sans modifier la logique de la fonction.
  • pmap / shard_map : parallélisation automatique sur plusieurs appareils (par exemple, les cœurs de TPU), qui constitue la base de l'entraînement distribué.

L'intégration fluide avec le modèle GSPMD (General-purpose SPMD) de XLA permet à JAX de paralléliser automatiquement les calculs sur de grands pods TPU avec un minimum de modifications de code. Dans la plupart des cas, la mise à l'échelle ne nécessite que des annotations de sharding de haut niveau.

Flax : création flexible de réseaux de neurones

Flax simplifie la création, le débogage et l'analyse des réseaux de neurones dans JAX en fournissant une approche intuitive et orientée objet pour la création de modèles. Bien que l'API fonctionnelle de JAX soit puissante, elle offre une abstraction basée sur les couches plus familière aux développeurs habitués aux frameworks tels que PyTorch, sans aucune perte de performances.

Cette conception simplifie la modification ou la combinaison des composants du modèle entraîné. Les techniques telles que LoRA et la quantification nécessitent des définitions de modèle manipulables, que l'API NNX de Flax fournit via une interface Pythonique. NNX encapsule l'état du modèle, réduit la charge cognitive de l'utilisateur et permet la traversée et la modification programmatiques de la hiérarchie du modèle.

Points forts :

  • API orientée objet intuitive : simplifie la construction de modèles et permet des cas d'utilisation avancés tels que le remplacement de sous-modules et l'initialisation partielle.
  • Cohérence avec Core JAX : Flax fournit des transformations liftées entièrement compatibles avec le paradigme fonctionnel de JAX, offrant toutes les performances de JAX avec une convivialité améliorée pour les développeurs.

Optax : stratégies composables de traitement et d'optimisation des gradients

Optax est une bibliothèque de traitement et d'optimisation des gradients pour JAX. Il est conçu pour fournir aux créateurs de modèles des blocs de construction qui peuvent être recombinés de manière personnalisée afin d'entraîner des modèles de deep learning, entre autres applications. Elle s'appuie sur les capacités de la bibliothèque JAX principale pour fournir une bibliothèque de fonctions de perte et d'optimiseur hautes performances et bien testée, ainsi que des techniques associées qui peuvent être utilisées pour entraîner des modèles de ML.

Motivation

Le calcul et la minimisation des pertes sont au cœur de ce qui permet l'entraînement des modèles de ML. Grâce à sa prise en charge de la différenciation automatique, la bibliothèque JAX principale fournit les capacités numériques nécessaires à l'entraînement des modèles, mais elle ne fournit pas d'implémentations standards des optimiseurs populaires (par exemple, RMSProp ou Adam) ni des pertes (par exemple,