Nota:
El acceso a esta página requiere autorización. Puede intentar iniciar sesión o cambiar directorios.
El acceso a esta página requiere autorización. Puede intentar cambiar los directorios.
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:
- Carga los datos de forma eficiente para que las GPUs no estén inactivas.
- El estado del modelo de checkpoint y del optimizador de forma eficiente a volúmenes del Catálogo de Unity.
- Recupera automáticamente tras una interrupción.
- Checkpoint en la tubería de datos para que la ejecución reanudada continúe entrenando con los datos correctos.
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.DataLoaderpara entrenamiento multiépoca. Fuerzapersistent_workers=Truecuandonum_workers > 0, así que el rastreador de desalojo de caché en memoria de cada trabajador sobrevive a través de épocas. El PyTorchDataLoaderde 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.UCVolumeDatasetParticiona los archivos usando un paso global a travésworld_size × num_workersde 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.
Páginas relacionadas
- Cargar datos en AI Runtime: cargar datos tabulares y no estructurados a través del Unity Catalog.
- Seguimiento y observabilidad de experimentos: seguimiento de experimentos en MLflow, visualización de registros y monitorización de recursos GPU.
-
Formación distribuida en cuadernos: el
@distributeddecorador y la formación multi-GPU.