Verbeter de trainingsprestaties en veerkracht op AI Runtime

Important

Deze functie bevindt zich in openbare preview-versie.

Naarmate een taak schaalt naar meer GPU's, neemt de kans op hardware- en softwarestoringen toe. Deze pagina behandelt strategieën om je trainingsruns sneller en toleranter te maken:

Met deze patronen is checkpointen van je model goedkoop, zodat je regelmatig checkpoint kunt gebruiken, goedkoop kunt hervatten en de effectieve rekenkracht van je GPU's kunt verbeteren.

Note

serverless_gpu.data.UCVolumeDataset, serverless_gpu.data.DataLoader, , serverless_gpu.data.UCVolumeWriteren serverless_gpu.data.UCVolumeReader vereisen GPU-omgeving 5 of hoger (Serverless GPU Python API 0.5.16 of hoger).

Laad data efficiënt om de inactieve GPU-tijd te minimaliseren

Een trainingsstap moet GPU-berekening overlappen met datavoorbereiding voor de volgende stap. Op AI Runtime verloopt alle data-toegang via Unity Catalog. Voor bestandsgebaseerde datasets in Unity Catalog-volumes gebruik serverless_gpu.data.UCVolumeDatasetje , waarbij elk bestand van de FUSE-mount wordt gekopieerd naar een snelle lokale cache bij eerste toegang en zo het gecachte lokale pad wordt opgeleverd.

Combineer het met serverless_gpu.data.DataLoader, een drop-in subklasse van de PyTorch DataLoader die is afgestemd op serverless GPU I/O, die bestanden gelijktijdig ophaalt en cachet terwijl de GPU rekent.

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

Het pad dat door serverless_gpu.data.UCVolumeDataset wordt gegeven is vluchtig. De cache verwijdert minst recent gedownloade bestanden zodra de vrije schijf onder een drempel zakt (standaard 10% van het cachebestandssysteem, overschrijfbaar met de SGC_FSLAYER_MIN_FREE_DISK_BYTES omgevingsvariabele), dus een pad kan worden verwijderd zodra je het volgende item optrekt. Open, decoder, of kopieer het in dezelfde lus-iteratie. Sla nooit een teruggestuurd pad op in een lijst of dict om later opnieuw te openen.

Decodeer bestanden door in een seconde serverless_gpu.data.UCVolumeDataset te wrappen IterableDataset die de padstroom verbruikt. De wrapper ontvangt reeds gecachede lokale paden, dus parsing raakt de FUSE-mount nooit:

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)

Twee vereisten bij het opschalen:

  • Gebruik het altijd serverless_gpu.data.DataLoader voor training over meerdere epochen. Het dwingt persistent_workers=True wanneer num_workers > 0, zodat de cache-eviction tracker van elke werknemer in het geheugen overleeft door tijdperken heen. De standaard PyTorch DataLoader forkt workers standaard elke epoche opnieuw, waardoor de gedeelde cachemap lekt totdat deze vol is.
  • Alle rangen moeten hetzelfde num_workersslagen . serverless_gpu.data.UCVolumeDataset Partitioneert bestanden met een globale stride over world_size × num_workers slots. Niet-overeenkomende waarden zorgen ervoor dat bestanden worden gedupliceerd of overgeslagen tussen rangen.

Wanneer torch.distributed geïnitialiseerd is, serverless_gpu.data.UCVolumeDataset leest het de rang tijdens iteratietijd en partitioneert bestanden automatisch tussen rangen, dus je hebt geen bestand nodig DistributedSampler voor bestandsgebaseerde volumegegevens.

Checkpoint met Gedistribueerd Checkpoint (DCP)

Gebruik PyTorch Distributed Checkpoint (DCP) in plaats van torch.save. Elke rang schrijft zijn eigen shard parallel in een checkpointdirectory, waarbij de volledige totale I/O-bandbreedte wordt gebruikt en de geheugenpiek wordt voorkomen van het verzamelen van alle toestanden in één rang. DCP slaat ook globale tensormetadata op, zodat een checkpoint dat op één aantal GPU's is opgeslagen, op een ander aantal kan worden hervat.

Op AI Runtime serverless_gpu.data.UCVolumeWriter zijn de serverless_gpu.data.UCVolumeReader DCP-opslagbackends. Ze plaatsen alle I/O via een snelle lokale directory (/tmpNVMe-ondersteund op AIR GPU-nodes) en uploaden naar of downloaden van een Unity Catalog-volume, wat sneller is dan shards direct naar de FUSE-mount schrijven.

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 is het gebruik waard, zelfs voor pure data-parallel (DDP) training, waarbij de gewichten over verschillende rangen worden gerepliceerd. DCP schrijft een enkele gededupliceerde kopie van de gerepliceerde gewichten, terwijl het toch de unieke status van elke rang vastlegt (datapositie en RNG-status, hieronder besproken), en het is dezelfde API die je nodig hebt als je later overstapt op FSDP of tensorparallelisme.

Asynchroon opslaan

Een synchrone save blokkeert training totdat de bytes in het volume duurzaam zijn. Voor een groot checkpoint is dat de idle-tijd van de GPU. dcp.async_save kopieert de status naar een stagingbuffer (snel) en uploadt vervolgens in de achtergrond terwijl de training doorgaat. Omdat elk checkpoint bijna geen GPU-tijd kost, kun je checkpoint veel vaker checkpointen, wat het verlies van werk na een onderbreking beperkt is.

Asynchrone saves vereisen een CPU-backend op de procesgroep, dus initialiseer deze met zowel gloo als 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)

