Hoppa till innehållet
AI-grafen
F· AI engineeringdeep-learning· ca 90 min· volatil — kontrolleras ofta· verifierad 2026-09-21

Mixed precision-träning

Kunna träna i bf16/fp16 med loss scaling och veta när det går fel.

Förkunskaper

Intuition

Mixed precision betyder att olika delar av träningen körs i olika precision:

DelPrecisionVarför
Matrismultiplikationerbf16/fp16snabbast, tål grov upplösning
Aktiveringarbf16/fp16minne
Huvudvikterfp32små uppdateringar får inte försvinna
Optimerartillståndfp32ackumuleras över tusentals steg
Softmax, normalisering, förlustfp32numeriskt känsliga

Vinsten är ungefär 2–3× snabbare träning och nästan halverat aktiveringsminne, med i praktiken oförändrad kvalitet.

Två sätt, och de skiljer sig i hur mycket besvär de orsakar:

fp16bf16
Loss scalingkrävsbehövs inte
Risk för överflödjanej
Precisionbättre mantissagrövre
HårdvaraVolta och senareAmpere och senare

Stöder hårdvaran bf16 — använd bf16. Det tar bort hela loss scaling-problematiken.

Formellt

Loss scaling för fp16, steg för steg:

  1. Multiplicera förlusten med SS (typiskt startvärde 2¹⁶).
  2. Kör backward() — alla gradienter är nu SS gånger större och hamnar inom fp16:s område.
  3. Unscale gradienterna (dividera med SS) före klippning och optimerarsteg.
  4. Upptäck inf/NaN i gradienterna: hoppa över steget och halvera SS.
  5. Har inga överflöden inträffat på ett tag: dubbla SS.

Det kallas dynamisk loss scaling och sköts av torch.amp.GradScaler. Punkt 3 är den som görs fel oftast — klipper man gradienterna före unscale_ klipper man mot en faktor 65 536 gånger för stor norm, och klippningen gör ingenting.

Felsökningsordning när mixed precision ger NaN:

#Kontroll
1Kör samma steg i fp32 — blir det NaN där också är felet inte precisionen
2Är förlustberäkningen i fp32? (autocast bör undantas kring den)
3Sker unscale_ före klippningen?
4Finns log(0), division med noll eller sqrt av negativt tal?
5Är någon normalisering i fp16 med mycket små varianser?
6Byt till bf16 — löser överflödsproblemen helt

Kontroll 1 först, alltid. Mixed precision får ofta skulden för fel som finns i fp32 också.

Vad som inte ska ligga i autocast:

Ligger utanförVarför
loss.backward()bakåtpasset ärver precisionen från framåtpasset
Optimerarstegetska köras i fp32
Förlustberäkning vid känsliga förlusternumerisk stabilitet

Inferens i låg precision är enklare: ingen optimerare, inga gradienter. bf16 eller fp16 rakt av fungerar nästan alltid, och därifrån går vägen vidare till int8 och int4 med kvantisering.

Mät alltid. Kör 200 steg i fp32 och 200 i mixed precision och jämför förlustkurvorna. Ligger de ovanpå varandra är allt bra; divergerar de finns ett problem som inte syns i en enskild mätpunkt.

Kod

import torch, torch.nn as nn

# bf16 — enklast, ingen scaler behövs
def trana_bf16(modell, dataloader, opt, epoker=1):
    for _ in range(epoker):
        for x, y in dataloader:
            with torch.autocast("cuda", dtype=torch.bfloat16):
                loss = nn.functional.cross_entropy(modell(x), y)
            opt.zero_grad(set_to_none=True)
            loss.backward()                       # UTANFÖR autocast
            nn.utils.clip_grad_norm_(modell.parameters(), 1.0)
            opt.step()

# fp16 — kräver GradScaler, och ordningen är kritisk
def trana_fp16(modell, dataloader, opt, epoker=1):
    scaler = torch.amp.GradScaler("cuda")
    for _ in range(epoker):
        for x, y in dataloader:
            with torch.autocast("cuda", dtype=torch.float16):
                loss = nn.functional.cross_entropy(modell(x), y)
            opt.zero_grad(set_to_none=True)
            scaler.scale(loss).backward()         # 1–2: skala och bakåt
            scaler.unscale_(opt)                  # 3: MÅSTE komma före klippning
            nn.utils.clip_grad_norm_(modell.parameters(), 1.0)
            scaler.step(opt)                      # 4: hoppar över steget vid inf/NaN
            scaler.update()                       # 5: justerar skalfaktorn

# Håll ett känsligt block i fp32
with torch.autocast("cuda", dtype=torch.bfloat16):
    h = modell.ryggrad(x)
    with torch.autocast("cuda", enabled=False):   # av för just detta
        logits = modell.huvud(h.float())
        loss = egen_kanslig_forlust(logits, y)

# Jämför kurvorna i stället för att lita på att det fungerar
def jamfor_precision(bygg_modell, dataloader, steg=200):
    kurvor = {}
    for namn, dtype in (("fp32", None), ("bf16", torch.bfloat16), ("fp16", torch.float16)):
        torch.manual_seed(0)
        m = bygg_modell()
        opt = torch.optim.AdamW(m.parameters(), lr=1e-4)
        scaler = torch.amp.GradScaler("cuda", enabled=(dtype is torch.float16))
        losses = []
        for i, (x, y) in enumerate(dataloader):
            if i >= steg:
                break
            with torch.autocast("cuda", dtype=dtype, enabled=dtype is not None):
                loss = nn.functional.cross_entropy(m(x), y)
            opt.zero_grad(set_to_none=True)
            scaler.scale(loss).backward()
            scaler.step(opt); scaler.update()
            losses.append(float(loss))
        kurvor[namn] = losses
    return kurvor

# Upptäck NaN tidigt i stället för hundra steg senare
def nan_vakt(modell):
    def krok(namn):
        def f(m, i, o):
            t = o[0] if isinstance(o, tuple) else o
            if torch.is_tensor(t) and not torch.isfinite(t).all():
                raise RuntimeError(f"icke-finita värden i {namn}")
        return f
    for namn, m in modell.named_modules():
        m.register_forward_hook(krok(namn))

nan_vakt är värd att ha under utveckling. Utan den märks ett NaN först när förlusten blir NaN, och då är det svårt att veta vilket lager som orsakade det.

Behärskning innebär

  • Sätter upp mixed precision-träning
  • Använder loss scaling korrekt
  • Felsöker NaN och divergens

Logga in för att göra övningarna och bygga upp din behärskning.

Källor

Alla källor och licenser