Migliora le prestazioni e la resilienza dell'allenamento su Runtime AI

Importante

Questa funzionalità è in Anteprima Pubblica.

Man mano che un lavoro si espande verso più GPU, aumenta la probabilità di guasti hardware e software. Questa pagina tratta strategie per rendere le tue corse di allenamento più veloci e più resistenti ai guasti:

Con questi pattern, il checkpoint del tuo modello è economico, quindi puoi fare checkpoint frequentemente, riprendere a basso costo e migliorare il calcolo efficace delle tue GPU.

Note

serverless_gpu.data.UCVolumeDataset, serverless_gpu.data.DataLoader, serverless_gpu.data.UCVolumeWriter, e serverless_gpu.data.UCVolumeReader richiedono l'ambiente GPU 5 o superiore (API Python GPU serverless 0.5.16 o superiore).

Carica i dati in modo efficiente per minimizzare il tempo inattivo della GPU

Un passaggio di addestramento dovrebbe sovrapporsi al calcolo GPU con la preparazione dei dati per il passaggio successivo. Su AI Runtime, tutto l'accesso ai dati passa attraverso Unity Catalog. Per i dataset basati su file nei volumi del Catalogo Unity, si usa serverless_gpu.data.UCVolumeDataset, che copia ogni file dal supporto FUSE a una cache locale veloce al primo accesso e fornisce il percorso locale memorizzato.

Abbinalo a serverless_gpu.data.DataLoader, una sottoclasse drop-in del PyTorch DataLoader ottimizzata per l'I/O GPU serverless che recupera e memorizza file contemporaneamente mentre la GPU calcola.

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

Il percorso lasciato da serverless_gpu.data.UCVolumeDataset è effimero. La cache elimina i file scaricati meno di recente una volta che il disco libero scende sotto una certa soglia (predefinito 10% del file system della cache, sovrascrivibile con la SGC_FSLAYER_MIN_FREE_DISK_BYTES variabile ambiente), quindi un percorso può essere cancellato non appena estrai l'elemento successivo. Aprilo, decodificalo o copialo nella stessa iterazione del ciclo. Mai memorizzare un percorso restituito in una lista o in un dettatore da riaprire in seguito.

Decodifica i file avvolgendo serverless_gpu.data.UCVolumeDataset un secondo IterableDataset che consuma il flusso di percorso. Il wrapper riceve percorsi locali già memorizzati nella cache, quindi il parsing non tocca mai il montaggio 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)

Due requisiti quando si scala la scala:

  • Da usare serverless_gpu.data.DataLoader sempre per l'addestramento multi-epoca. Forza persistent_workers=True quando num_workers > 0, così il tracker di sfratti della cache in memoria di ogni lavoratore sopravvive attraverso le epoche. Il PyTorch DataLoader stock ri-forka i lavoratori in ogni epoca di default, il che fa trapelare la directory della cache condivisa finché non si riempie.
  • Tutti i gradi devono superare lo stesso num_workers. serverless_gpu.data.UCVolumeDataset Partiziona i file usando un passo globale attraverso world_size × num_workers gli slot. Valori non corrispondenti causano duplicazioni o salti tra i file di fascia.

Quando torch.distributed viene inizializzato, serverless_gpu.data.UCVolumeDataset legge il rank al momento dell'iterazione e suddivide automaticamente i file tra i ranghi, quindi non serve un DistributedSampler volume dati basati su file.

Checkpoint con Checkpoint Distribuito (DCP)

Usa il PyTorch Distributed Checkpoint (DCP) invece di torch.save. Ogni rango scrive il proprio shard in parallelo in una directory checkpoint, utilizzando tutta la larghezza di banda aggregata di I/O ed evitando il picco di memoria dovuto a raccogliere tutti gli stati in un solo rango. DCP memorizza anche metadati tensori globali, così un checkpoint salvato su un numero di GPU può essere ripreso su un numero diverso.

Su AI Runtime, serverless_gpu.data.UCVolumeWriter e serverless_gpu.data.UCVolumeReader sono i backend di archiviazione DCP. Inviano tutto l'I/O tramite una directory locale veloce (/tmpsupportata da NVMe sui nodi GPU AIR) e caricano o scaricano da un volume del Catalogo Unity, il che è più veloce che scrivere gli shard direttamente sul supporto 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"],
)

Il DCP vale la pena di essere utilizzato anche per l'addestramento puramente data-parallelo (DDP), dove i pesi vengono replicati tra i ranghi. DCP scrive una singola copia deduplicata dei pesi replicati pur catturando lo stato unico di ogni rango (posizione dati e stato RNG, trattato sotto), ed è la stessa API necessaria se successivamente passi a FSDP o parallelismo tensoriale.

Salva in modo asincrono

Un salvataggio sincrono blocca l'allenamento finché i byte non sono durevoli nel volume. Per un checkpoint grande, quello è il tempo di inattività della GPU. dcp.async_save copie lo stato in un buffer di staging (fast) e poi caricato in background mentre l'addestramento continua. Poiché ogni checkpoint costa quasi nessun tempo GPU, puoi permetterti di fare checkpoint molto più spesso, ed è proprio per questo che Bounds ha perso lavoro dopo un'interruzione.

I salvataggi asincroni richiedono un backend CPU sul gruppo di processi, quindi inizializzalo con entrambi gloo e 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)

