Hoppa till innehållet
AI-grafen
E· Universitetmodelltraning-finjustering· ca 60 min· utvecklande· verifierad 2026-09-20

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:

PostStorlekBeror på
ParametrarP×P \times bytesmodellen
GradienterP×P \times bytesmodellen
Optimerartillstånd2P×2P \times 4 B för Adammodellen
Aktiveringar∝ batchstorlek × sekvenslängd × djupbatchen

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:

PostMinne
Parametrar (2 B)2 GB
Gradienter (2 B)2 GB
Adams m och v (4 B vardera)8 GB
Summa innan aktiveringar12 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 kk av dem.

Beffektiv=Bmikro×k×antal GPU:erB_{\text{effektiv}} = B_{\text{mikro}} \times k \times \text{antal GPU:er}

Matematiskt är det nästan identiskt med en stor batch. Två detaljer:

  1. Dividera förlusten med kk, annars blir gradienten kk gånger för stor.
  2. 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.

MinneTid
UtanO(L)O(L)1×
Med, varje lagerO(L)O(\sqrt{L})~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ärdKostnad
1Minska mikrobatchen, öka ackumuleringeningen (nästan)
2bf16/fp16 i stället för fp32ingen kvalitetsförlust i praktiken
3Activation checkpointing~30 % längre tid
48-bitars optimerareliten kvalitetsrisk
5Kortare sekvenserberor på uppgiften
6Sharding (ZeRO/FSDP) över flera GPU:erkräver fler kort
7Mindre modellsista utvägen

Två vanliga misstag som ger OOM utan att det är modellens fel:

  • Att ackumulera tensorer med gradientgraf. total_loss += loss sparar hela grafen för varje batch. Använd loss.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

Alla källor och licenser