Aller au contenu principal

Améliorer les performances et la résilience de l'entraînement sur AI Runtime

info

Aperçu

Cette fonctionnalité est en Aperçu public.

À mesure qu'un Job monte en charge sur davantage de GPU, la probabilité de défaillance matérielle et outil/solution/technologie/plateforme augmente. Cette page couvre les stratégies pour rendre vos exécutions d'entraînement plus rapides et plus tolérantes aux pannes :

Avec ces modèles, la création de points de contrôle de votre modèle est peu coûteuse ; vous pouvez donc créer des points de contrôle fréquemment, reprendre à moindre coût et améliorer le compute effectif de vos GPU.

remarque

serverless_gpu.data.UCVolumeDataset, serverless_gpu.data.DataLoader, serverless_gpu.data.UCVolumeWriter et serverless_gpu.data.UCVolumeReader nécessitent un environnement GPU 5 ou supérieur (API Python GPU serverless 0.5.16 ou supérieure).

Chargez les données efficacement pour minimiser le temps d'inactivité du GPU

Une étape d'entraînement doit chevaucher le compute GPU avec la préparation des données pour l'étape suivante. Sur AI Runtime, tout accès aux données passe par Unity Catalog. Pour les datasets basés sur des fichiers dans les volumes Unity Catalog, utilisez serverless_gpu.data.UCVolumeDataset, qui copie chaque fichier du montage FUSE vers un cache local rapide lors du premier accès et renvoie le chemin local mis en cache.

Associez-le à serverless_gpu.data.DataLoader, une sous-classe prête à l'emploi du DataLoader de PyTorch optimisée pour les E/S GPU Serverless, qui récupère et met en cache les fichiers simultanément pendant que le GPU compute.

Python
import serverless_gpu.data

dataset = serverless_gpu.data.UCVolumeDataset("/Volumes/my-catalog/my-schema/my-volume/data")

loader = serverless_gpu.data.DataLoader(
dataset,
batch_size=64,
)

for batch in loader:
local_paths = batch # open these immediately; see the caching note below
...
attention

