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

Important

Cette fonctionnalité est disponible en préversion publique.

À mesure qu’un poste s’étend à plus de GPU, la probabilité de défaillance matérielle et logicielle augmente. Cette page présente des stratégies pour rendre vos séances d’entraînement plus rapides et plus tolérantes aux fautes :

Avec ces modèles, le contrôle ponctuel de votre modèle est peu coûteux, vous pouvez donc faire ces contrôles fréquemment, reprendre à moindre coût et améliorer la capacité de calcul effective de vos GPU.

Note

serverless_gpu.data.UCVolumeDataset, serverless_gpu.data.DataLoader, serverless_gpu.data.UCVolumeWriter, et serverless_gpu.data.UCVolumeReader nécessitent l’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 recouper le calcul GPU avec la préparation des données pour l’étape suivante. Sur AI Runtime, tout l’accès aux données passe par Unity Catalog. Pour les ensembles de données basés sur des fichiers dans les volumes du catalogue Unity, utilisez serverless_gpu.data.UCVolumeDataset, qui copie chaque fichier du montage FUSE vers un cache local rapide dès le premier accès et donne le chemin local mis en cache.

Associez-le à serverless_gpu.data.DataLoader, une sous-classe drop-in du PyTorch DataLoader réglée pour l’E/S GPU serverless qui récupère et met en cache les fichiers simultanément pendant que le GPU calcule.

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
    ...

Warning

Le chemin emprunté serverless_gpu.data.UCVolumeDataset est éphémère. Le cache élimine les fichiers les moins récemment téléchargés dès que le disque libre descend sous un seuil (par défaut 10% du système de fichiers cache, écrasable par la SGC_FSLAYER_MIN_FREE_DISK_BYTES variable environnement), donc un chemin peut être supprimé dès que vous retirez 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 retourné dans une liste ou un dict pour le rouvrir plus tard.

Décodez les fichiers en encapsulant serverless_gpu.data.UCVolumeDataset dans un second IterableDataset qui consomme le flux de chemins d’accès. L’enveloppe reçoit des chemins locaux déjà mis en cache, donc 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 pour le scale-out :

  • Utilisez toujours serverless_gpu.data.DataLoader pour l’apprentissage multi-époque. Cela force persistent_workers=True quand num_workers > 0, donc le suivi d’éviction du cache en mémoire de chaque collaborateur survit à travers les époques. La version standard de DataLoader de PyTorch duplique à nouveau les collaborateurs à chaque époque par défaut, ce qui entraîne une fuite dans le répertoire de cache partagé jusqu’à saturation.
  • Tous les rangs doivent passer le même num_workers. serverless_gpu.data.UCVolumeDataset partitionne les fichiers à l’aide d’un pas global sur world_size × num_workers emplacements. Des valeurs inadaptées provoquent la duplication ou le saut des fichiers entre rangs.

Lorsque torch.distributed est initialisé, serverless_gpu.data.UCVolumeDataset lit le rang au moment de l’itération et répartit automatiquement les fichiers entre les rangs, vous n’avez donc pas besoin de DistributedSampler pour les données volumiques basées sur des fichiers.

Point de contrôle avec point de contrôle distribué (DCP)

Utilisez le point de contrôle distribué (DCP) de PyTorch plutôt que torch.save. Chaque rang écrit en parallèle sa propre partition dans un répertoire de points de contrôle, en exploitant toute la bande passante d’E/S agrégée et en évitant le pic de mémoire causé par la collecte de tous les états vers un seul rang. DCP stocke également des métadonnées tensorielles globales, de sorte qu’un point de contrôle sauvegardé sur un nombre de GPU peut être repris sur un autre nombre.

Sur AI Runtime, serverless_gpu.data.UCVolumeWriter et serverless_gpu.data.UCVolumeReader ce sont les backends de stockage DCP. Ils font transiter toutes les E/S par un répertoire local rapide (/tmp, sur stockage NVMe sur les nœuds GPU AIR), puis chargent ou téléchargent les données dans un volume Unity Catalog, ce qui est plus rapide que d’écrire directement les fragments 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"],
)

La DCP mérite d’être utilisée même pour l’apprentissage pur de données en parallèle (DDP), où les poids sont répliqués sur tous les rangs. DCP écrit une seule copie dédupliquée des poids répliqués tout en capturant l’état unique de chaque rang (position des données et état RNG, abordé ci-dessous), et c’est la même API dont vous avez besoin si vous passez plus tard à FSDP ou au parallélisme tensoriel.

Enregistrer de manière asynchrone

Une sauvegarde synchrone bloque l’apprentissage jusqu’à ce que les octets aient été écrits de façon durable sur le volume. Pour un gros point de contrôle, c’est le temps d’inactivité du GPU. dcp.async_save copie l’état dans un tampon intermédiaire (rapide), puis le charge en arrière-plan pendant que l’apprentissage se poursuit. Comme chaque point de contrôle ne mobilise presque aucun temps de calcul GPU, vous pouvez vous permettre d’effectuer des points de contrôle bien 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érez automatiquement depuis le dernier point de contrôle valide

