Gradient Checkpointing
Saiba como o checkpointing de gradientes reduz o uso de memória GPU ao recalcular ativações durante a retropropagação, com exemplos em PyTorch, vantagens e desvantagens e orientações práticas de treinamento.
O checkpointing de gradientes é uma técnica de treino que poupa memória: guarda apenas ativações intermédias selecionadas da passagem direta e volta a calcular as restantes durante a retropropagação. Também chamado checkpointing de ativações, troca computação adicional por um menor consumo máximo de memória. Apesar do nome, esta técnica cria checkpoints de ativações, não de gradientes dos parâmetros nem de ficheiros do modelo, o que a torna especialmente útil quando os tensores de ativação impedem que uma rede neuronal caiba na memória da GPU disponível.
Como funciona o checkpointing de gradientes#
Durante uma passagem direta padrão, uma rede neuronal calcula tensores intermédios chamados ativações. O sistema de diferenciação automática mantém as ativações necessárias para calcular os gradientes mais tarde, conforme descrito na mecânica do autograd do PyTorch. Redes profundas, lotes grandes, imagens de alta resolução e sequências de entrada longas podem fazer com que estes tensores guardados consumam muita memória.
O checkpointing de gradientes divide a rede em segmentos:
- A passagem direta guarda as entradas ou as ativações de fronteira dos segmentos selecionados.
- As restantes ativações intermédias dentro desses segmentos são descartadas.
- Durante a passagem inversa, cada segmento com checkpoint volta a executar a passagem direta para reconstruir os valores em falta.
- As ativações reconstruídas são usadas imediatamente para calcular os gradientes.
A API de checkpointing de ativações do PyTorch disponibiliza este comportamento através de torch.utils.checkpoint. Há conceitos equivalentes no checkpointing de gradientes e na rematerialização do JAX e no checkpointing de tapes do TensorFlow.
A colocação dos checkpoints determina a compensação. Em geral, usar checkpoints em mais regiões poupa mais memória, mas repete mais operações. Uma colocação seletiva junto de blocos que usam muitas ativações pode proporcionar um melhor equilíbrio do que recalcular a rede inteira.
Compensações e técnicas relacionadas#
O checkpointing de gradientes normalmente não altera os parâmetros do modelo, os gradientes nem os estados do otimizador. O seu principal objetivo é reduzir a memória usada pelas ativações, e o treino torna-se mais lento porque alguns cálculos da passagem direta são executados duas vezes. O resultado exato depende da arquitetura, dos limites dos checkpoints, da forma do lote e do hardware.
Difere de várias técnicas relacionadas:
- A acumulação de gradientes processa vários microlotes antes de atualizar os pesos do modelo, criando um lote efetivo maior sem carregar todas as amostras em simultâneo. O checkpointing de gradientes, por sua vez, reduz as ativações mantidas para cada microlote.
- A precisão mista usa tipos de dados de menor precisão em operações selecionadas. O fluxo de trabalho de precisão mista automática do PyTorch pode reduzir o consumo de memória e acelerar operações compatíveis, enquanto o checkpointing acrescenta deliberadamente computação.
- A redução do tamanho do lote reduz o consumo de memória ao processar menos amostras de cada vez. O checkpointing pode permitir que se use um lote maior ou uma resolução de entrada superior.
- Os checkpoints de treino guardam os pesos e o estado do otimizador para recuperação ou inferência posterior. O guia de checkpoints de modelos do PyTorch descreve este mecanismo de persistência, que não está relacionado com o recálculo de ativações.
O código com checkpoints deve comportar-se de forma consistente nas execuções originais e recalculadas da passagem direta. Um estado mutável, transferências entre dispositivos ou aleatoriedade não controlada dentro de uma região com checkpoint podem causar erros ou gradientes incorretos. Em transformers do tipo descodificador, também pode ser necessário desativar uma cache de chaves e valores durante o treino com checkpoints, pois o estado em cache da inferência pode entrar em conflito com a reconstrução do grafo da passagem direta.
Exemplo com PyTorch#
O exemplo seguinte cria um checkpoint para um bloco que consome muita memória durante um passo de treino:
import torch
from torch import nn
from torch.utils.checkpoint import checkpoint
torch.manual_seed(0)
block = nn.Sequential(
nn.Linear(1024, 4096),
nn.ReLU(),
nn.Linear(4096, 1024),
)
optimizer = torch.optim.AdamW(block.parameters())
inputs = torch.randn(8, 1024, requires_grad=True)
targets = torch.zeros_like(inputs)
optimizer.zero_grad(set_to_none=True)
outputs = checkpoint(block, inputs, use_reentrant=False)
loss = nn.functional.mse_loss(outputs, targets)
loss.backward()
optimizer.step()Só são mantidas as entradas do bloco e as informações de fronteira necessárias; as ativações internas são reconstruídas durante loss.backward(). Em projetos reais, compara as execuções com e sem checkpointing usando uma medição como o pico de memória de GPU alocada do PyTorch.
Aplicações no mundo real#
-
Visão computacional de alta resolução: A segmentação médica, a deteção aérea e a inspeção industrial podem exigir o treino com imagens grandes, cujos mapas de características consomem mais memória do que os pesos do modelo. Criar checkpoints para etapas selecionadas da rede base pode preservar a resolução da imagem ou permitir mais amostras por lote, em vez de reduzir drasticamente a resolução das entradas.
-
Treino de transformers com sequências longas: O armazenamento de ativações cresce rapidamente com o aumento do comprimento da sequência e do número de camadas. Recalcular blocos transformer pode permitir que contextos mais longos ou microlotes maiores caibam no mesmo acelerador. O guia da NVIDIA sobre recálculo de ativações ilustra o recálculo total e seletivo de camadas transformer.
Orientações práticas#
Usa o checkpointing de gradientes quando o perfil de desempenho mostrar que as ativações, e não os parâmetros ou o estado do otimizador, são o principal fator de consumo de memória. Começa pelos blocos repetidos de maiores dimensões e mede o pico de memória, o tempo por iteração e o comportamento na validação antes de alargar a utilização.
Nas configurações documentadas de treino do Ultralytics YOLO, os primeiros controlos de memória a considerar incluem a precisão mista automática, o tamanho da imagem e o tamanho físico do lote. A referência do Ultralytics AutoBatch explica a seleção automática do lote com base na memória da GPU disponível. Quando o hardware local continua a ser insuficiente, o treino na cloud da Ultralytics Platform disponibiliza GPUs configuráveis na cloud para execuções de treino geridas. O checkpointing de gradientes pode complementar estes controlos quando uma arquitetura PyTorch personalizada exige uma gestão mais precisa da memória das ativações.