Le chemin généré par serverless_gpu.data.UCVolumeDataset est éphémère . Le cache évince les fichiers téléchargés les moins récemment une fois que l'espace disque libre tombe en dessous d'un threshold (default 10 % du système de fichiers du cache, remplaçable avec la variable d'environnement SGC_FSLAYER_MIN_FREE_DISK_BYTES), un chemin peut donc être supprimé dès que vous récupérez l'élément suivant. Ouvrez-le, décodez-le ou copiez-le dans la même itération de boucle. Ne stockez jamais un chemin renvoyé dans une liste ou un dictionnaire pour le rouvrir plus tard.

Décodez les fichiers en enveloppant serverless_gpu.data.UCVolumeDataset dans un second IterableDataset qui consomme le Stream de chemin. Le wrapper reçoit des chemins locaux déjà mis en cache, de sorte que l'analyse ne touche jamais au montage FUSE :

Python
from torch.utils.data import IterableDataset
from PIL import Image
import torchvision.transforms.functional as TF

class ImageDataset(IterableDataset):
"""Decodes each cached file path from UCVolumeDataset into a tensor."""

def __init__(self, path_dataset: serverless_gpu.data.UCVolumeDataset):
self._path_dataset = path_dataset

def __iter__(self):
for local_path in self._path_dataset:
image = Image.open(local_path).convert("RGB")
yield TF.to_tensor(image)


path_dataset = serverless_gpu.data.UCVolumeDataset("/Volumes/my-catalog/my-schema/my-volume/images")
dataset = ImageDataset(path_dataset)
loader = serverless_gpu.data.DataLoader(dataset, batch_size=64)

Deux exigences lors de la montée en charge :

  • Utilisez toujours serverless_gpu.data.DataLoader pour une formation multi-époque. Cela force persistent_workers=True lorsque num_workers > 0, de sorte que le suivi d'éviction du cache en mémoire de chaque worker survit à travers les époques. Le DataLoader PyTorch standard re-forke les Worker à chaque époque par default, ce qui entraîne une fuite du répertoire de cache partagé jusqu'à ce qu'il soit plein.
  • Tous les rangs doivent transmettre le même num_workers. serverless_gpu.data.UCVolumeDataset partitionne les fichiers en utilisant un pas global sur world_size × num_workers emplacements. Des valeurs incohérentes entraînent la duplication ou l'omission de fichiers entre les rangs.

Lorsque torch.distributed est initialisé, serverless_gpu.data.UCVolumeDataset lit le rang au moment de l'itération et partitionne automatiquement les fichiers entre les rangs ; vous n'avez donc pas besoin d'un DistributedSampler pour les données de volume basées sur des fichiers.

Création de point de contrôle avec Distributed Checkpoint (DCP)

Utilisez PyTorch Distributed Checkpoint (DCP) plutôt que torch.save. Chaque rang écrit sa propre partition en parallèle dans un répertoire de points de contrôle, en utilisant la bande passante E/S globale totale et en évitant le pic de mémoire lié au rassemblement de tout l'état sur un seul rang. DCP stocke également les métadonnées globales des tenseurs, de sorte qu'un point de contrôle enregistré sur un certain nombre de GPU peut être repris sur un nombre différent.

Sur AI Runtime, serverless_gpu.data.UCVolumeWriter et serverless_gpu.data.UCVolumeReader sont les backends de stockage DCP. Ils organisent toutes les E/S via un répertoire local rapide (/tmp, pris en charge par NVMe sur les nœuds GPU AIR) et effectuent l'upload ou le download vers ou depuis un volume Unity Catalog, ce qui est plus rapide que l'écriture directe des partitions sur le montage FUSE.

Python
import torch.distributed.checkpoint as dcp
from torch.distributed.checkpoint.state_dict import get_state_dict, set_state_dict
import serverless_gpu.data

checkpoint_path = "/Volumes/my-catalog/my-schema/my-volume/checkpoints/step_1000"

# Save
model_sd, optim_sd = get_state_dict(model, optimizer)
state_dict = {"model": model_sd, "optim": optim_sd, "step": 1000}
dcp.save(state_dict, storage_writer=serverless_gpu.data.UCVolumeWriter(checkpoint_path))

# Load
model_sd, optim_sd = get_state_dict(model, optimizer)
state_dict = {"model": model_sd, "optim": optim_sd}
dcp.load(state_dict, storage_reader=serverless_gpu.data.UCVolumeReader(checkpoint_path))
set_state_dict(
model,
optimizer,
model_state_dict=state_dict["model"],
optim_state_dict=state_dict["optim"],
)

DCP vaut la peine d'être utilisé même pour un entraînement purement parallèle aux données (DDP), où les poids sont répliqués entre les rangs. DCP écrit une copie unique dédoublonnée des poids répliqués tout en capturant l'état unique de chaque rang (position des données et état du RNG, traités ci-dessous), et il s'agit de la même API dont vous aurez besoin si vous passez ultérieurement à FSDP ou au parallélisme tensoriel.

Enregistrer de manière asynchrone

Une sauvegarde synchrone bloque l'entraînement jusqu'à ce que les octets soient durables dans le volume. Pour un point de contrôle volumineux, il s'agit du temps d'inactivité du GPU. dcp.async_save copie l'état dans un tampon de transfert (rapide) puis effectue l'upload en arrière-plan pendant que l'entraînement se poursuit. Comme chaque point de contrôle ne coûte pratiquement aucun temps GPU, vous pouvez vous permettre de créer des points de contrôle beaucoup plus souvent, ce qui limite le travail perdu après une interruption.

Les sauvegardes asynchrones nécessitent un backend CPU sur le groupe de processus ; initialisez-le donc avec gloo et nccl:

Python
import torch.distributed as dist
import torch.distributed.checkpoint as dcp
from torch.distributed.checkpoint.state_dict import get_state_dict
import serverless_gpu.data

dist.init_process_group(backend="cpu:gloo,cuda:nccl")

checkpoint_future = None

def save_async(step, model, optimizer):
global checkpoint_future
# Ensure the previous async save finished before starting a new one.
if checkpoint_future is not None:
checkpoint_future.result()

model_sd, optim_sd = get_state_dict(model, optimizer)
state_dict = {"model": model_sd, "optim": optim_sd, "step": step}
writer = serverless_gpu.data.UCVolumeWriter(f"/Volumes/my-catalog/my-schema/my-volume/checkpoints/step_{step}")
checkpoint_future = dcp.async_save(state_dict, storage_writer=writer)

Récupération automatique à partir du point de contrôle valide le plus récent

Une exécution peut être interrompue pendant la sauvegarde, laissant un répertoire de point de contrôle partiel. serverless_gpu.data.UCVolumeWriter publie le fichier .metadata dans le volume uniquement après la fin de l'upload des fichiers de données shard, la présence de .metadata est donc un signal fiable indiquant qu'une sauvegarde est terminée. Utilisez-le pour sélectionner le point de contrôle valide le plus récent lors du redémarrage.

Python
import os

def find_latest_valid(checkpoint_root):
"""Return the newest checkpoint directory that finished writing, or None."""
candidates = sorted(
(d for d in os.listdir(checkpoint_root) if d.startswith("step_")),
key=lambda d: int(d.split("_")[1]),
reverse=True,
)
for name in candidates:
path = os.path.join(checkpoint_root, name)
if os.path.exists(os.path.join(path, ".metadata")): # save completed
return path
return None # nothing valid; start fresh

Une boucle d'entraînement résiliente sélectionne le dernier point de contrôle valide, effectue une restauration à partir de celui-ci et crée fréquemment des points de contrôle. L'intervalle des points de contrôle limite le travail perdu après une interruption ; ainsi, des sauvegardes asynchrones fréquentes et peu coûteuses maintiennent la recomputation à un niveau faible :

Python
CHECKPOINT_EVERY = 100

latest = find_latest_valid("/Volumes/my-catalog/my-schema/my-volume/checkpoints")
start_step = 0
if latest is not None:
model_sd, optim_sd = get_state_dict(model, optimizer)
state = {"model": model_sd, "optim": optim_sd, "step": 0}
dcp.load(state, storage_reader=serverless_gpu.data.UCVolumeReader(latest))
set_state_dict(
model,
optimizer,
model_state_dict=state["model"],
optim_state_dict=state["optim"],
)
start_step = state["step"]

for step in range(start_step, total_steps):
train_step(...)
if step % CHECKPOINT_EVERY == 0:
save_async(step, model, optimizer) # inexpensive, so run it often

Cette boucle crée des points de contrôle de l'état du modèle et de l'optimiseur. Il ne restaure pas encore la position du pipeline de données, ce qui est traité dans la section suivante.

Créer un point de contrôle du pipeline de données

Un point de contrôle de modèle capture l'état du modèle et de l'optimiseur, mais pas la position de votre pipeline de données au sein du dataset. Supposons que vous restauriez le modèle à l'étape 1 900, mais que le dataloader redémarre depuis le début du dataset. L'exécution reprise ré-entraîne sur des exemples déjà vus lors de cette époque et ignore les exemples proches du point d'interruption, biaisant silencieusement la distribution des données sans erreur.

Pour reprendre sur les bonnes données, suivez votre position dans le dataset dans le cadre de votre propre état d'entraînement et restaurez-la lors de la reprise. Il y a quatre points à prendre en considération :

Suivre un décalage d'échantillon ou de partition

Enregistrez votre progression dans l'époque en utilisant un index d'échantillon global, un nombre de batch ou une liste d'ID de partitions consommés dans le dictionnaire d'état du point de contrôle, et passez à cette position lors de la reprise. Cela permet de garder la position des données sous votre contrôle explicite plutôt que de compter sur un dataloader pour sérialiser son état interne.

Python
# Include the data position in the checkpoint state dict:
state_dict = {
"model": model_sd,
"optim": optim_sd,
"step": step,
"epoch": epoch,
"samples_seen": samples_seen, # your own counter, advanced each batch
}

Pour un dataset de type map avec un échantillonneur déterministe, ignorez les batchs déjà consommés lors de cette époque lors de la reprise. Comme l'ordre de l'échantillonneur est déterministe pour un (seed, epoch) donné (voir Rendre le pipeline déterministe), l'avance rapide reproduit la position exacte :

Python
resume_batch = state["samples_seen"] // batch_size

for epoch in range(start_epoch, num_epochs):
sampler.set_epoch(epoch)
for batch_idx, batch in enumerate(loader):
# On the resumed epoch only, skip batches already processed.
if epoch == start_epoch and batch_idx < resume_batch:
continue
train_step(batch)
samples_seen += batch_size

Pour un dataset en streaming partitionné, suivez plutôt l'ensemble des partitions terminées et ne transmettez à l'exécution reprise que les partitions restantes. Cela évite de rejouer la totalité des batchs d'une époque juste pour atteindre le point d'interruption :

Python
# Filter the shard list down to work not yet done, then build the loader from it.
remaining = [s for s in all_shards if s not in state["completed_shards"]]
dataset = ShardDataset(remaining)

Créer un point de contrôle de l'état interne d'un dataset personnalisé

Si vous écrivez votre propre dataset, donnez-lui des méthodes pour sérialiser et restaurer sa propre position, et intégrez cet état dans le point de contrôle. Cela permet de conserver la logique de reprise à proximité de la logique d'itération, et le dataset sait exactement ce dont il a besoin pour effectuer une avance rapide (partition actuelle, décalage au sein de celle-ci, contenu du tampon de mélange, etc.) plutôt que de laisser la boucle d'entraînement le reconstruire à partir d'un compteur externe.

Python
from torch.utils.data import IterableDataset

class ResumableShardDataset(IterableDataset):
"""A streaming dataset that can checkpoint and restore its own position."""

def __init__(self, shards):
self._shards = shards
self._shard_idx = 0 # position advanced during __iter__
self._offset = 0

def state_dict(self):
return {"shard_idx": self._shard_idx, "offset": self._offset}

def load_state_dict(self, state):
self._shard_idx = state["shard_idx"]
self._offset = state["offset"]

def __iter__(self):
for i in range(self._shard_idx, len(self._shards)):
self._shard_idx = i
for j, example in enumerate(self._read_shard(self._shards[i])):
if j < self._offset:
continue # skip examples already consumed from this shard
self._offset = j + 1
yield example
self._offset = 0
Python
# Save and restore the dataset position with the rest of the checkpoint state.
state_dict["dataset"] = dataset.state_dict()
# On resume:
dataset.load_state_dict(state["dataset"])

Redémarrer à partir d'une limite d'époque

Si le saut en avant n'est pas pratique, créez des points de contrôle uniquement aux limites des époques et reprenez au start de l'époque suivante. L'exécution voit alors chaque exemple exactement une fois par époque malgré l'interruption, au prix d'une perte pouvant aller jusqu'à une époque de progression par défaillance. C'est plus simple lorsque les époques sont courtes par rapport au taux de défaillance.

Rendre le pipeline déterministe

Chacune de ces stratégies ne reprend sur les bonnes données que si le pipeline est reproductible à partir de l'état enregistré. Le mélange et l'augmentation s'appuient sur des RNG ; veillez donc à les initialiser et à conserver leur état dans le point de contrôle. Sinon, l'ordre de mélange et d'augmentation après le redémarrage ne correspond pas à l'ordre précédant l'interruption, et un décalage d'avance rapide pointe vers les mauvais échantillons.

Pour reprendre en milieu de pipeline plutôt qu’uniquement à la limite d’une époque, les RNG pilotant le mélange et l’augmentation doivent eux-mêmes être capables de créer des points de contrôle. La définition d’une graine (seeding) seule reproduit la séquence depuis le start, mais pas le point où vous avez été interrompu. Utilisez des objets RNG dont vous pouvez sérialiser l’état interne, et créez un point de contrôle de cet état afin que chaque RNG reprenne exactement là où il s’est arrêté. Se reposer sur un générateur de nombres aléatoires (RNG) global qui n’est réinitialisé qu’au start de l’époque rejoue les mêmes tirages depuis le début de l’époque, ce qui ne correspond plus à un décalage de saut en milieu d’époque.

Initialiser toutes les sources de hasard :

Python
import random
import numpy as np
import torch

def seed_everything(seed: int):
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)

