Pular para o conteúdo principal

Melhore o desempenho e a resiliência do treinamento no AI Runtime

info

Visualização

Esse recurso está em Pré-lançamento público.

À medida que um Job escala para mais GPUs, a probabilidade de falha de hardware e software aumenta. Esta página aborda estratégias para tornar suas execuções de treinamento mais rápidas e mais tolerantes a falhas:

Com esses padrões, criar pontos de verificação do seu modelo é barato, portanto, você pode criar pontos de verificação com frequência, retomar de forma econômica e melhorar o compute efetivo das suas GPUs.

nota

serverless_gpu.data.UCVolumeDataset, serverless_gpu.data.DataLoader, serverless_gpu.data.UCVolumeWriter e serverless_gpu.data.UCVolumeReader exigem ambiente de GPU 5 ou acima (API Python de GPU Serverless 0.5.16 ou acima).

Carregue dados de forma eficiente para minimizar o tempo parado da GPU

Um passo de treinamento deve sobrepor o compute de GPU com a preparação de dados para o próximo o passo. No AI Runtime, todo acesso a dados passa pelo Unity Catalog. Para datasets baseados em arquivos em volumes do Unity Catalog, use serverless_gpu.data.UCVolumeDataset, que copia cada arquivo da montagem FUSE para um cache local rápido no primeiro acesso e fornece o caminho local em cache.

Combine-o com serverless_gpu.data.DataLoader, uma subclasse drop-in do DataLoader do PyTorch ajustada para I/O de GPU serverless que busca e armazena arquivos em cache simultaneamente enquanto a GPU realiza o compute.

Python
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
...
atenção

O caminho gerado por serverless_gpu.data.UCVolumeDataset é efêmero . O cache remove os arquivos download menos recentemente assim que o disco livre cai abaixo de um limite (default de 10% do sistema de arquivos de cache, substituível com a variável de ambiente SGC_FSLAYER_MIN_FREE_DISK_BYTES), portanto, um caminho pode ser excluído assim que você extrair o próximo item. Abra, decodifique ou copie-o na mesma iteração de loop. Nunca armazene um caminho retornado em uma lista ou dicionário para reabri-lo mais tarde.

Decodifique arquivos envolvendo serverless_gpu.data.UCVolumeDataset em um segundo IterableDataset que consome a transmissão de caminhos. O wrapper recebe caminhos locais já armazenados em cache, portanto, a análise nunca acessa a montagem FUSE:

Python
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 realizar o dimensionamento (scale-out):

  • Sempre use serverless_gpu.data.DataLoader para treinamento de múltiplas épocas. Isso força persistent_workers=True quando num_workers > 0, para que o rastreador de eliminação de cache em memória de cada worker sobreviva entre as épocas. O PyTorch DataLoader padrão cria novos workers a cada época por default, o que causa vazamento no diretório de cache compartilhado até que ele fique cheio.
  • Todos os ranks devem passar o mesmo num_workers. serverless_gpu.data.UCVolumeDataset particiona arquivos usando um passo global em world_size × num_workers slots. Valores incompatíveis fazem com que arquivos sejam duplicados ou ignorados entre os ranks.

Quando torch.distributed é inicializado, serverless_gpu.data.UCVolumeDataset lê o rank no momento da iteração e particiona os arquivos entre os ranks automaticamente, portanto, você não precisa de um DistributedSampler para dados de volume baseados em arquivo.

Ponto de verificação com Distributed Checkpoint (DCP)

Use o Distributed Checkpoint (DCP) do PyTorch em vez de torch.save. Cada rank grava seu próprio fragmento em paralelo em um diretório de ponto de verificação, usando a largura de banda total de E/S agregada e evitando o pico de memória de reunir todo o estado em um único rank. O DCP também armazena metadados globais de tensores, para que um checkpoint salvo em um determinado número de GPUs possa ser retomado em um número diferente.

No AI Runtime, serverless_gpu.data.UCVolumeWriter e serverless_gpu.data.UCVolumeReader são os back-ends de armazenamento DCP. Eles organizam toda a E/S por meio de um diretório local rápido (/tmp, com suporte a NVMe em nós de GPU AIR) e fazem upload ou download de um volume do Unity Catalog, o que é mais rápido do que gravar fragmentos diretamente na montagem FUSE.

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

Vale a pena usar o DCP mesmo para treinamento puramente paralelo de dados (DDP), onde os pesos são replicados entre os ranks. O DCP grava uma única cópia desduplicada dos pesos replicados enquanto ainda captura o estado exclusivo de cada rank (posição dos dados e estado do RNG, abordados abaixo), e é a mesma API de que você precisará se decidir migrar posteriormente para FSDP ou paralelismo de tensor.

Salvar de forma assíncrona

Um salvamento síncrono bloqueia o treinamento até que os bytes estejam duráveis no volume. Para um ponto de verificação grande, isso é tempo parado da GPU. dcp.async_save copia o estado para um buffer de preparo (rápido) e, em seguida, faz o upload em segundo plano enquanto o treinamento continua. Como cada ponto de verificação quase não custa tempo de GPU, você pode criar pontos de verificação com muito mais frequência, o que limita o trabalho perdido após uma interrupção.

