Flyttalsformat: fp32, fp16, bf16
Kunna förklara skillnaden mellan precisionerna och när mixed precision är säkert.
Förkunskaper
Intuition
Ett flyttal består av tre delar: tecken, exponent (hur stort) och mantissa (hur exakt).
| Format | Bitar | Tecken | Exponent | Mantissa | Största tal | Decimalsiffror |
|---|---|---|---|---|---|---|
| fp32 | 32 | 1 | 8 | 23 | ~3,4·10³⁸ | ~7 |
| fp16 | 16 | 1 | 5 | 10 | 65 504 | ~3 |
| bf16 | 16 | 1 | 8 | 7 | ~3,4·10³⁸ | ~2 |
| fp8 (E4M3) | 8 | 1 | 4 | 3 | 448 | ~1 |
Den avgörande raden är bf16. Den har samma exponent som fp32 och offrar i stället mantissa. Konsekvensen är att den täcker samma talområde — inga överflöden — men med grövre upplösning.
Och det visade sig vara precis rätt avvägning för djupinlärning: räckvidd spelar mer roll än precision. Gradienter kan vara mycket små eller mycket stora, men de behöver sällan sju decimalers noggrannhet.
Formellt
Varför fp16 är besvärligt. Med 5 exponentbitar är minsta normala tal cirka och största 65 504. Gradienter i djupa nät hamnar rutinmässigt under det undre gränsvärdet och blir noll.
Därför krävs loss scaling med fp16: multiplicera förlusten med en stor faktor (t.ex. 2¹⁶) före bakåtpasset, så att gradienterna hamnar inom det representerbara området, och dividera tillbaka före optimerarsteget.
Med bf16 behövs det inte — exponentområdet är detsamma som fp32. Det är hela skälet till att bf16 tagit över där hårdvaran stöder det.
Operationer som måste göras i högre precision:
| Operation | Varför |
|---|---|
| Ackumulering av summor | tusentals additioner av små tal — cancellation |
| Softmax | exponentiering och normalisering |
| Layer/BatchNorm | varians av små skillnader |
| Optimerarens tillstånd | små uppdateringar ackumuleras över tusentals steg |
| Förlustberäkning | samma skäl |
Därför är «mixed precision» blandad: matrismultiplikationerna görs i 16 bitar, allt känsligt i 32.
Den lömskaste effekten är «swamping»: adderas ett litet tal till ett stort försvinner det helt. I bf16 med 7 mantissabitar gäller 1.0 + 0.004 == 1.0. En optimerare som ackumulerar små uppdateringar i låg precision tappar dem alltså tyst — vilket är exakt varför huvudvikterna alltid hålls i fp32 även vid bf16-träning.
Hårdvarustöd avgör vad som är praktiskt möjligt:
| Format | Stöd |
|---|---|
| fp16 | NVIDIA Volta och senare |
| bf16 | NVIDIA Ampere (A100) och senare, TPU, moderna CPU:er |
| fp8 | NVIDIA Hopper (H100) och senare |
Praktisk regel: använd bf16 om hårdvaran stöder det, annars fp16 med loss scaling. fp8 för inferens och, i växande grad, för träning av mycket stora modeller.
Kod
import numpy as np, torch
for namn, dt in (("fp32", torch.float32), ("fp16", torch.float16), ("bf16", torch.bfloat16)):
i = torch.finfo(dt)
print(f"{namn}: max {i.max:.3e} min normal {i.tiny:.3e} eps {i.eps:.3e}")
# fp32: max 3.403e+38 min normal 1.175e-38 eps 1.192e-07
# fp16: max 6.550e+04 min normal 6.104e-05 eps 9.766e-04
# bf16: max 3.390e+38 min normal 1.175e-38 eps 7.812e-03
# ↑ samma räckvidd som fp32, men grövre steg
# Overflow i fp16, inte i bf16
x = torch.tensor(70000.0)
print(x.half(), x.bfloat16()) # tensor(inf, dtype=float16) tensor(70144., ...)
# Underflow: en typisk gradient i fp16
g = torch.tensor(1e-8)
print(g.half(), g.bfloat16()) # tensor(0., dtype=float16) tensor(9.9970e-09, ...)
# ↑ gradienten försvinner helt i fp16 — därför behövs loss scaling
# Loss scaling: skala upp, räkna, skala ner
SKALA = 2 ** 16
print((g * SKALA).half(), (g * SKALA).half().float() / SKALA)
# tensor(0.0007, dtype=float16) tensor(9.9977e-09) ← räddad
# Swamping: små tal försvinner vid addition
for namn, dt in (("fp32", torch.float32), ("bf16", torch.bfloat16)):
a = torch.tensor(1.0, dtype=dt)
b = torch.tensor(0.004, dtype=dt)
print(f"{namn}: 1.0 + 0.004 = {(a + b).item()}")
# fp32: 1.0 + 0.004 = 1.003999948501587
# bf16: 1.0 + 0.004 = 1.0 ← tillägget försvann helt
# Ackumulering: summera 100 000 små tal
n = 100_000
for namn, dt in (("fp16", torch.float16), ("bf16", torch.bfloat16), ("fp32", torch.float32)):
s = torch.zeros((), dtype=dt)
for _ in range(1000):
s += torch.tensor(0.001, dtype=dt)
print(f"{namn}: 1000 × 0.001 = {s.item():.4f} (exakt 1.0000)")
# ↑ därför görs ackumulering alltid i fp32
Behärskning innebär
- Förklarar exponent och mantissa
- Jämför fp32, fp16 och bf16
- Vet vilka operationer som kräver hög precision
Logga in för att göra övningarna och bygga upp din behärskning.
Källor
- Goldberg — What Every Computer Scientist Should Know About Floating-Point Arithmetic — fri läsning
- arXiv — Mixed Precision Training — arXiv (öppen åtkomst; licens per artikel)
- PyTorch — tutorials (BSD-3) — BSD-3-Clause