Melhorar o desempenho e a resiliência do treinamento em tempo de execução de IA

Importante

Esse recurso está em Visualização Pública.

À medida que um trabalho escala para mais GPUs, a probabilidade de falha de hardware e software aumenta. Esta página aborda estratégias para tornar seus treinos mais rápidos e mais tolerantes a falhas:

Com esses padrões, checkpointing do seu modelo é barato, então você pode checkpoint frequentemente, retomar barato e melhorar o processamento efetivo das suas GPUs.

Note

serverless_gpu.data.UCVolumeDataset, serverless_gpu.data.DataLoader, serverless_gpu.data.UCVolumeWriter, e serverless_gpu.data.UCVolumeReader requerem ambiente de 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 treinamento deve sobrepor a computação da GPU com a preparação de dados para a próxima etapa. No AI Runtime, todo acesso a dados passa pelo Unity Catalog. Para conjuntos de dados 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 gera o caminho local cacheado.

Pare-o com serverless_gpu.data.DataLoader, uma subclasse drop-in do PyTorch DataLoader ajustada para I/O de GPU serverless, que busca e armazena arquivos 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. O cache elimina os arquivos menos recentemente baixados assim que o disco livre cai abaixo de um limite (padrão 10% do sistema de arquivos de cache, sobrescrevendo com a SGC_FSLAYER_MIN_FREE_DISK_BYTES variável ambiente), então um caminho pode ser deletado assim que você puxar o próximo item. Abra, decodifica ou copia na mesma iteração do ciclo. Nunca armazene um caminho retornado em uma lista ou dito para reabrir depois.

Decodifica arquivos envolvendo serverless_gpu.data.UCVolumeDataset um segundo IterableDataset que consome o fluxo de caminho. O wrapper recebe caminhos locais já cacheados, então a análise nunca toca a 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:

  • Sempre use serverless_gpu.data.DataLoader para treinamento multi-época. Ele força persistent_workers=True quando num_workers > 0, para que o rastreador de cache-eviction em memória de cada trabalhador sobreviva através das épocas. O PyTorch DataLoader padrão re-forka os trabalhadores a cada época por padrão, o que vaza o diretório de cache compartilhado até que ele preenchia.
  • Todas as patentes devem passar da mesma num_workersforma. serverless_gpu.data.UCVolumeDataset Particiona arquivos usando um passo global entre world_size × num_workers slots. Valores incompatíveis fazem com que arquivos sejam duplicados ou pulados entre as patentes.

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

Ponto de controle com Ponto de Controle Distribuído (DCP)

Use o PyTorch Distributed Checkpoint (DCP) em vez de torch.save. Cada rank escreve seu próprio shard em paralelo em um diretório de checkpoint, usando toda a largura de banda agregada de E/S e evitando o pico de memória de reunir todos os estados em um único rank. O DCP também armazena metadados globais de tensorial, então um checkpoint salvo em um número de GPUs pode ser retomado em outro número.

No AI Runtime, serverless_gpu.data.UCVolumeWriter e serverless_gpu.data.UCVolumeReader são os backends de armazenamento DCP. Eles escalam toda a E/S por um diretório local rápido (/tmprespaldado por NVMe nos nós da GPU AIR) e fazem upload ou download de um volume do Unity Catalog, o que é mais rápido do que gravar shards diretamente no suporte 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 treinamento puramente paralelo de dados (DDP), onde os pesos são replicados entre patentes. O DCP grava uma única cópia desduplicada dos pesos replicados enquanto ainda captura o estado único de cada rank (posição dos dados e estado RNG, abordado abaixo), e é a mesma API que você precisa se depois migrar para FSDP ou paralelismo tensorial.

Salvar de forma assíncrona

Um save síncrono bloqueia o treinamento até que os bytes fiquem duráveis no volume. Para um checkpoint grande, isso é o tempo ocioso da GPU. dcp.async_save Cópias para um buffer de staging (FAST) e depois fazem upload em segundo plano enquanto o treinamento continua. Como cada checkpoint custa quase nenhum tempo de GPU, você pode se dar ao luxo de fazer checkpoint com muito mais frequência, que é o que a Bounds perde trabalho após uma interrupção.

