Hoppa till innehållet
AI-grafen
E· Universitetdeep-learning· ca 60 min· utvecklande· verifierad 2026-09-20

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:

g←g⋅min⁡(1,τ∥g∥)g \leftarrow g \cdot \min\left(1, \frac{\tau}{\|g\|}\right)

Riktningen bevaras — bara längden kortas. Det är den viktiga detaljen och skillnaden mot att klippa varje element för sig.

MetodGörBevarar riktning
Normklippningskalar hela gradientenja
Värdeklippningklipper varje element till [−c, c]nej

Normklippning (clip_grad_norm_) är nästan alltid rätt val.

Formellt

Var det behövs:

ModellTypisk tröskelVarför
RNN och LSTM1,0–5,0långa beroendekedjor ger multiplikativ tillväxt
Transformer1,0standard i nästan alla recept
RL (PPO)0,5belöningar kan variera kraftigt
CNNofta onödigtnormaliseringslagren 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:

SymtomVerkligt problem
Klipps varje steglr för hög, eller tröskeln för snäv
Normen växer stadigtinstabil arkitektur, saknad normalisering
Plötsliga toppar på enstaka batchardålig data — avvikare, felmärkning, extremvärden
NaN trots klippningdivision 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

Alla källor och licenser