Salvamentos assíncronos exigem um backend de CPU no grupo de processos, portanto, inicialize-o com gloo e nccl:

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

Recupere automaticamente a partir do ponto de verificação válido mais recente

Uma execução pode ser interrompida durante o salvamento, deixando um diretório de ponto de verificação parcial. serverless_gpu.data.UCVolumeWriter publica o arquivo .metadata no volume somente após o upload dos arquivos de dados do fragmento ser concluído, portanto, a presença de .metadata é um sinal confiável de que um salvamento foi concluído. Use-o para selecionar o ponto de verificação válido mais recente na reinicialização.

Python
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 loop de treinamento resiliente seleciona o ponto de verificação válido mais recente, restaura a partir dele e cria pontos de verificação frequentemente. O intervalo do ponto de verificação limita o trabalho perdido após uma interrupção, portanto, salvamentos assíncronos frequentes e baratos mantêm a recomputação pequena:

Python
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 loop cria pontos de verificação do estado do modelo e do otimizador. Ele ainda não restaura a posição do pipeline de dados, o que é abordado na próxima seção.

Fazer checkpoint do pipeline de dados

Um ponto de verificação de modelo captura o estado do modelo e do otimizador, mas não a posição do seu pipeline de dados dentro do dataset. Suponha que você restaure o modelo no passo 1.900, mas o dataloader reinicie do início do dataset. A execução retomada treina novamente com exemplos já vistos nesta época e ignora exemplos próximos ao ponto de interrupção, enviesando silenciosamente a distribuição de dados sem erro.

Para retomar nos dados corretos, acompanhe sua posição no dataset como parte do seu próprio estado de treinamento e restaure-a ao retomar. Há quatro pontos a considerar:

Rastrear um deslocamento de amostra ou fragmento

Registre o progresso na época usando um índice global de amostras, contagem de lotes ou lista de IDs de fragmentos consumidos no dicionário de estado do ponto de verificação e pule para essa posição ao retomar. Isso mantém a posição dos dados sob seu controle explícito, em vez de depender de um dataloader para serializar seu estado interno.

Python
# 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 dataset estilo mapa com um amostrador determinístico, ignore os lotes já consumidos nesta época ao retomar. Como a ordem do amostrador é determinística para um determinado (seed, epoch) (consulte Tornar o pipeline determinístico), o avanço rápido reproduz a posição exata:

Python
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 dataset fragmentado e de transmissão, rastreie o conjunto de fragmentos concluídos e forneça à execução retomada apenas os fragmentos restantes. Isso evita repetir uma quantidade de lotes de uma época inteira apenas para chegar ao ponto de interrupção:

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

Criar ponto de verificação do estado interno de um dataset personalizado

Se você criar seu próprio dataset, forneça métodos para serializar e restaurar sua própria posição, e incorpore esse estado ao checkpoint. Isso mantém a lógica de retomada próxima à lógica de iteração e o dataset sabe exatamente o que precisa para avançar rapidamente (fragmento atual, deslocamento dentro dele, conteúdo do buffer de embaralhamento, etc.) em vez de o loop de treinamento reconstruí-lo a partir de um contador externo.

Python
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
Python
# 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 um limite de época

Se o avanço rápido for impraticável, crie pontos de verificação apenas nos limites das épocas e retome no início da próxima época. A execução então vê cada exemplo exatamente uma vez por época durante a interrupção, ao custo de perder até uma época de progresso por falha. Isso é mais simples quando as épocas são curtas em relação à taxa de falha.

Torne o pipeline determinístico

Qualquer uma das estratégias só retoma os dados corretos se o pipeline for reproduzível a partir do estado salvo. O embaralhamento e a aumentação utilizam RNGs, portanto, defina suas sementes e mantenha seus estados no ponto de verificação. Caso contrário, a ordem de embaralhamento e aumentação após a reinicialização não corresponde à ordem anterior à interrupção, e um deslocamento de salto aponta para as amostras incorretas.

Para retomar no meio de um pipeline em vez de apenas em um limite de época, os RNGs que controlam o embaralhamento e a aumentação devem, eles próprios, ser passíveis de checkpoint. Apenas definir a seed reproduz a sequência desde o começar, mas não o ponto em que você foi interrompido. Use objetos RNG cujo estado interno você possa serializar e crie um ponto de verificação desse estado para que cada RNG continue exatamente de onde parou. Depender de um RNG global que só é semeado novamente ao começar a época reproduz os mesmos sorteios desde o início da época, o que não corresponde mais a um deslocamento de salto no meio da época.

Defina a seed de todas as fontes de aleatoriedade:

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

Salve e restaure o estado do RNG junto com o modelo para que as sequências de aumento e embaralhamento continuem perfeitamente:

Python
# 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 você usar um DistributedSampler em vez de um dataloader com estado, chame 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 épocas reproduz a permutação exata:

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

Para a correção do pipeline de dados, você precisa que a ordem dos dados e a transmissão de aumento sejam reproduzíveis, o que o seeding e o checkpointing de RNG acima fornecem. Geralmente, você não precisa de passagens diretas idênticas bit a bit; torch.use_deterministic_algorithms(True) força kernels determinísticos, mas pode reduzir o throughput e não cobre todas as operações.

Páginas relacionadas