Nota
O acesso a esta página requer autorização. Pode tentar iniciar sessão ou alterar os diretórios.
O acesso a esta página requer autorização. Pode tentar alterar os diretórios.
À medida que um trabalho escala para mais GPUs, a probabilidade de falhas de hardware e software aumenta. Esta página aborda estratégias para tornar as suas corridas de treino mais rápidas e tolerantes a falhas:
- Carregue os dados de forma eficiente para que as GPUs não fiquem paradas.
- Modelo de checkpoint e estado do otimizador de forma eficiente para volumes do Catálogo Unity.
- Recupera automaticamente após uma interrupção.
- Checkpoint o pipeline de dados para que a retomada continue a treinar com os dados corretos.
Com estes padrões, fazer checkpointing do teu modelo é barato, por isso podes fazer checkpoints frequentemente, retomar barato e melhorar o cálculo efetivo das tuas GPUs.
Note
serverless_gpu.data.UCVolumeDataset, serverless_gpu.data.DataLoader, serverless_gpu.data.UCVolumeWriter, e serverless_gpu.data.UCVolumeReader requerem o ambiente GPU 5 ou superior (API Python da GPU sem servidor 0.5.16 ou superior).
Carregue os dados de forma eficiente para minimizar o tempo ocioso da GPU
Uma etapa de treino deve sobrepor a computação da GPU com a preparação de dados para a etapa seguinte. No AI Runtime, todo o acesso a dados passa pelo Unity Catalog. Para conjuntos de dados baseados em ficheiros em volumes do Unity Catalog, use serverless_gpu.data.UCVolumeDataset, que copia cada ficheiro da montagem FUSE para uma cache local rápida no primeiro acesso e gera o caminho local em cache.
Combine-o com serverless_gpu.data.DataLoader, uma subclasse drop-in do PyTorch DataLoader ajustada para I/O de GPU serverless, que recolhe e armazena ficheiros em cache simultaneamente enquanto a GPU faz o cálculo.
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
O caminho dado por serverless_gpu.data.UCVolumeDataset é efémero. A cache expulsa os ficheiros menos recentemente descarregados assim que o disco livre cai abaixo de um determinado limiar (padrão 10% do sistema de ficheiros da cache, sobrescrevendo com a SGC_FSLAYER_MIN_FREE_DISK_BYTES variável ambiente), pelo que um caminho pode ser eliminado assim que retirar o próximo item. Abrir, decodificar ou copiar na mesma iteração do ciclo. Nunca guarde um caminho devolvido numa lista ou dito para reabrir mais tarde.
Decodifica ficheiros envolvendo serverless_gpu.data.UCVolumeDataset um segundo IterableDataset que consome o fluxo de caminho. O wrapper recebe caminhos locais já armazenados em cache, pelo que a análise nunca toca na montagem 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)
Dois requisitos ao escalar:
- Usa
serverless_gpu.data.DataLoadersempre para treino multi-época. Forçapersistent_workers=Truequandonum_workers > 0, para que o rastreador de cache-expulsão em memória de cada trabalhador sobreviva através das épocas. O PyTorchDataLoaderoriginal re-forka os trabalhadores em todas as épocas por defeito, o que faz com que o diretório de cache partilhado vaze até preencher. - Todas as patentes devem passar da mesma
num_workersforma.serverless_gpu.data.UCVolumeDatasetParticiona ficheiros usando um passo global atravésworld_size × num_workersdos slots. Valores incompatíveis fazem com que os ficheiros sejam duplicados ou saltados entre classificações.
Quando torch.distributed está inicializado, serverless_gpu.data.UCVolumeDataset lê o rank na altura da iteração e particiona os ficheiros automaticamente entre ranks, por isso não precisa DistributedSampler de um para dados de volume baseados em ficheiros.
Ponto de controlo com Ponto de Controlo Distribuído (DCP)
Use o PyTorch Distributed Checkpoint (DCP) em vez de torch.save. Cada rank escreve o seu próprio shard em paralelo num diretório de checkpoint, usando toda a largura de banda agregada de I/O e evitando o pico de memória de reunir todos os estados num só rank. O DCP também armazena metadados tensoriais globais, pelo que um checkpoint guardado num número de GPUs pode ser retomado num número diferente.
No AI Runtime, serverless_gpu.data.UCVolumeWriter e serverless_gpu.data.UCVolumeReader são os backends de armazenamento DCP. Eles fazem todo o I/O através de um diretório local rápido (/tmpcom NVMe apoiado nos nós da GPU AIR) e fazem upload ou download a partir de um volume do Unity Catalog, o que é mais rápido do que escrever shards diretamente na montagem 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"],
)
O DCP vale a pena ser usado mesmo para treino puramente paralelo de dados (DDP), onde os pesos são replicados entre patentes. O DCP escreve uma única cópia desduplicada dos pesos replicados, mantendo ainda assim o estado único de cada rank (posição de dados e estado RNG, abordado abaixo), e é a mesma API que precisa se mais tarde passar para FSDP ou paralelismo tensorial.
Guardar de forma assíncrona
Um save síncrono bloqueia o treino até que os bytes sejam duráveis no volume. Para um checkpoint grande, isso é o tempo de inatividade da GPU.
dcp.async_save Cópias para um buffer de staging (FAST) e depois carregam em segundo plano enquanto o treino continua. Como cada checkpoint custa quase nenhum tempo de GPU, podes dar-te ao luxo de fazer checkpoint muito mais vezes, que é o que os limites perdem trabalho após uma interrupção.
Gravações assíncronas requerem um backend de CPU no grupo de processos, por isso inicialize-o com ambos 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)
Recuperar automaticamente do ponto de controlo válido mais recente
Uma execução pode ser interrompida a meio da gravação, deixando um diretório parcial de checkpoint.
serverless_gpu.data.UCVolumeWriter Publica o .metadata ficheiro no volume apenas depois de os ficheiros de dados do fragmento terminarem de ser carregados, por isso a presença de .metadata é um sinal fiável de que um save foi concluído. Usa-o para selecionar o ponto de controlo válido mais recente ao 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
Um ciclo de treino resiliente seleciona o ponto de controlo válido mais recente, restaura a partir dele e os pontos de controlo frequentemente. Os limites do intervalo de checkpoint perdiam trabalho após uma interrupção, por isso os saves assíncronos frequentes e baratos mantêm a recomputação pequena:
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 modelo de checkpoints de ciclo e estado do otimizador. Ainda não restaura a posição do pipeline de dados, que é abordada na secção seguinte.
Checkpoint: o pipeline de dados
Um checkpoint de modelo capta o estado do modelo e do otimizador, mas não a posição do seu pipeline de dados dentro do conjunto de dados. Suponha que restaura o modelo no passo 1.900, mas o dataloader reinicia do início do conjunto de dados. A corrida retomada é retreinada em exemplares já vistos nesta época e salta exemplos perto do ponto de interrupção, polarizando silenciosamente a distribuição de dados sem erro.
Para retomar com os dados corretos, acompanha a tua posição no conjunto de dados como parte do teu próprio estado de treino e restaura-a no currículo. Há quatro pontos a considerar:
Rastrear um deslocamento de amostra ou fragmento
Regista até onde avançaste na época usando um índice global de amostras, contagem de lotes ou lista de IDs de fragmentos consumidos no estado do checkpoint, e avança para essa posição quando retomares. Isto mantém a posição dos dados sob o seu controlo explícito, em vez de depender de um dataloader para serializar o seu 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 um conjunto de dados em estilo mapa com um amostrador determinístico, ignore os lotes já consumidos nesta época ao retomar. Como a ordem do sampler é determinística para um dado (seed, epoch) ( ver Tornar o pipeline determinístico), o avanço rápido reproduz a posição exata:
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 um conjunto de dados em shards e streaming, regista o conjunto de fragmentos completos e entrega à execução retomada apenas os fragmentos que restam. Isto evita repetir lotes de uma época inteira só para chegar ao ponto de interrupção:
# 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 um conjunto de dados personalizado
Se escreveres o teu próprio conjunto de dados, dá-lhe métodos para serializar e restaurar a sua própria posição, e integra esse estado no checkpoint. Isto mantém a lógica de retomada ao lado da lógica de iteração e o conjunto de dados sabe exatamente o que precisa de avançar rapidamente (fragmento atual, offset dentro dele, conteúdo do buffer de embaralhamento, e assim sucessivamente) em vez de o ciclo de treino reconstruí-lo a partir de um 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 a partir de uma fronteira de época
Se o skip-ahead for impraticável, faça checkpoint apenas nos limites da época e retome no início da próxima época. A corrida passa então a ver cada exemplo exatamente uma vez por época durante a interrupção, ao custo de perder até uma época de progresso por falha. Isto é mais simples quando as épocas são curtas em relação à taxa de falha.
Tornar o pipeline determinístico
Qualquer uma das estratégias só retoma com os dados corretos se o pipeline for reproduzível a partir do estado guardado. O baralho e a aumentação consomem RNGs, por isso sementem-nos e levam o estado deles no checkpoint. Caso contrário, a ordem de baralhamento e aumento após o reinício não corresponde à ordem anterior à interrupção, e um deslocamento de salto aponta para as amostras erradas.
Para retomar a meio de um pipeline em vez de apenas numa fronteira de época, os RNGs que impulsionam o baralhamento e aumento têm de ser, por si só, checkpoints. A sementação sozinha reproduz a sequência desde o início, mas não o ponto em que foste interrompido. Usa objetos RNG cujo estado interno possas serializar, e faz checkpoint nesse estado para que cada RNG continue exatamente de onde ficou. Confiar num RNG global que só é re-seed no início da época repete os mesmos draws do início da época, o que já não corresponde a um deslocamento skip-ahead de meio da época.
Seme todas as fontes de aleatoriedade:
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)
Guarde e restaure o estado RNG juntamente com o modelo para que as sequências de aumento e embaralhamento continuem sem interrupções:
# 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 usares um DistributedSampler dataloader em vez de um dataloader com estado, liga sampler.set_epoch(epoch) no início de cada época. Como o embaralhamento é uma função determinística de (seed, epoch), restaurar o contador de época reproduz a permutação exata:
for epoch in range(start_epoch, num_epochs):
sampler.set_epoch(epoch) # deterministic reshuffle per epoch
for batch in loader:
...
Note
Para a correção do pipeline de dados, é necessário que a ordem dos dados e o fluxo de aumento sejam reproduzíveis, o que o checkpoint de seed e RNG acima fornecem. Geralmente, não precisas de passes para a frente idênticos a bits; torch.use_deterministic_algorithms(True) força kernels determinísticos mas pode reduzir o débito e não cobre todas as operações.
Páginas relacionadas
- Carregar dados no AI Runtime: carregar dados tabulares e não estruturados através do Unity Catalog.
- Rastreio e observabilidade de experiências: rastreio de experiências MLflow, visualização de registos e monitorização de recursos da GPU.
-
Formação distribuída em cadernos: o
@distributeddecorador e a formação multi-GPU.