Revenir au contenu principal
Data Science et ML

Entraînement PyTorch rapide et tolérant aux pannes sur le Runtime IA

Comment les choix de chargement des données et de points de contrôle déterminent l'utilisation des GPU, le coût de récupération et la facture d'entraînement à grande échelle, et les API AI Runtime pour y parvenir.

par Bruce Fontaine

  • À grande échelle, les pannes de GPU sont la norme et non l'exception ; le code doit être conçu pour y survivre.
  • Les sauvegardes asynchrones et distribuées de points de contrôle de Torch rendent cette opération presque gratuite, ce qui permet des sauvegardes plus fréquentes et réduit le coût de récupération.
  • La sauvegarde des points de contrôle du modèle seule ne suffit pas ; celle du pipeline de données évite la corruption silencieuse des données d'entraînement lors de la reprise.

À grande échelle, l'efficacité de votre entraînement est déterminée par une seule métrique : le « goodput », c'est-à-dire la proportion de temps que vos GPU passent à effectuer des calculs productifs plutôt qu'à attendre ou à se remettre de pannes. Comme les pannes de GPU sont monnaie courante à grande échelle, la capacité à se remettre rapidement et automatiquement d'une panne est le seul moyen de maintenir un bon niveau de goodput et de gérer vos dépenses totales en GPU.

Deux sous-systèmes conditionnent la réussite de cette récupération, et pourtant, ils sont souvent relégués au second plan : le pipeline de données qui alimente vos accélérateurs et le mécanisme de checkpointing qui capture l'état pour qu'un travail puisse reprendre. Si l'un ou l'autre est mal configuré, chaque panne vous coûtera beaucoup plus de temps d'inactivité des GPU qu'elle ne le devrait. Même en dehors des scénarios de panne, un pipeline de données qui ne parvient pas à suivre le rythme de vos accélérateurs privera silencieusement vos GPU de ressources et réduira le goodput tout aussi sûrement qu'un plantage. Nous allons passer en revue les mécanismes et les compromis de chacun d'eux, et voir comment ils influencent votre goodput et vos dépenses totales en GPU. Consultez le guide d'accompagnement sur les performances et la résilience de l'entraînement pour obtenir des exemples et des indications de code.

Pour l'aspect infrastructure de ce même problème, à savoir comment un parc détecte et isole les GPU défaillants avant qu'ils n'interrompent un travail, consultez l'article complémentaire, Comment nous assurons la fiabilité des GPU au sein de Databricks AI.

Pourquoi les pannes sont inévitables à grande échelle

À mesure que le nombre de GPU d'un travail augmente, la probabilité qu'il se termine sans interruption diminue rapidement. Un modèle d'estimation rapide issu de l'article complémentaire de Databricks suppose que chaque GPU présente un taux de panne annuel d'environ 1 %. Selon cette hypothèse, l'article note qu'« un travail sur 256 GPU s'exécutant pendant 30 jours a environ 19 % de chances de subir une panne. Avec 1 024 GPU, ce risque grimpe à 57 % », et il ne s'agit là que de problèmes liés à l'infrastructure.

Pour ancrer cette estimation dans la réalité, le supercalculateur Delta de 608 GPU H100 a connu des pannes toutes les 1,9 heure, ce qui signifie que pour un travail sur 32 GPU, le temps moyen avant panne serait de 36 heures. Le point clé à retenir est que votre travail d'entraînement finira probablement par échouer à un moment donné, et que prendre les bonnes d'écisions peut rendre votre modèle résilient et réduire le temps total perdu lorsque cela se produit.

Impact 1 : Le format de checkpoint détermine la fréquence de sauvegarde que vous pouvez vous permettre

C'est au niveau du checkpointing que se joue la résilience, et le mécanisme que vous choisissez a un effet de premier ordre sur la fréquence à laquelle vous pouvez sauvegarder. C'est le levier le plus important pour votre goodput : si vous effectuez un checkpoint une fois par jour, une panne nécessite de réexécuter en moyenne 12 heures de travail en double pour revenir à l'état initial avant la panne.

Le goulot d'étranglement monolithique de torch.save

