Numerisk stabilitet och flyttal
Kunna känna igen overflow, underflow och cancellation och använda log-sum-exp-tricket.
Förkunskaper
- CPython — grundernakrävs
- DLogaritmerkrävsÖva i Mattegrafen ↗
Intuition
Datorer räknar inte exakt med decimaltal. 0.1 + 0.2 blir 0.30000000000000004, eftersom 0,1 inte går att skriva exakt i bas 2 — precis som 1/3 inte går att skriva exakt i bas 10.
Fyra fällor:
| Fälla | Vad som händer | Exempel |
|---|---|---|
| Overflow | talet blir för stort | exp(1000) → inf |
| Underflow | talet blir 0 | 0.1 ** 400 → 0.0 |
| Cancellation | precision försvinner vid subtraktion | (1 + 1e-16) - 1 → 0.0 |
| Jämförelse | == på flyttal | 0.1 + 0.2 == 0.3 → False |
Räckvidden för de vanliga formaten:
| Format | Största | Minsta normala | Decimalsiffror |
|---|---|---|---|
| float64 | ~1,8·10³⁰⁸ | ~2,2·10⁻³⁰⁸ | ~16 |
| float32 | ~3,4·10³⁸ | ~1,2·10⁻³⁸ | ~7 |
| float16 | ~65 504 | ~6,1·10⁻⁵ | ~3 |
float16 svämmar över redan vid 65 504. Det är därför mixed precision-träning behöver loss scaling.
Formellt
Log-sum-exp-tricket är det viktigaste enskilda knepet. Du vill beräkna
Med svämmar över. Men eftersom
för vilket som helst, väljer man . Då är den största exponenten exakt 0, och . Ingen overflow, samma svar.
Samma sak i softmax:
Det här är inte en detalj — det är inbyggt i varje seriös implementation av softmax och korsentropi. Det är också därför cross_entropy i PyTorch tar logits, inte sannolikheter: den kombinerar log-softmax och NLL i en numeriskt stabil operation. Att först köra softmax och sedan log är både långsammare och instabilt.
Cancellation uppstår när två nästan lika tal subtraheras: de gemensamma siffrorna tar ut varandra och bara skräpet i slutet återstår.
| Naivt | Stabilt | Varför |
|---|---|---|
log(1 + x) | log1p(x) | för små x är 1 + x == 1 |
exp(x) - 1 | expm1(x) | samma problem |
sqrt(x²+y²) | hypot(x, y) | undviker overflow i mellansteget |
mean((x-m)²) i ett svep | Welfords algoritm | undviker cancellation i varians |
Aldrig == på flyttal. Använd math.isclose eller np.allclose med en tolerans. Och i tester: ange både absolut och relativ tolerans, eftersom den ena räcker för små tal och den andra för stora.
Praktisk regel i ML: arbeta i log-rummet så länge du kan, konvertera tillbaka så sent som möjligt, och lita på bibliotekens logsumexp, log_softmax och cross_entropy i stället för att skriva dem själv.
Kod
import math, numpy as np
from scipy.special import logsumexp
# 1. Overflow
x = np.array([1000.0, 1001.0, 1002.0])
try:
print(np.log(np.exp(x).sum()))
except Exception:
pass
print(np.log(np.exp(x).sum())) # inf (med varning)
print(logsumexp(x)) # 1002.4076059644443 ← stabilt
# Samma sak för hand
c = x.max()
print(c + np.log(np.exp(x - c).sum())) # 1002.4076059644443
# 2. Underflow
print(math.prod([0.1] * 400)) # 0.0 — all information borta
print(sum(math.log(0.1) for _ in range(400))) # -921.03
# 3. Cancellation
print((1 + 1e-16) - 1) # 0.0
print(math.log(1 + 1e-16), math.log1p(1e-16)) # 0.0 1e-16
# 4. Jämförelse
print(0.1 + 0.2 == 0.3) # False
print(0.1 + 0.2) # 0.30000000000000004
print(math.isclose(0.1 + 0.2, 0.3)) # True
# 5. Softmax: naiv mot stabil
def softmax_naiv(z):
e = np.exp(z); return e / e.sum()
def softmax(z):
e = np.exp(z - np.max(z)); return e / e.sum()
z = np.array([1000.0, 1001.0, 1002.0])
print(softmax_naiv(z)) # [nan nan nan]
print(softmax(z)) # [0.09003057 0.24472847 0.66524096]
# 6. Varför cross_entropy tar logits, inte sannolikheter
import torch
logits = torch.tensor([[1000.0, 1001.0, 1002.0]])
mal = torch.tensor([2])
print(torch.nn.functional.cross_entropy(logits, mal)) # tensor(0.4076)
p = torch.softmax(logits, -1)
print(torch.nn.functional.nll_loss(torch.log(p), mal)) # tensor(0.4076) men instabil väg
Behärskning innebär
- Känner igen overflow, underflow och cancellation
- Använder log-sum-exp
- Vet varför flyttalsjämförelser är farliga
Logga in för att göra övningarna och bygga upp din behärskning.