Gradient Checkpointing
Scopri come il checkpointing del gradiente riduce l'uso della memoria GPU ricalcolando le attivazioni durante la retropropagazione, con esempi PyTorch, compromessi e indicazioni pratiche per l'addestramento.
Il gradient checkpointing è una tecnica di addestramento che riduce l'uso di memoria: conserva solo alcune attivazioni intermedie del passaggio in avanti e ricalcola le altre durante la retropropagazione. Chiamata anche checkpointing delle attivazioni, scambia un aumento dei calcoli con un picco di utilizzo della memoria inferiore. Nonostante il nome, la tecnica crea checkpoint delle attivazioni, non dei gradienti dei parametri né dei file del modello; è quindi particolarmente utile quando i tensori delle attivazioni impediscono a una rete neurale di rientrare nella memoria GPU disponibile.
Come funziona il gradient checkpointing#
Durante un normale passaggio in avanti, una rete neurale calcola tensori intermedi chiamati attivazioni. Il sistema di differenziazione automatica conserva le attivazioni necessarie a calcolare successivamente i gradienti, come descritto nei meccanismi di autograd di PyTorch. Reti profonde, batch di grandi dimensioni, immagini ad alta risoluzione e sequenze di input lunghe possono far sì che questi tensori occupino molta memoria.
Il gradient checkpointing divide la rete in segmenti:
- Il passaggio in avanti conserva gli input o le attivazioni ai confini dei segmenti selezionati.
- Le altre attivazioni intermedie all'interno di questi segmenti vengono eliminate.
- Durante il passaggio all'indietro, ogni segmento con checkpoint viene eseguito di nuovo in avanti per ricostruire i valori mancanti.
- Le attivazioni ricostruite vengono usate immediatamente per calcolare i gradienti.
L'API di PyTorch per il checkpointing delle attivazioni espone questo comportamento tramite torch.utils.checkpoint. Concetti equivalenti sono noti come gradient checkpointing e rimaterializzazione in JAX e checkpointing delle tape in TensorFlow.
La posizione dei checkpoint determina il compromesso. In genere, creare checkpoint per più regioni consente di risparmiare più memoria, ma comporta la ripetizione di più operazioni. Una collocazione selettiva attorno ai blocchi che richiedono molte attivazioni può offrire un equilibrio migliore rispetto al ricalcolo dell'intera rete.
Compromessi e tecniche correlate#
Il gradient checkpointing in genere lascia invariati i parametri del modello, i gradienti e gli stati dell'ottimizzatore. Il suo obiettivo principale è la memoria delle attivazioni e l'addestramento diventa più lento perché alcuni calcoli in avanti vengono eseguiti due volte. Il risultato esatto dipende dall'architettura, dai confini dei checkpoint, dalla forma del batch e dall'hardware.
Si differenzia da diverse tecniche correlate:
- La accumulazione del gradiente elabora più microbatch prima di aggiornare i pesi del modello, creando un batch effettivo più grande senza caricare tutti i campioni simultaneamente. Il gradient checkpointing riduce invece le attivazioni conservate per ogni microbatch.
- La precisione mista usa tipi di dati a precisione inferiore per alcune operazioni. Il flusso di lavoro di precisione mista automatica di PyTorch può ridurre l'uso di memoria e accelerare le operazioni compatibili, mentre il checkpointing aggiunge intenzionalmente calcoli.
- La riduzione della dimensione del batch riduce l'uso di memoria elaborando meno campioni contemporaneamente. Il checkpointing può rendere fattibile una dimensione del batch o una risoluzione di input maggiore.
- I checkpoint di addestramento salvano i pesi e lo stato dell'ottimizzatore per il ripristino o per inferenze successive. La guida ai checkpoint dei modelli PyTorch descrive questo meccanismo di persistenza, che non è correlato al ricalcolo delle attivazioni.
Il codice sottoposto a checkpoint deve essere funzionalmente coerente tra l'esecuzione originale e quella ricalcolata del passaggio in avanti. Stato mutabile, trasferimenti tra dispositivi o casualità non controllata all'interno di una regione con checkpoint possono causare errori o gradienti errati. Nei Transformer in stile decoder, potrebbe essere necessario disabilitare anche una cache chiave-valore durante l'addestramento con checkpoint, perché lo stato memorizzato nella cache per l'inferenza può entrare in conflitto con la ricostruzione del grafo del passaggio in avanti.
Esempio con PyTorch#
L'esempio seguente applica il checkpointing a un blocco che richiede molta memoria durante un passaggio di addestramento:
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()Vengono conservati solo gli input del blocco e le informazioni di confine necessarie; le attivazioni interne vengono ricostruite durante loss.backward(). Nei progetti reali, confronta le esecuzioni con e senza checkpointing usando una misurazione come la memoria GPU massima allocata in PyTorch.
Applicazioni nel mondo reale#
-
Visione artificiale ad alta risoluzione: La segmentazione medica, il rilevamento aereo e l'ispezione industriale possono richiedere l'addestramento su immagini di grandi dimensioni, le cui mappe delle caratteristiche occupano più memoria dei pesi del modello. Applicare il checkpointing a fasi selezionate del backbone può preservare la risoluzione delle immagini o consentire più campioni per batch, invece di ridimensionare drasticamente gli input.
-
Addestramento di Transformer su sequenze lunghe: Lo spazio occupato dalle attivazioni cresce rapidamente all'aumentare della lunghezza della sequenza e del numero di layer. Ricalcolare i blocchi Transformer può permettere di usare contesti più lunghi o microbatch più grandi sullo stesso acceleratore. La guida NVIDIA al ricalcolo delle attivazioni illustra il ricalcolo completo e selettivo per i layer Transformer.
Indicazioni pratiche#
Usa il gradient checkpointing quando la profilazione mostra che sono le attivazioni, e non i parametri o lo stato dell'ottimizzatore, a dominare l'uso di memoria. Inizia dai blocchi ripetuti più grandi e misura la memoria di picco, il tempo di iterazione e il comportamento in validazione prima di estenderne l'applicazione.
Per le configurazioni documentate di addestramento Ultralytics YOLO, i controlli iniziali per la memoria includono la precisione mista automatica, le dimensioni delle immagini e la dimensione fisica del batch. La documentazione di riferimento di Ultralytics AutoBatch spiega la selezione automatica del batch in base alla memoria GPU disponibile. Quando l'hardware locale non è sufficiente, l'addestramento cloud sulla Ultralytics Platform offre GPU cloud configurabili per sessioni di addestramento gestite. Il gradient checkpointing può integrare questi controlli quando un'architettura PyTorch personalizzata richiede una gestione più precisa della memoria delle attivazioni.









