Gradient Checkpointing
Aprende cómo el *gradient checkpointing* reduce la memoria de la GPU recomputando las activaciones durante la propagación hacia atrás, con ejemplos en PyTorch, ventajas y desventajas, y pautas prácticas de entrenamiento.
El gradiente acumulado (gradient checkpointing) es una técnica de entrenamiento que ahorra memoria y que almacena solo las activaciones intermedias seleccionadas de la pasada hacia adelante y vuelve a calcular las demás durante la propagación hacia atrás. También llamado checkpointing de activaciones, intercambia cálculo adicional por un menor uso de memoria máxima. A pesar de su nombre, la técnica aplica puntos de control a las activaciones en lugar de a los gradientes de los parámetros o a los archivos del modelo, lo que la hace especialmente valiosa cuando los tensores de activación evitan que una red neuronal quepa en la memoria de la GPU disponible.
Cómo funciona el Gradient Checkpointing#
Durante una pasada hacia adelante estándar, una red neuronal calcula tensores intermedios llamados activaciones. El sistema de diferenciación automática retiene las activaciones necesarias para calcular gradientes más tarde, como se describe en la mecánica del autograd de PyTorch. Las redes profundas, los lotes grandes, las imágenes de alta resolución y las secuencias de entrada largas pueden hacer que estos tensores guardados consuman una memoria considerable.
El gradient checkpointing divide la red en segmentos:
- La pasada hacia adelante almacena las entradas o las activaciones de los límites para los segmentos seleccionados.
- Otras activaciones intermedias dentro de esos segmentos se descartan.
- Durante la pasada hacia atrás, cada segmento con punto de control se ejecuta hacia adelante de nuevo para reconstruir los valores faltantes.
- Las activaciones reconstruidas se utilizan inmediatamente para calcular los gradientes.
La API de checkpointing de activaciones de PyTorch expone este comportamiento a través de torch.utils.checkpoint. Conceptos equivalentes aparecen como gradient checkpointing y rematerialización en JAX y checkpointing de cintas en TensorFlow.
La ubicación de los puntos de control determina el equilibrio. Establecer puntos de control en más regiones generalmente ahorra más memoria pero repite más operaciones. Una ubicación selectiva alrededor de los bloques con alta carga de activaciones puede proporcionar un mejor equilibrio que volver a calcular toda la red.
Compromisos y técnicas relacionadas#
El gradient checkpointing normalmente deja los parámetros del modelo, los gradientes y los estados del optimizador sin cambios. Su objetivo principal es la memoria de activaciones, y el entrenamiento se vuelve más lento porque algunos cálculos hacia adelante se ejecutan dos veces. El resultado exacto depende de la arquitectura, los límites de los puntos de control, la forma del lote y el hardware.
Difiere de varias técnicas relacionadas:
- La acumulación de gradientes procesa múltiples microlotes antes de actualizar los pesos del modelo, creando un lote efectivo más grande sin cargar cada muestra simultáneamente. En su lugar, el gradient checkpointing reduce las activaciones retenidas para cada microlote.
- La precisión mixta utiliza tipos de datos de menor precisión para operaciones seleccionadas. El flujo de trabajo de precisión mixta automática de PyTorch puede reducir la memoria y acelerar las operaciones compatibles, mientras que el checkpointing agrega deliberadamente cálculo.
- La reducción del tamaño de lote disminuye la memoria procesando menos muestras juntas. El checkpointing puede permitir que un lote más grande o una resolución de entrada sigan siendo viables.
- Los puntos de control de entrenamiento guardan los pesos y el estado del optimizador para su recuperación o inferencia posterior. La guía de puntos de control de modelos de PyTorch describe este mecanismo de persistencia, el cual no está relacionado con el recálculo de activaciones.
El código con puntos de control debe ser funcionalmente coherente entre sus ejecuciones hacia adelante original y recalculada. El estado mutable, las transferencias de dispositivos o la aleatoriedad descontrolada dentro de una región con puntos de control pueden causar errores o gradientes incorrectos. En los transformadores de estilo descodificador, es posible que también sea necesario deshabilitar una caché de clave-valor durante el entrenamiento con puntos de control, ya que el estado de inferencia almacenado en caché puede entrar en conflicto con la reconstrucción del grafo de avance.
Ejemplo en PyTorch#
El siguiente ejemplo aplica puntos de control a un bloque que consume mucha memoria durante un paso de entrenamiento:
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()Solo se retienen las entradas del bloque y la información de contorno requerida; sus activaciones internas se reconstruyen durante loss.backward(). Los proyectos reales deben comparar las ejecuciones con y sin checkpointing utilizando una medida como la memoria de GPU asignada máxima en PyTorch.
Aplicaciones en el mundo real#
-
Visión artificial de alta resolución: La segmentación médica, la detección aérea y la inspección industrial pueden entrenarse con imágenes grandes cuyos mapas de características consumen más memoria que los pesos del modelo. Aplicar puntos de control en etapas seleccionadas del backbone puede preservar la resolución de la imagen o permitir muestras adicionales por lote en lugar de reducir drásticamente el tamaño de las entradas.
-
Entrenamiento de Transformer de secuencia larga: El almacenamiento de activaciones crece rápidamente a medida que aumentan la longitud de la secuencia y el recuento de capas. Volver a calcular los bloques de Transformer puede hacer que contextos más largos o microlotes más grandes quepan en el mismo acelerador. La guía de recomputación de activaciones de NVIDIA ilustra la recomputación total y selectiva para las capas de Transformer.
Guía práctica#
Usa el gradient checkpointing cuando la perfilación muestre que las activaciones, en lugar de los parámetros o el estado del optimizador, dominan la memoria. Comienza con bloques grandes repetidos y evalúa el rendimiento de la memoria máxima, el tiempo de iteración y el comportamiento de validación antes de ampliar la cobertura.
Para las configuraciones de entrenamiento de Ultralytics YOLO documentadas, los controles de memoria de primera línea incluyen la precisión mixta automática, el tamaño de imagen y el tamaño de lote físico. La referencia de Ultralytics AutoBatch explica la selección automática de lotes según la memoria de la GPU disponible. Cuando el hardware local sigue siendo insuficiente, la entrenamiento en la nube de Ultralytics Platform proporciona GPUs en la nube configurables para ejecuciones de entrenamiento administradas. El gradient checkpointing puede complementar estos controles cuando una arquitectura de PyTorch personalizada requiere una gestión más fina de la memoria de activaciones.