Herstel automatisch van het meest recente geldige checkpoint

Een run kan midden in het opslaan worden onderbroken, waardoor er een gedeeltelijke checkpointmap overblijft. serverless_gpu.data.UCVolumeWriter publiceert het .metadata bestand pas op het volume nadat de shard-databestanden zijn geüpload, zodat de aanwezigheid van .metadata een betrouwbaar signaal is dat een save is voltooid. Gebruik het om bij het herstarten het meest recente geldige checkpoint te selecteren.

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

Een veerkrachtige trainingslus selecteert het meest recente geldige checkpoint, herstelt daarvan en checkpoints regelmatig. Het checkpointinterval verliest werk na een onderbreking, dus frequente, goedkope asynchrone saves houden de herberekening klein:

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

Deze lus checkpoints modeleren de status van de optimizer. De positie van de datapijplijn wordt nog niet hersteld, die in de volgende sectie wordt behandeld.

Controleer de datapijplijn

Een modelcheckpoint legt de status van model en optimizer vast, maar niet de positie van je datapipeline binnen de dataset. Stel dat je het model herstelt bij stap 1.900, maar de dataloader opnieuw opstart vanaf het begin van de dataset. De hervatte run hertraint op voorbeelden die deze periode al zijn gezien en slaat voorbeelden nabij het onderbrekingspunt over, waardoor de dataverdeling stilletjes wordt verstoord zonder fouten.

Om te hervatten met de juiste data, volg je je positie in de dataset als onderdeel van je eigen trainingsstaat en herstel je die op je cv. Er zijn vier punten om rekening mee te houden:

Volg een sample of shard-offset

Noteer hoe ver je in het tijdperk bent gekomen door een globale steekproefindex, batchtelling of lijst van verbruikte shard-ID's in het checkpoint-toestanddict te gebruiken, en sla die positie over wanneer je weer opgaat. Dit houdt de datapositie expliciet onder jouw controle in plaats van te vertrouwen op een dataloader om de interne status te serialiseren.

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

Voor een map-achtige dataset met een deterministische sampler, sla bij het hervatten van de reeds verbruikte batches over. Omdat de volgorde van de samplers deterministisch is voor een gegeven (seed, epoch) (zie Make the pipeline deterministisch), reproduceert fastforwarding de exacte positie:

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

Voor een gesharded, streaming dataset volg je in plaats daarvan de set voltooide shards en geef je de hervatte run alleen de shards die overblijven. Dit voorkomt dat je een hele periode aan batches opnieuw hoeft te spelen om het onderbrekingspunt te bereiken:

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

Controleer de interne status van een aangepaste dataset

Als je je eigen dataset schrijft, geef het dan methoden om zijn eigen positie te serialiseren en te herstellen, en vouw die toestand in het checkpoint. Dit houdt de hervatingslogica naast de iteratielogica en de dataset weet precies wat hij moet doorspoelen (huidige shard, offset erin, shuffle buffer-inhoud, enzovoort) in plaats van dat de trainingslus het reconstrueert vanuit een externe teller.

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

Begin vanaf een tijdperkgrens

Als skip-ahead onpraktisch is, controleer dan alleen bij epochgrenzen en hervat aan het begin van de volgende epoch. De run ziet vervolgens elk voorbeeld precies één keer per epoch over de onderbreking, met als prijs dat je tot één epoch aan voortgang per mislukking verliest. Dit is het eenvoudigst wanneer tijdperken kort zijn ten opzichte van het faalpercentage.

Maak de pijplijn deterministisch

Beide strategieën worden alleen voortgezet op de juiste data als de pijplijn reproduceerbaar is vanuit de opgeslagen toestand. Shuffling en augmentatie putten uit RNG's, dus zaad ze en neem hun staat mee in het checkpoint. Anders komt de volgorde van schudden en augmentatie na herstart niet overeen met de volgorde vóór de onderbreking, en wijst een skip-ahead offset naar de verkeerde samples.

Om halverwege een pijplijn te hervatten in plaats van alleen op een epochgrens, moeten de RNG's die schudden en augmentatie aansturen zelf checkpoint-baar zijn. Alleen seeding reproduceert de sequentie vanaf het begin, maar niet het punt waar je werd onderbroken. Gebruik RNG-objecten waarvan je de interne staat kunt serialiseren, en checkpoint die toestand zodat elke RNG precies verdergaat waar hij gebleven was. Vertrouwen op een globale RNG die pas bij het begin van het tijdperk opnieuw wordt geïmplanteerd, worden dezelfde trekkingen van het begin van het tijdperk opnieuw afgespeeld, wat niet langer overeenkomt met een midden-epoch-skip-ahead offset.

Saai alle bronnen van willekeur:

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)

Sla de RNG-status op en herstel samen met het model, zodat augmentatie- en schudsequenties naadloos verlopen:

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

Als je in DistributedSampler plaats van een stateful dataloader gebruikt, roep sampler.set_epoch(epoch) dan aan het begin van elk epoch. Omdat de schud een deterministische functie van (seed, epoch)is, reproduceert het herstellen van de epoch-teller de exacte permutatie:

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

Note

Voor de correctheid van de datapijplijn moet de datavolgorde en augmentatiestroom reproduceerbaar zijn, wat de seeding en RNG-checkpoints hierboven bieden. Je hebt over het algemeen geen bitwise-identieke voorwaartse passes nodig; torch.use_deterministic_algorithms(True) dwingt deterministische kernels, maar kan de doorvoer verminderen en dekt niet elke bewerking.