Le premier checkpoint que la plupart des équipes écrivent est un simple torch.save sur le rank 0. Selon la façon dont votre modèle est entraîné, deux problèmes peuvent se poser :

  1. Pour l'entraînement distribué, il rassemble tous les états sur le rank 0 et écrit un seul fichier.
  2. Un seul processus écrit l'intégralité du checkpoint de manière synchrone. Cela peut être bloqué par des éléments tels que les transferts réseau lors de la sauvegarde vers des magasins d'objets distants comme Unity Catalog (UC).

image6.png

Ce comportement bloquant laisse vos GPU inactifs, ce qui réduit votre goodput. Mais il existe un moyen de réduire le temps que votre GPU passe à effectuer le checkpointing : l'API de checkpoint distribué de Torch.

Checkpoint distribué (DCP) : chaque rank écrit son propre shard

Le checkpoint distribué de PyTorch inverse cette logique. Chaque rank écrit son propre shard distinct en parallèle, accompagné d'un petit fichier .metadata décrivant comment les shards s'assemblent pour former les tenseurs complets.

image1.png

Le temps de sauvegarde diminue d'environ 1/N avec le nombre de ranks et, comme le fichier .metadata enregistre la disposition globale, le même checkpoint peut être rechargé sur un nombre différent de GPU. DCP replanifie les octets dont chaque nouveau rank a besoin, de sorte que la récupération sur un cluster à capacité réduite après la perte de nœuds fonctionne parfaitement.

DCP en vaut la peine, même pour les travaux de parallélisme de données classiques

On suppose souvent que DCP est réservé aux modèles partitionnés (sharded), et qu'un travail de parallélisme de données (DDP), où chaque rank détient une réplique identique des poids, n'a rien à y gagner. C'est faux : DCP partitionne l'état du modèle et l'écrit en parallèle sur chaque worker, même pour les tâches d'entraînement DDP.

C'est également la même API dont vous aurez besoin le jour où vous passerez à FSDP ou au parallélisme de tenseurs. L'adopter tôt vous évite donc d'avoir à réécrire le code de résilience au pire moment possible.

Les sauvegardes asynchrones rendent la fréquence presque gratuite

Même avec des écritures parallèles, une sauvegarde synchrone bloque l'entraînement jusqu'à ce que les octets soient sécurisés dans le stockage, ce qui, pour un checkpoint volumineux vers un volume distant, représente des dizaines de secondes d'inactivité de l'accélérateur. async_save divise l'opération : une copie rapide vers un tampon de transition (staging buffer), puis un chargement en arrière-plan qui se superpose à la poursuite de l'entraînement.

image4.png

La boucle d'entraînement ne paie que pour la copie de transition, pas pour le chargement. Un checkpoint qui coûtait auparavant des dizaines de secondes d'inactivité ne coûte désormais presque rien, ce qui rend précisément abordable le checkpointing fréquent décrit dans la section suivante.

Sur AI Runtime, UCVolumeWriter et UCVolumeReader implémentent DCP sur les volumes UC, en faisant transiter les I/O par le NVMe local et en ne marquant un checkpoint comme terminé qu'une fois ses données entièrement enregistrées. Consultez le guide sur les performances et la résilience pour obtenir tous les détails et des exemples de code.

Travail d'entraînementGain de temps de async_save par rapport à torch.save
LLM DDP avec 2,8 milliards de paramètres sur 32xH1001,8x (36 s vs 66 s)
LLM FSPD avec 20 milliards de paramètres sur 32xH10058x (522 s vs 9 s)

Les données ci-dessus excluent le temps de stockage réseau pour torch.save.

Impact 2 : La fréquence de checkpoint détermine votre coût de récupération

C'est là que les éléments s'accumulent. Lorsqu'un travail échoue, il perd tout depuis le dernier checkpoint valide et doit tout recalculer. Ainsi, le travail perdu attendu par panne correspond à environ la moitié de l'intervalle de checkpoint, et les sauvegardes asynchrones économiques vous permettent de réduire cet intervalle.

Diviser l'intervalle par 10 réduit d'autant le temps de récupération attendu. Rappelez-vous le chiffre de Llama 3 d'environ 8,6 interruptions par jour : avec ce taux de panne, effectuer un checkpoint toutes les 2 heures signifie que vous perdrez environ 8,6 heures par jour en réentraînement, soit un goodput de 64 %. En effectuant un checkpoint toutes les 30 minutes, vous ne perdez que 2,15 heures, soit un goodput de 91 %.

