Améliorer les performances et la résilience de l'entraînement sur AI Runtime
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 :
- Chargez les données efficacement afin que les GPU ne soient pas inactifs.
- Enregistrez efficacement l'état du modèle et de l'optimiseur dans les volumes Unity Catalog.
- Récupérez automatiquement après une interruption.
- Créer un point de contrôle du pipeline de données afin qu'une exécution reprise continue l'entraînement sur les bonnes données.
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.
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.
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
...
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 :
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.DataLoaderpour une formation multi-époque. Cela forcepersistent_workers=Truelorsquenum_workers > 0, de sorte que le suivi d'éviction du cache en mémoire de chaque worker survit à travers les époques. LeDataLoaderPyTorch 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.UCVolumeDatasetpartitionne les fichiers en utilisant un pas global surworld_size × num_workersemplacements. 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.
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:
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.
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 :
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.
# 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 :
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 :
# 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.
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
# 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 :
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 :
# 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 :
for epoch in range(start_epoch, num_epochs):
sampler.set_epoch(epoch) # deterministic reshuffle per epoch
for batch in loader:
...
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
- Charger des données sur AI Runtime: chargement de données tabulaires et non structurées via Unity Catalog.
- Suivi des expérimentations et observabilité: suivi des expérimentations MLflow, visualisation des Logs et monitoring des ressources GPU.
- Formation distribuée dans les notebooks: le décorateur
@distributedet la formation multi-GPU.