Saves assíncronos exigem um backend de CPU no grupo de processos, então inicialize-o com 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)

Recupere automaticamente do ponto de controle válido mais recente

Uma execução pode ser interrompida no meio do save, deixando um diretório parcial de checkpoint. serverless_gpu.data.UCVolumeWriter Publica o .metadata arquivo no volume somente depois que os arquivos de dados do fragmento terminam de ser enviados, então a presença de .metadata é um sinal confiável de que um save foi concluído. Use para selecionar o checkpoint 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 treinamento resiliente seleciona o checkpoint válido mais recente, restaura a partir dele e os checkpoints frequentemente. Os limites do intervalo de checkpoint perdiam trabalho após uma interrupção, então saves assíncronos frequentes e baratos mantêm o recálculo pequeno:

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 modela checkpoints e o estado do otimizador. Ainda não restaura a posição do pipeline de dados, que será abordada na próxima seção.

Checkpoint: o pipeline de dados

Um checkpoint de modelo captura 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 você restaure o modelo no passo 1.900, mas o dataloader reinicie do início do conjunto de dados. A corrida retomada reinicia os treinos em exemplares já vistos nessa época e pula exemplos próximos ao ponto de interrupção, polarizando silenciosamente a distribuição dos dados sem erro.

Para retomar com os dados corretos, acompanhe sua posição no conjunto de dados como parte do seu próprio estado de treinamento e restaure no currículo. Há quatro pontos a considerar:

Rastreie um deslocamento de amostra ou fragmento

Registre até onde você avançou na época usando um índice global de amostra, contagem de lotes ou lista de IDs de fragmentos consumidos no dito do estado do checkpoint, e avance para essa posição quando 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.

# 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 no estilo mapa com um sampler determinístico, ignore os lotes já consumidos nesta época ao retomar. Como a ordem do sampler é determinística para um dado (seed, epoch) ( veja Make the pipeline deterministic), 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 com fragmentos e streaming, acompanhe o conjunto de fragmentos completos e entregue à execução retomada apenas os fragmentos que restarem. Isso evita rejogar 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 o estado interno de um conjunto de dados personalizado

Se você escrever seu próprio conjunto de dados, dê a ele métodos para serializar e restaurar sua própria posição, e dobre esse estado no checkpoint. Isso mantém a lógica de retomar ao lado da lógica de iteração e o conjunto de dados sabe exatamente o que precisa avançar rapidamente (fragmento atual, offset dentro dele, conteúdo do buffer de embaralhamento, e assim por diante), em vez do ciclo de treinamento reconstruindo 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 pular adiante for impraticável, faça checkpoint apenas nos limites da época e retome no início da próxima época. A execução então mostra 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 nos dados corretos se o pipeline for reproduzível a partir do estado salvo. Embaralhar e aumentar tiram RNGs, então seed eles e carregam o estado deles no checkpoint. Caso contrário, a ordem de embaralhamento e aumento após o reinício não corresponde à ordem anterior à interrupção, e um deslocamento de pular aponta para as amostras erradas.

Para retomar no meio de um pipeline, e não apenas em uma fronteira de época, os RNGs que impulsionam embaralhamento e aumento devem ser checkpoint. A semente sozinha reproduz a sequência do início, mas não o ponto em que você foi interrompido. Use objetos RNG cujo estado interno você possa serializar, e faça checkpoint nesse estado para que cada RNG continue exatamente de onde parou. Confiar em um RNG global que só é re-seed no início da época rejoga os mesmos draws do início da época, que não corresponde mais a um offset skip-ahead da época média.

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)

Salve e restaure o estado RNG junto com o modelo para que as sequências de aumento e embaralhamento continuem sem 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"])

Se você usar um DistributedSampler dataloader em vez de um dataloader com estado, ligue 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, você precisa que a ordem dos dados e o fluxo de aumento sejam reproduzíveis, o que o seeding e o checkpoint RNG acima fornecem. Geralmente, você não precisa de passes diretos idênticos a bits; torch.use_deterministic_algorithms(True) força kernels determinísticos, mas pode reduzir o throughput e não cobre todas as operações.