La récupération doit également être automatique. Au redémarrage, le travail doit trouver le checkpoint le plus récent dont l'écriture est terminée, ignorer ceux qui ont été interrompus par le plantage, et reprendre à partir de là sans intervention humaine. DCP rend cela fiable : le fichier .metadata n'est écrit qu'une fois tous les shards enregistrés, sa présence est donc un indicateur fiable que « cette sauvegarde est terminée ».

image5.png

Impact 3 : Le chargement des données détermine si vos GPU restent inactifs

Un travail d'entraînement progresse à la vitesse de son entrée la plus lente. Lorsque les accélérateurs attendent le lot suivant, votre goodput diminue car vos GPU sont tout simplement inactifs. La seule façon de résoudre ce problème est de s'assurer que votre pipeline d'entrée superpose la préparation des données de l'étape suivante aux calculs de l'étape en cours, comme le montre la figure ci-dessous :

image2.png

Nous constatons souvent que les clients qui passent à une superposition du chargement des données et du calcul enregistrent une baisse de 20 à 50 % du temps d'exécution réel.

Le coût de la lecture directe depuis un stockage distant

Sur une plateforme gouvernée, les données d'entraînement résident dans un stockage d'objets distant. Sur AI Runtime, les volumes Unity Catalog (UC) sont présentés sous forme de montages réseau.

Lire les fichiers directement depuis ce montage à chaque accès lie le temps de votre étape à la latence du réseau et télécharge à nouveau les mêmes fichiers à chaque époque. La solution consiste à utiliser un chargeur de données (dataloader) qui copie chaque fichier sur un stockage local rapide lors du premier accès, sert les lectures suivantes depuis ce cache local et récupère les fichiers suivants en parallèle pendant que le GPU effectue les calculs.

image7.png

Avec AI Runtime, UCVolumeDataset et DataLoader font exactement cela (voir le guide pour des exemples de code). UCVolumeDataset lit en continu les fichiers d'un volume UC, met chacun d'eux en cache sur le NVMe local lors du premier accès, et partitionne les fichiers entre les rangs et les workers pour que chaque accélérateur obtienne une tranche disjointe et sans chevauchement. Notre DataLoader est une sous-classe prête à l'emploi de DataLoader de PyTorch dont les valeurs par défaut sont optimisées pour ce chemin, de sorte que les fichiers sont récupérés et mis en cache de manière simultanée pendant que le GPU effectue ses calculs, plutôt qu'un par un sur le thread d'entraînement.

Exemple : entraînement d'un modèle d'image à partir de fichiers UC

Prenons l'exemple d'une charge de travail simple de classification d'images : décoder des fichiers JPEG à partir d'un volume UC, appliquer des augmentations et entraîner un modèle de vision. Voyons deux façons de procéder avec les mêmes GPU, modèle et taille de lot (batch size) : le Dataset PyTorch standard lisant à partir d'un volume UC, par rapport à UCVolumeDataset combiné avec les valeurs par défaut du DataLoader de Databricks.

Métrique (par GPU, régime permanent)PyTorch standard DataLoader, lecture directe depuis UCUCVolumeDataset + Databricks DataLoader
Débit de l'époque 1 (images/sec)57.2417
Débit de l'époque 2 (images/sec)371.66590
Utilisation du GPU (%)12.6%53.3%

Plus besoin de deviner où passe le temps

Lors du développement de DataLoader, nous avons veillé à ce qu'il enregistre ses métriques dans MLflow, ce qui permet de voir d'un coup d'œil si votre pipeline de données bloque l'entraînement.

image8.png

La métrique fetch_seconds mesure explicitement le temps nécessaire au dataloader pour produire un lot (batch), période pendant laquelle votre GPU reste inactif.

Impact 4 : Oublier le pipeline de données corrompt silencieusement votre modèle