Enregistrez et restaurez l'état du générateur de nombres aléatoires (RNG) avec le modèle afin que les séquences d'augmentation et de mélange se poursuivent de manière transparente :

Python
# Save
state_dict["rng"] = {
"python": random.getstate(),
"numpy": np.random.get_state(),
"torch": torch.get_rng_state(),
"cuda": torch.cuda.get_rng_state_all(),
}

# Load
rng = state["rng"]
random.setstate(rng["python"])
np.random.set_state(rng["numpy"])
torch.set_rng_state(rng["torch"])
torch.cuda.set_rng_state_all(rng["cuda"])

Si vous utilisez un DistributedSampler au lieu d'un dataloader avec état, appelez sampler.set_epoch(epoch) au start de chaque époque. Comme le mélange est une fonction déterministe de (seed, epoch), la restauration du compteur d'époque reproduit la permutation exacte :

Python
for epoch in range(start_epoch, num_epochs):
sampler.set_epoch(epoch) # deterministic reshuffle per epoch
for batch in loader:
...
remarque

Pour la correction du pipeline de données, vous devez vous assurer que l'ordre des données et le Stream d'augmentation sont reproductibles, ce que permettent le seeding et le checkpointing RNG ci-dessus. Vous n'avez généralement pas besoin de passes avant identiques au niveau du bit ; torch.use_deterministic_algorithms(True) force des kernels déterministes mais peut réduire le throughput et ne couvre pas toutes les opérations.

Pages associées