Une exécution peut être interrompue en cours de sauvegarde, laissant un répertoire partiel de points de contrôle. serverless_gpu.data.UCVolumeWriter Publie le .metadata fichier dans le volume seulement après que les fichiers de données fragmentaires aient terminé leur chargement, donc la présence de .metadata est un signal fiable qu’une sauvegarde a été terminée. Utilisez-le pour sélectionner le dernier point de contrôle valide au 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, restaure à partir de celui-ci, puis effectue fréquemment des points de contrôle. L’intervalle entre les points de contrôle limite le travail perdu après une interruption, de sorte que des sauvegardes asynchrones fréquentes et peu coûteuses réduisent au minimum le recalcul :

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 enregistre l’état du modèle et de l’optimiseur. Il ne rétablit pas encore la position du pipeline de données, qui est abordée dans la section suivante.

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

Un point de contrôle modèle capture l’état du modèle et de l’optimiseur, mais pas la position de votre pipeline de données dans le jeu de données. Supposons que vous restauriez le modèle à l’étape 1 900 mais que le dataloader redémarre depuis le début du jeu de données. L’exécution reprise entraîne à nouveau le modèle sur des exemples déjà vus au cours de cette époque et ignore des exemples proches du point d’interruption, ce qui biaise silencieusement la distribution des données sans générer d’erreur.

Pour reprendre avec les bonnes données, enregistrez votre position dans le jeu de données dans votre propre état d’apprentissage et restaurez-la lors de la reprise. Il y a quatre points à considérer :

Suivi du décalage d’un échantillon ou d’une partition

Notez jusqu’où vous êtes avancé dans l’époque en utilisant un index global d’échantillons, un nombre de lots ou une liste des identifiants de partitions consommées dans le dictionnaire d’état du point de contrôle, puis passez directement à cette position lorsque vous reprenez. Cela maintient la position des données sous votre contrôle explicite plutôt que de dépendre d’un dataloader pour sérialiser leur é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 jeu de données de type carte avec un échantillonneur déterministe, ignorez les lots déjà consommés dans cette époque lors de la reprise. Parce que l’ordre de l’échantillonneur est déterministe pour un élément donné (seed, epoch) ( voir Rendre le pipeline déterministe), l’avancement 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 jeu de données partitionné et diffusé en continu, suivez plutôt le jeu des partitions déjà traitées et ne transmettez à l’exécution reprise que les partitions restantes. Cela évite de rejouer les lots d’une époque entière 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 jeu de données personnalisé

Si vous écrivez votre propre jeu de données, donnez-lui des méthodes pour sérialiser et restaurer sa propre position, puis intégrez cet état dans le point de contrôle. Cela maintient la logique de reprise à côté de la logique d’itération, et le jeu de données sait exactement ce qu’il doit faire avancer rapidement (partition actuelle, décalage au sein de celle-ci, contenu du tampon de lecture aléatoire, etc.), plutôt que la boucle d’apprentissage doive les reconstituer à 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 l’avance rapide est impraticable, n’effectuez des points de contrôle qu’à la fin de chaque époque et reprenez au début de l’époque suivante. L’exécution voit alors chaque exemple exactement une fois par époque à travers l’interruption, au risque de perdre jusqu’à une époque de progrès par échec. C’est le plus simple lorsque les époques sont courtes par rapport au taux de défaillance.

Rendez le pipeline déterministe

L’une ou l’autre stratégie ne reprend à partir des bonnes données que si le pipeline est reproductible à partir d’un état sauvegardé. La lecture aléatoire et les augmentations s’appuient sur des RNG ; donc initialisez-les et enregistrez leur état dans le point de contrôle. Sinon, l’ordre de lecture aléatoire et des augmentations après le redémarrage ne correspond pas à l’ordre avant l’interruption, et un décalage d’avance rapide pointe vers des échantillons incorrects.

Pour reprendre à mi-parcours d’un pipeline plutôt qu’à une limite d’époque, les RNG qui entraînent la lecture aléatoire et les augmentations doivent eux-mêmes être contrôlables ponctuellement. L’initialisation seule reproduit la séquence depuis le début, mais pas à partir du point où elle a été interrompue. Utilisez des objets RNG dont vous pouvez sérialiser l’état interne, et faites un checkpoint pour que chaque RNG continue exactement là où il s’était arrêté. S’appuyer sur un RNG global qui n’est réinitialisé qu’au début de l’époque rejoue les mêmes augmentations à partir du début de l’époque, ce qui ne correspond plus à un décalage d’avance rapide au milieu d’une époque.

Semez toutes les sources d’aléa :

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 RNG en même temps que le modèle afin que les séquences d’augmentation et de lecture aléatoire se poursuivent sans interruption :

# 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 chargeur de données avec état, appelez sampler.set_epoch(epoch) au début de chaque époque. Parce que la lecture aléatoire est une fonction déterministe de (seed, epoch), la restauration du compteur d’époque reproduit exactement la même permutation :

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

Note

Pour l’exactitude du pipeline de données, il faut que l’ordre des données et le flux d’augmentation soient reproductibles, ce que fournissent les initialisations et les points de contrôle du RNG ci-dessus. En général, vous n’avez pas besoin de passes avant identiques bit par bit ; torch.use_deterministic_algorithms(True) force des noyaux déterministes mais peut réduire le débit et ne couvre pas toutes les opérations.