Gradientklippning
Kunna använda gradientklippning och förklara när den behövs.
Förkunskaper
Intuition
Ibland blir gradienten plötsligt enorm — en brant klippa i förlustlandskapet — och ett enda steg slungar iväg vikterna så långt att modellen aldrig återhämtar sig. Förlusten går till NaN och körningen är förlorad.
Gradientklippning sätter ett tak. Är gradientens norm större än tröskeln skalas hela gradienten ner så att normen blir exakt tröskeln:
Riktningen bevaras — bara längden kortas. Det är den viktiga detaljen och skillnaden mot att klippa varje element för sig.
| Metod | Gör | Bevarar riktning |
|---|---|---|
| Normklippning | skalar hela gradienten | ja |
| Värdeklippning | klipper varje element till [−c, c] | nej |
Normklippning (clip_grad_norm_) är nästan alltid rätt val.
Formellt
Var det behövs:
| Modell | Typisk tröskel | Varför |
|---|---|---|
| RNN och LSTM | 1,0–5,0 | långa beroendekedjor ger multiplikativ tillväxt |
| Transformer | 1,0 | standard i nästan alla recept |
| RL (PPO) | 0,5 | belöningar kan variera kraftigt |
| CNN | ofta onödigt | normaliseringslagren håller redan gradienterna i schack |
Var i koden — ordningen spelar roll:
loss.backward()
unscale om AMP används
clip_grad_norm_(parameters, max_norm)
optimizer.step()
Klipper du före backward() klipper du ingenting (gradienterna finns inte än). Klipper du efter step() är det för sent. Och med mixed precision måste gradienterna skalas tillbaka innan klippningen, annars klipper du mot en förstorad norm.
Välj tröskeln av data, inte av magkänsla. Logga gradientnormen i några hundra steg utan klippning och sätt tröskeln vid ungefär 90:e percentilen. Då klipps bara de verkliga utstickarna, och de normala stegen påverkas inte.
Klippning är ett skyddsnät, inte en lösning. Måste du klippa i varje steg är tröskeln för låg eller något annat är fel:
| Symtom | Verkligt problem |
|---|---|
| Klipps varje steg | lr för hög, eller tröskeln för snäv |
| Normen växer stadigt | instabil arkitektur, saknad normalisering |
| Plötsliga toppar på enstaka batchar | dålig data — avvikare, felmärkning, extremvärden |
| NaN trots klippning | division med noll, log(0), eller NaN redan i indata |
Den tredje raden är värd att undersöka: logga vilken batch som orsakade toppen och titta på den. Ofta är det en enda trasig datapunkt, och att laga datan är bättre än att klippa hårdare.
Kod
import torch, torch.nn as nn
# Grundmönstret
for x, y in dataloader:
loss = nn.functional.cross_entropy(modell(x), y)
opt.zero_grad()
loss.backward()
norm = nn.utils.clip_grad_norm_(modell.parameters(), max_norm=1.0) # returnerar normen FÖRE
opt.step()
# Med mixed precision: unscale först, annars klipper du mot en förstorad norm
scaler = torch.amp.GradScaler()
for x, y in dataloader:
with torch.autocast("cuda", dtype=torch.bfloat16):
loss = nn.functional.cross_entropy(modell(x), y)
opt.zero_grad()
scaler.scale(loss).backward()
scaler.unscale_(opt) # ← nödvändigt
nn.utils.clip_grad_norm_(modell.parameters(), 1.0)
scaler.step(opt); scaler.update()
# Välj tröskel av data: logga normer utan klippning först
import numpy as np
def matt_gradientnormer(modell, dataloader, opt, steg=200):
normer = []
for i, (x, y) in enumerate(dataloader):
if i >= steg:
break
loss = nn.functional.cross_entropy(modell(x), y)
opt.zero_grad(); loss.backward()
normer.append(float(torch.nn.utils.clip_grad_norm_(
modell.parameters(), max_norm=float("inf")))) # mät utan att klippa
opt.step()
n = np.array(normer)
return {"median": round(float(np.median(n)), 3),
"p90": round(float(np.percentile(n, 90)), 3),
"p99": round(float(np.percentile(n, 99)), 3),
"max": round(float(n.max()), 3),
"förslag_tröskel": round(float(np.percentile(n, 90)), 2)}
# Håll koll på hur ofta det faktiskt klipps
klippta = totalt = 0
for x, y in dataloader:
...
norm = nn.utils.clip_grad_norm_(modell.parameters(), 1.0)
klippta += int(float(norm) > 1.0); totalt += 1
print(f"klippt i {klippta / totalt:.1%} av stegen")
# > 50 % betyder att tröskeln är för låg eller lärhastigheten för hög
Den sista mätningen är den som gör klippning till ett medvetet val i stället för en besvärjelse man kopierat in.
Behärskning innebär
- Använder gradientklippning på rätt sätt
- Väljer tröskel utifrån observerade normer
- Vet när klippning döljer ett annat problem
Logga in för att göra övningarna och bygga upp din behärskning.
Källor
- arXiv — On the difficulty of training Recurrent Neural Networks — arXiv (öppen åtkomst; licens per artikel)
- PyTorch — tutorials (BSD-3) — BSD-3-Clause
- Dive into Deep Learning (CC BY-SA 4.0) — CC BY-SA 4.0