Il existe un dernier bug de résilience qui ne produit aucun message d'erreur, aucun plantage ni aucun échec de tâche, mais simplement un modèle légèrement moins performant qu'il ne devrait l'être. Cela se produit lorsque vous sauvegardez un point de contrôle (checkpoint) du modèle, de l'optimiseur et de l'étape, mais pas de la position de votre pipeline de données au sein du jeu de données.

Imaginez une tâche interrompue au milieu d'une époque. Elle restaure correctement le modèle et reprend la boucle d'entraînement, mais le dataloader recommence depuis le début du jeu de données.

La tâche reprise s'entraîne à nouveau sur des exemples qu'elle a déjà vus au cours de cette époque et ignore potentiellement ceux qu'elle n'avait pas encore atteints. Avec les nombreux redémarrages devenus courants à grande échelle, cela biaise silencieusement la distribution de vos données. Le modèle continue de s'entraîner, mais sur un mauvais échantillonnage de vos données. C'est précisément le genre de défaillance silencieuse la plus coûteuse, car la tâche se termine correctement et personne ne remarque de problème avant d'obtenir des métriques décevantes.

La solution consiste à traiter la position des données comme faisant partie du point de contrôle (checkpoint). Selon votre pipeline, cela signifie suivre un décalage (offset) d'échantillon ou de fragment (shard) et passer directement à cette position lors de la reprise, faire en sorte qu'un jeu de données personnalisé sérialise sa propre position, ou créer des points de contrôle aux limites des époques. Tout cela repose sur un prérequis unique : le déterminisme. Le mélange (shuffling) et l'augmentation s'appuient sur des générateurs de nombres aléatoires. Ces graines (seeds) et états de RNG doivent donc également faire partie du point de contrôle, sans quoi l'ordre des données après un redémarrage ne correspondra pas à l'ordre précédent, et une position sauvegardée pointera vers les mauvais échantillons.

La graine (seed), l'ordre reproductible et le pipeline de données reprenable sont trois expressions d'une même idée. Le guide détaille chaque stratégie avec des exemples de code.

Résumé

Un entraînement rapide et tolérant aux pannes repose sur quelques décisions clés qui se complètent :

  1. Utilisez Distributed Checkpoint au lieu de torch.save, même pour le DDP, afin que les sauvegardes soient parallèles et peu coûteuses plutôt que de constituer un goulot d'étranglement séquentiel.
  2. Sauvegardez de manière asynchrone pour que les points de contrôle soient presque gratuits, ce qui vous permet de sauvegarder souvent.
  3. Restaurez automatiquement le point de contrôle valide le plus récent, de sorte qu'une panne ne coute que quelques minutes de calcul au lieu de plusieurs heures.
  4. Superposez le chargement des données et le calcul en mettant en cache et en pré-récupérant (prefetching) les données depuis le stockage distant, afin que les accélérateurs ne restent jamais inactifs en attendant des entrées. Ce sont des heures-GPU économisées à chaque étape.
  5. Créez des points de contrôle pour le pipeline de données et l'état de la RNG, afin qu'une tâche reprise continue sur les bonnes données au lieu de corrompre silencieusement votre modèle.

Le principe unificateur : des points de contrôle fréquents, économiques et complets transforment une panne matérielle d'un événement fatal pour la tâche en une simple erreur d'arrondi, tandis qu'un pipeline d'entrée superposé maintient les accélérateurs occupés dans l'intervalle. Les sauvegardes peu coûteuses (asynchrones) rendent la fréquence abordable ; les sauvegardes complètes (modèle, données et RNG) garantissent une restauration correcte. Avec ces deux éléments en place, et un parc de machines qui détecte et isole le matériel défaillant, votre temps d'entraînement effectif frôle le maximum autorisé par le matériel, quelle que soit l'instabilité du cluster sous-jacent.

Références

Prêt à essayer ? Consultez le Guide sur les performances et la résilience de l'entraînement dans la documentation de Databricks AI Runtime pour obtenir le code complet, et lisez Comment nous assurons la fiabilité des GPU au sein de Databricks AI pour découvrir le volet infrastructure.

(Cet article de blog a été traduit à l'aide d'outils basés sur l'intelligence artificielle) Article original

Recevez les derniers articles dans votre boîte mail

Abonnez-vous à notre blog et recevez les derniers articles directement dans votre boîte mail.