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:
| Del | Precision | Varför |
|---|---|---|
| Matrismultiplikationer | bf16/fp16 | snabbast, tål grov upplösning |
| Aktiveringar | bf16/fp16 | minne |
| Huvudvikter | fp32 | små uppdateringar får inte försvinna |
| Optimerartillstånd | fp32 | ackumuleras över tusentals steg |
| Softmax, normalisering, förlust | fp32 | numeriskt 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:
| fp16 | bf16 | |
|---|---|---|
| Loss scaling | krävs | behövs inte |
| Risk för överflöd | ja | nej |
| Precision | bättre mantissa | grövre |
| Hårdvara | Volta och senare | Ampere 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:
- Multiplicera förlusten med (typiskt startvärde 2¹⁶).
- Kör
backward()— alla gradienter är nu gånger större och hamnar inom fp16:s område. - Unscale gradienterna (dividera med ) före klippning och optimerarsteg.
- Upptäck
inf/NaNi gradienterna: hoppa över steget och halvera . - Har inga överflöden inträffat på ett tag: dubbla .
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 |
|---|---|
| 1 | Kö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) |
| 3 | Sker unscale_ före klippningen? |
| 4 | Finns log(0), division med noll eller sqrt av negativt tal? |
| 5 | Är någon normalisering i fp16 med mycket små varianser? |
| 6 | Byt 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ör | Varför |
|---|---|
loss.backward() | bakåtpasset ärver precisionen från framåtpasset |
| Optimerarsteget | ska köras i fp32 |
| Förlustberäkning vid känsliga förluster | numerisk 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
- arXiv — Mixed Precision Training — arXiv (öppen åtkomst; licens per artikel)
- PyTorch — tutorials (BSD-3) — BSD-3-Clause
- NVIDIA — Mixed Precision Training Guide — dokumentation, fri läsning