GPU-minne, gradientackumulering och batchar
Kunna beräkna minnesbehov och använda gradientackumulering och checkpointing.
Förkunskaper
Intuition
«CUDA out of memory» är det vanligaste felet i djupinlärning. Det går att räkna ut i förväg vad som ryms.
Fyra poster tar minne under träning:
| Post | Storlek | Beror på |
|---|---|---|
| Parametrar | bytes | modellen |
| Gradienter | bytes | modellen |
| Optimerartillstånd | 4 B för Adam | modellen |
| Aktiveringar | ∝ batchstorlek × sekvenslängd × djup | batchen |
De tre första är fasta — de beror bara på modellen. Den fjärde växer med batchen, och det är därför «minska batchstorleken» är det första man gör.
Exempel, en modell med 1 miljard parametrar i bf16 med Adam:
| Post | Minne |
|---|---|
| Parametrar (2 B) | 2 GB |
| Gradienter (2 B) | 2 GB |
| Adams m och v (4 B vardera) | 8 GB |
| Summa innan aktiveringar | 12 GB |
På ett 24 GB-kort återstår alltså cirka 12 GB för aktiveringar — det är de som avgör hur stor batch du kan köra.
Formellt
Gradientackumulering ger stor effektiv batch utan stort minne: kör flera små batchar, summera gradienterna, och ta ett steg först efter av dem.
Matematiskt är det nästan identiskt med en stor batch. Två detaljer:
- Dividera förlusten med , annars blir gradienten gånger för stor.
- BatchNorm normaliserar per mikrobatch, så statistiken skiljer sig från en äkta stor batch. LayerNorm och GroupNorm påverkas inte — ännu ett skäl till att transformerar använder LayerNorm.
Activation checkpointing byter beräkning mot minne: spara bara aktiveringarna vid vissa punkter och räkna om resten i bakåtpasset.
| Minne | Tid | |
|---|---|---|
| Utan | 1× | |
| Med, varje lager | ~1,3× |
Trettio procent längre träningstid för att kunna köra en modell som annars inte får plats är nästan alltid en bra affär.
Åtgärdslistan vid OOM, i ordning efter hur mycket de kostar:
| # | Åtgärd | Kostnad |
|---|---|---|
| 1 | Minska mikrobatchen, öka ackumuleringen | ingen (nästan) |
| 2 | bf16/fp16 i stället för fp32 | ingen kvalitetsförlust i praktiken |
| 3 | Activation checkpointing | ~30 % längre tid |
| 4 | 8-bitars optimerare | liten kvalitetsrisk |
| 5 | Kortare sekvenser | beror på uppgiften |
| 6 | Sharding (ZeRO/FSDP) över flera GPU:er | kräver fler kort |
| 7 | Mindre modell | sista utvägen |
Två vanliga misstag som ger OOM utan att det är modellens fel:
- Att ackumulera tensorer med gradientgraf.
total_loss += losssparar hela grafen för varje batch. Användloss.item()eller.detach(). - Att glömma
torch.no_grad()vid utvärdering. Då byggs en graf som aldrig används, och minnet dubblas.
Kod
import torch, torch.nn as nn
def minnesuppskattning(parametrar, bytes_per_param=2, optimerare="adam"):
"""GB för parametrar, gradienter och optimerartillstånd — allt utom aktiveringar."""
per_opt = {"adam": 8, "adamw": 8, "sgd_momentum": 4, "sgd": 0, "adam8bit": 2}[optimerare]
byte = parametrar * (bytes_per_param * 2 + per_opt)
return round(byte / 1024**3, 2)
for n, namn in ((125e6, "125M"), (1e9, "1B"), (7e9, "7B")):
print(f"{namn:>5}: adam {minnesuppskattning(n):>6.2f} GB "
f"adam8bit {minnesuppskattning(n, optimerare='adam8bit'):>6.2f} GB")
# 125M: 1.40 GB 0.70 GB
# 1B: 11.18 GB 5.59 GB
# 7B: 78.23 GB 39.12 GB
# Gradientackumulering
ACK = 8 # effektiv batch = mikrobatch × 8
opt.zero_grad()
for i, (x, y) in enumerate(dataloader):
loss = nn.functional.cross_entropy(modell(x), y) / ACK # ← dividera!
loss.backward()
if (i + 1) % ACK == 0:
nn.utils.clip_grad_norm_(modell.parameters(), 1.0)
opt.step()
opt.zero_grad()
# Activation checkpointing
from torch.utils.checkpoint import checkpoint
class Stapel(nn.Module):
def __init__(self, block, anvand_checkpoint=True):
super().__init__()
self.block = nn.ModuleList(block)
self.cp = anvand_checkpoint
def forward(self, x):
for b in self.block:
x = checkpoint(b, x, use_reentrant=False) if self.cp and self.training else b(x)
return x
# Mät faktiskt toppminne
def toppminne(fn):
torch.cuda.reset_peak_memory_stats()
fn()
return round(torch.cuda.max_memory_allocated() / 1024**3, 2)
# De två klassiska misstagen
total = 0.0
for x, y in dataloader:
loss = nn.functional.cross_entropy(modell(x), y)
total += loss.item() # RÄTT — .item() bryter grafen
# total += loss # FEL — sparar hela grafen, minnet växer varje batch
modell.eval()
with torch.no_grad(): # RÄTT — utan detta byggs en graf som aldrig används
for x, y in val_loader:
modell(x)
Behärskning innebär
- Beräknar minnesbehovet för en träningskörning
- Använder gradientackumulering
- Vet vad activation checkpointing kostar och ger
Logga in för att göra övningarna och bygga upp din behärskning.
Källor
- PyTorch — tutorials (BSD-3) — BSD-3-Clause
- arXiv — Training Deep Nets with Sublinear Memory Cost — arXiv (öppen åtkomst; licens per artikel)
- arXiv — ZeRO: Memory Optimizations Toward Training Trillion Parameter Models — arXiv (öppen åtkomst; licens per artikel)