Mejorar el rendimiento y la resiliencia del entrenamiento en tiempo de ejecución con IA

Importante

Esta característica está en versión preliminar pública.

A medida que un trabajo escala a más GPUs, aumenta la probabilidad de fallos de hardware y software. Esta página cubre estrategias para que tus entrenamientos sean más rápidos y tolerantes a fallos:

Con estos patrones, hacer checkpointing de tu modelo es barato, así que puedes hacer checkpoint con frecuencia, reanudar barato y mejorar el cálculo efectivo de tus GPUs.

Nota

serverless_gpu.data.UCVolumeDataset, serverless_gpu.data.DataLoader, serverless_gpu.data.UCVolumeWriter, y serverless_gpu.data.UCVolumeReader requieren un entorno GPU 5 o superior (API Python de GPU sin servidor 0.5.16 o superior).

Carga los datos de forma eficiente para minimizar el tiempo de inactividad de la GPU

Un paso de entrenamiento debe solaparse entre el cálculo de la GPU y la preparación de datos para el siguiente paso. En AI Runtime, todo el acceso a los datos pasa por Unity Catalog. Para conjuntos de datos basados en archivos en volúmenes del Catálogo de Unity, usa serverless_gpu.data.UCVolumeDataset, que copia cada archivo desde el montaje FUSE a una caché local rápida en el primer acceso y proporciona la ruta local almacenada en caché.

Combínalo con serverless_gpu.data.DataLoader, una subclase de entrada directa del PyTorch DataLoader ajustada para E/S de GPU serverless que recupera y almacena archivos en caché simultáneamente mientras la GPU calcula.

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

El camino que serverless_gpu.data.UCVolumeDataset demuestra es efímero. La caché expulsa los archivos menos descargados una vez que el disco libre baja de un umbral (por defecto 10% del sistema de archivos de caché, sobrescribible con la SGC_FSLAYER_MIN_FREE_DISK_BYTES variable de entorno), por lo que una ruta puede eliminarse en cuanto extraes el siguiente elemento. Ábrelo, decodificalo o cópialo en la misma iteración de bucle. Nunca guardes una ruta retornada en una lista o dictado para reabrirla más tarde.

Decodifica archivos envolviendo serverless_gpu.data.UCVolumeDataset un segundo IterableDataset que consume el flujo de ruta. El envoltorio recibe rutas locales ya almacenadas en caché, por lo que el análisis nunca toca el montaje 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)

Dos requisitos al escalar:

  • Siempre úsalo serverless_gpu.data.DataLoader para entrenamiento multiépoca. Fuerza persistent_workers=True cuando num_workers > 0, así que el rastreador de desalojo de caché en memoria de cada trabajador sobrevive a través de épocas. El PyTorch DataLoader de serie vuelve a bifurcar a los trabajadores en cada época por defecto, lo que filtra el directorio de caché compartido hasta que se llena.
  • Todos los rangos deben pasar igual num_workers. serverless_gpu.data.UCVolumeDataset Particiona los archivos usando un paso global a través world_size × num_workers de las ranuras. Los valores desajustados hacen que los archivos se dupliquen o se salten entre rangos.

Cuando torch.distributed se inicializa, serverless_gpu.data.UCVolumeDataset lee el rango en tiempo de iteración y divide los archivos automáticamente entre rangos, así que no necesitas un DistributedSampler para datos de volumen basados en archivos.

Punto de control con Punto de Control Distribuido (DCP)

Utiliza el punto de control distribuido (DCP) de PyTorch en lugar de torch.save. Cada rango escribe su propio fragmento en paralelo en un directorio de puntos de control, usando todo el ancho de banda agregado de E/S y evitando el pico de memoria de reunir todos los estados en un solo rango. DCP también almacena metadatos globales de tensores, por lo que un punto de control guardado en un número de GPUs puede retomarse en otro número.

En AI Runtime, serverless_gpu.data.UCVolumeWriter y serverless_gpu.data.UCVolumeReader son los backends de almacenamiento DCP. Ellos pasan toda la E/S a través de un directorio local rápido (/tmprespaldado por NVMe en nodos de la GPU AIR) y suben o descargan desde un volumen del Catálogo de Unity, lo cual es más rápido que escribir fragmentos directamente en el soporte 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 merece la pena incluso para entrenamiento puramente paralelo de datos (DDP), donde los pesos se replican entre rangos. DCP escribe una única copia deduplicada de los pesos replicados mientras captura el estado único de cada rango (posición de datos y estado RNG, explicado más abajo), y es la misma API que necesitas si más adelante pasas a FSDP o paralelismo tensorial.

Guardar de forma asíncrona

Una partida de guardado síncrona bloquea el entrenamiento hasta que los bytes son duraderos en el volumen. Para un punto de control grande, eso es el tiempo de reposo de la GPU. dcp.async_save Copia el estado en un búfer de etapas (rápido) y luego se sube en segundo plano mientras continúa el entrenamiento. Como cada punto de control cuesta casi nada de tiempo de GPU, puedes permitirte hacer puntos de control mucho más a menudo, que es lo que hace que Bounds pierda trabajo tras una interrupción.

Los guardados asíncronos requieren un backend de CPU en el grupo de procesos, así que inicialízalo con ambosgloo: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)

