Gradient checkpointing: como poupa memória da GPU?
Tradução automática
Este artigo foi traduzido automaticamente a partir da versão original em inglês.
O gradient checkpointing guarda menos ativações intermédias durante a passagem para a frente e volta a calcular as que faltam durante o cálculo para trás. Utilize-o quando as ativações impedirem que o lote de treino ou o comprimento da sequência caibam na memória da GPU e o cálculo adicional for aceitável.
Reduz a memória ocupada pelas ativações guardadas. Não elimina os pesos do modelo, o estado do otimizador nem todos os buffers temporários.
O que é calculado novamente?
O cálculo para trás precisa de resultados intermédios da passagem para a frente para calcular os gradientes. Sem checkpointing, o framework conserva os resultados necessários. Com checkpointing, guarda determinadas entradas ou resultados intermédios e volta a executar as operações afetadas quando os seus valores são necessários.
Por exemplo, um framework pode guardar a entrada de um grupo de camadas de um transformer e descartar as suas ativações internas. Durante o cálculo para trás, volta a executar esse grupo a partir da entrada guardada. A entrada guardada é um checkpoint de ativações, diferente de um checkpoint do modelo escrito em disco para recuperação.
A análise original de Chen et al. apresenta uma estratégia com segmentos uniformes que usa O(sqrt(n)) de memória de ativações para uma rede com n camadas e exige aproximadamente uma passagem adicional para a frente por passo de treino. Os transformers reais têm ativações de tamanhos diferentes, diversas implementações de atenção e políticas dos frameworks. Meça a memória e o tempo obtidos em vez de usar a estimativa assintótica como previsão para uma implementação em produção.
Ative a integração prevista pelo seu ambiente
Para os modelos suportados pelo Hugging Face Trainer, gradient_checkpointing=True ativa o mecanismo geral. Com FSDP, verifique antes a integração de checkpointing de ativações do framework. A documentação atual do Transformers Trainer recomenda activation_checkpointing para FSDP porque o gradient checkpointing genérico pode causar um all-gather adicional durante o cálculo para trás.
Compare a mesma configuração de treino com e sem checkpointing. Registe o pico de memória atribuída na GPU, os exemplos ou tokens por segundo e se o comprimento de sequência pretendido cabe na memória. Verifique também se os gradientes continuam válidos quando o modelo usa operações personalizadas ou treino de adaptadores.
Se o estado do otimizador ocupar a maior parte da memória, considere antes dividir esse estado entre dispositivos ou usar um ajuste eficiente em parâmetros; voltar a calcular as ativações não elimina essa atribuição de memória.
A secção sobre gradient checkpointing do LLM Engineering Guide explica o cálculo por segmentos.