Recupera automaticamente dall'ultimo checkpoint valido

Una run può essere interrotta a metà salvataggio, lasciando una directory parziale di checkpoint. serverless_gpu.data.UCVolumeWriter pubblica il .metadata file nel volume solo dopo che i file di dati dello shard hanno terminato di caricarsi, quindi la presenza di .metadata è un segnale affidabile che un salvataggio è stato completato. Usalo per selezionare il checkpoint valido più recente al riavvio.

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

Un ciclo di addestramento resiliente seleziona l'ultimo checkpoint valido, ripristina da esso e i checkpoint frequentemente. I limiti dell'intervallo dei checkpoint hanno perso lavoro dopo un'interruzione, quindi i frequenti e economici salvataggi asincroni mantengono il calcolo piccolo:

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

Questo modello di checkpoint a loop e stato dell'ottimizzatore. Non ripristina ancora la posizione della pipeline dati, che sarà trattata nella sezione successiva.

Checkpoint della pipeline dati

Un checkpoint del modello cattura lo stato del modello e dell'ottimizzatore, ma non la posizione della tua pipeline dati all'interno del dataset. Supponiamo di ripristinare il modello al passo 1.900 ma il dataloader si riavvia dall'inizio del dataset. La corsa ripresa si riprende su esemplari già visti in questa epoca e salta esempi vicino al punto di interruzione, polarizzando silenziosamente la distribuzione dei dati senza errore.

Per riprendere con i dati corretti, traccia la tua posizione nel dataset come parte del tuo stato di addestramento e ripristinala nel curriculum. Ci sono quattro punti da considerare:

Traccia un campione o un offset di frammento

Registra fino a che punto sei avanzato nell'epoca usando un indice globale di campioni, un conteggio dei lotti o una lista degli ID degli shard consumati nel codice di stato del checkpoint, e salta a quella posizione quando riprendi. Questo mantiene la posizione dei dati sotto il tuo controllo esplicito invece di affidarsi a un dataloader per serializzare il proprio stato interno.

# 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
}

Per un dataset in stile mappa con un campionatore deterministico, si saltano i lotti già consumati in questa epoca al momento della ripresa. Poiché l'ordine del campionatore è deterministico per un dato (seed, epoch) ( vedi Make the pipeline deterministica), il fast-forwarding riproduce la posizione esatta:

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

Per un dataset sharded e in streaming, traccia invece il set di shard completati e lascia alla ripresa solo gli shard rimasti. Questo evita di rigiocare un'intera epoca di lotti solo per raggiungere il punto di interruzione:

# 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)

Checkpoint: lo stato interno di un dataset personalizzato

Se scrivi il tuo dataset, dagli metodi per serializzare e ripristinare la propria posizione, e integra quello stato nel checkpoint. Questo mantiene la logica del riprendimento accanto a quella delle iterazioni e il dataset sa esattamente cosa deve accelerare (shard corrente, offset al suo interno, contenuti del buffer di mescolare, ecc.) invece che il ciclo di addestramento debba ricostruirlo da un contatore esterno.

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"])

Ricominciare da un confine epocale

Se il skip-ahead non è pratico, fai checkpoint solo ai confini dell'epoca e riprendi all'inizio della successiva epoca. La run vede quindi ogni esempio esattamente una volta per epoca attraverso l'interruzione, al costo di perdere fino a un'epoca di progresso per ogni fallimento. Questo è più semplice quando le epoche sono brevi rispetto al tasso di guasto.

Rendere la pipeline deterministica

Entrambe le strategie riprendono solo sui dati corretti se la pipeline è riproducibile dallo stato salvato. Il rimescolamento e l'aumento attingono dagli RNG, quindi li seminano e portano il loro stato nel checkpoint. Altrimenti l'ordine di mescolamento e aumento dopo il riavvio non corrisponde a quello precedente all'interruzione, e uno spostamento salto-avanti indica i campioni sbagliati.

Per riprendere a metà di una pipeline e non solo a un confine epocale, gli RNG che guidano rimescolamenti e aumenti devono essere essi stessi controllabili. Il seeding da solo riproduce la sequenza dall'inizio, ma non il punto in cui sei stato interrotto. Usa oggetti RNG il cui stato interno puoi serializzare e fai checkpoint su quello stato così ogni RNG continua esattamente da dove aveva lasciato. Affidarsi a un RNG globale che viene riassegnato solo all'inizio dell'epoca, rigioca gli stessi draw dell'inizio dell'epoca, che non corrispondono più a un offset skip-ahead di metà epoca.

Semina tutte le fonti di casualità:

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)

Salva e ripristina lo stato RNG insieme al modello affinché le sequenze di aumento e mescolare continuino senza soluzione di continuità:

# 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"])

Se usi un DistributedSampler dataloader invece di uno con stato, chiama sampler.set_epoch(epoch) all'inizio di ogni epoca. Poiché il mescolamento è una funzione deterministica di (seed, epoch), ripristinare il contatore epoche riproduce la permutazione esatta:

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

Note

Per la correttezza della pipeline dati è necessario che l'ordine dei dati e il flusso di aumento siano riproducibili, cosa che forniscono il seeding e il checkpoint RNG sopra. In generale non servono passaggi in avanti identici a bit; torch.use_deterministic_algorithms(True) Forza kernel deterministici ma può ridurre la capacità di produzione e non copre ogni operazione.