Recupérate automáticamente desde el último punto de control válido

Una ejecución puede ser interrumpida a mitad de guardado, dejando un directorio parcial de puntos de control. serverless_gpu.data.UCVolumeWriter publica el .metadata archivo en el volumen solo después de que los archivos de datos del fragmento terminen de subirse, por lo que la presencia de .metadata es una señal fiable de que se ha completado un guardado. Úsala para seleccionar el punto de control válido más reciente al reiniciar.

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 bucle de entrenamiento resiliente selecciona el último punto de control válido, restaura desde él y los puntos de control con frecuencia. Los límites de intervalos de puntos de control perdieron trabajo tras una interrupción, por lo que las partidas asíncronas frecuentes y baratas mantienen el recálculo pequeño:

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

Este ciclo de puntos de control y estado del optimizador. Aún no restaura la posición de la tubería de datos, que se cubre en la siguiente sección.

Checkpoint: la tubería de datos

Un punto de control de modelo captura el estado del modelo y del optimizador, pero no la posición de tu pipeline de datos dentro del conjunto de datos. Supongamos que restauras el modelo en el paso 1.900 pero el cargador de datos se reinicia desde el inicio del conjunto de datos. La reanudación de la carrera se reinicia en los ejemplares ya vistos en esta época y salta ejemplos cerca del punto de interrupción, sesgando silenciosamente la distribución de datos sin error.

Para retomar con los datos correctos, registra tu posición en el conjunto de datos como parte de tu propio estado de entrenamiento y regíralo en el currículum. Hay cuatro puntos a considerar:

Rastrea un desfasamiento de muestra o fragmento

Registra hasta dónde has avanzado en la época usando un índice global de muestras, recuento de lotes o lista de IDs de fragmentos consumidos en el dictado de estado del punto de control, y salta a esa posición cuando reanudes. Esto mantiene la posición de los datos bajo tu control explícito en lugar de depender de un cargador de datos para serializar su estado 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
}

Para un conjunto de datos tipo mapa con un muestreador determinista, salta los lotes ya consumidos en esta época al reanudar. Debido a que el orden del muestreador es determinista para un dado (seed, epoch) ( véase Haz la tubería determinista), el avance rápido reproduce la posición exacta:

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

Para un conjunto de datos en shards y streaming, haz un seguimiento del conjunto de fragmentos completados y entrega a la ejecución reanudada solo los fragmentos que quedan. Esto evita tener que repetir lotes de toda una época solo para llegar al punto de interrupción:

# 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 Estado interno de un conjunto de datos personalizado

Si escribes tu propio conjunto de datos, dale métodos para serializar y restaurar su propia posición, y integra ese estado en el punto de control. Esto mantiene la lógica de recurrir junto a la lógica de iteración y el conjunto de datos sabe exactamente qué necesita avanzar rápidamente (fragmento actual, desplazamiento dentro de él, contenido del búfer de barajado, etc.) en lugar de que el bucle de entrenamiento lo reconstruya desde un contador externo.

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

Reiniciar desde un límite de época

Si el salto es poco práctico, haz un punto de control solo en los límites de la época y reanuda al inicio de la siguiente época. La partida entonces muestra cada ejemplo exactamente una vez por época a lo largo de la interrupción, a costa de perder hasta una época de progreso por fallo. Esto es más sencillo cuando las épocas son cortas en relación con la tasa de fallo.

Haz que la tubería sea determinista

Cualquiera de las dos estrategias solo se reanuda con los datos correctos si la tubería es reproducible desde el estado guardado. Barajar y aumentar obtienen cartas de RNG, así que siembras y lleva su estado en el punto de control. De lo contrario, el orden de barajado y aumento tras el reinicio no coincide con el orden previo a la interrupción, y un desplazamiento de salto apunta a las muestras incorrectas.

Para reanudar a mitad de una cadena en lugar de solo en un límite de época, los RNGs que impulsan la baraja y la aumentación deben ser controlables por puntos de control. Sembrar por sí solo reproduce la secuencia desde el principio, pero no el momento en el que te interrumpieron. Usa objetos RNG cuyo estado interno puedas serializar, y pon ese estado de control para que cada RNG continúe exactamente donde lo dejó. Confiar en un RNG global que solo se re-sembra al inicio de la época se repite las mismas cartas que al inicio de la época, lo que ya no coincide con un desplazamiento de salto de anticipación a mitad de época.

Sembra todas las fuentes de aleatoriedad:

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)

Guarda y restaura el estado RNG junto al modelo para que las secuencias de aumento y barajado continúen sin problemas:

# 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 usas un DistributedSampler cargador de datos en lugar de uno con estado, llama sampler.set_epoch(epoch) al inicio de cada época. Debido a que el barajado es una función determinista de (seed, epoch), restaurar el contador de época reproduce la permutación exacta:

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

Nota

Para la corrección de la pipeline de datos necesitas que el orden de datos y el flujo de aumento sean reproducibles, lo que proporcionan los puntos de siembra y el RNG anteriores. Generalmente no necesitas pases directos idénticos a bit a bit; torch.use_deterministic_algorithms(True) fuerza núcleos deterministas pero puede reducir el rendimiento y no cubre todas las operaciones.