Sparse autoencoders för features
Kunna träna en SAE på aktiveringar och tolka features.
Förkunskaper
- EAutoencoderskrävs
- FAktiveringar och linjära sonderkrävs
Intuition
En modells neuroner är polysemantiska: samma neuron reagerar på kattansikten, bilfronter och bokstaven A. Det gör dem nästan omöjliga att tolka.
Förklaringen är superposition: modellen behöver representera fler begrepp än den har dimensioner, och packar därför in flera per riktning. Det fungerar eftersom de flesta begrepp är sällsynta och sällan förekommer samtidigt.
Glesa autoencoders vänder på flaskhalsen. En vanlig autoencoder komprimerar till färre dimensioner. En SAE expanderar till fler — 8 till 64 gånger så många — men tvingar representationen att vara gles: bara en handfull är aktiva åt gången.
aktivering (512) ──encoder──→ features (16 384, ~20 aktiva) ──decoder──→ rekonstruktion (512)
Idén är att om modellen internt representerar tusentals begrepp i superposition, så kan en tillräckligt bred och gles bas plocka isär dem — och då blir varje riktning monosemantisk.
Formellt
Arkitektur och förlust. Encodern är ett linjärt lager med ReLU, decodern ett linjärt lager utan olinjäritet:
Tre implementationsdetaljer som avgör om det fungerar:
- Normalisera decoderns kolumner till enhetsnorm. Annars kan nätet kringgå L1-straffet genom att göra litet och stort.
- Subtrahera före encodern. Det centrerar indata och förbättrar rekonstruktionen märkbart.
- Hantera döda features. Många features slutar aktiveras helt. Motmedel: «resampling» — återinitiera döda features mot exempel med hög rekonstruktionsförlust.
Nyare varianter ersätter L1 med en hård gränsning: top-k SAE behåller bara de största aktiveringarna och sätter resten till noll. Det ger direkt kontroll över glesheten och tar bort -inställningen, som annars är känslig. JumpReLU och gated SAE är andra varianter som minskar den systematiska underskattning L1 orsakar.
Utvärdering — fyra mått, och inget av dem räcker ensamt:
| Mått | Mäter |
|---|---|
| Rekonstruktionsförlust | hur mycket information som bevaras |
| L0 (antal aktiva) | faktisk gleshet, typiskt 20–100 |
| Förlorad korsentropi | ersätt aktiveringen med rekonstruktionen i modellen och mät hur mycket sämre den blir |
| Tolkningsbarhet | kan en förklaring av featuren förutsäga när den aktiveras? |
Det tredje måttet är det ärligaste: en SAE med låg rekonstruktionsförlust men stor CE-förlust har missat just det som modellen faktiskt använder.
Den kritiska hållningen. SAE:er ger vackra, till synes tolkbara features — och det finns skäl till försiktighet:
| Invändning | Innebörd |
|---|---|
| Ingen grundsanning | vi vet inte vilka «verkliga» features som finns |
| Featurespliting | ett begrepp delas i flera features när bredden ökar |
| Rekonstruktionen är ofullständig | en del av modellens beteende fångas inte |
| Korrelation, inte orsak | att en feature aktiveras betyder inte att den används |
Den sista kräver samma botemedel som alltid i interpretability: intervention. Klampa featuren till noll eller till ett högt värde och mät om modellens beteende ändras. Utan det har man beskrivit ett mönster, inte en mekanism.
Vad SAE:er ändå gett: de har gjort det möjligt att hitta och namnge tiotusentals features i produktionsmodeller, och att styra beteende genom att förstärka eller dämpa enskilda riktningar. Det är ett verkligt framsteg — men fältet är ungt, och de negativa resultaten är lika viktiga att läsa som de positiva.
Kod
import torch, torch.nn as nn, torch.nn.functional as F
class SAE(nn.Module):
def __init__(self, d_modell=512, expansion=32, top_k=None):
super().__init__()
d_sae = d_modell * expansion
self.top_k = top_k
self.b_dec = nn.Parameter(torch.zeros(d_modell))
self.W_enc = nn.Parameter(torch.empty(d_sae, d_modell))
self.b_enc = nn.Parameter(torch.zeros(d_sae))
self.W_dec = nn.Parameter(torch.empty(d_modell, d_sae))
nn.init.kaiming_uniform_(self.W_enc)
self.W_dec.data = self.W_enc.data.T.clone()
self.normalisera_decoder()
@torch.no_grad()
def normalisera_decoder(self):
self.W_dec.data /= self.W_dec.data.norm(dim=0, keepdim=True) + 1e-8
def koda(self, x):
f = F.relu((x - self.b_dec) @ self.W_enc.T + self.b_enc)
if self.top_k: # top-k: hård gleshet, ingen lambda
varden, index = f.topk(self.top_k, dim=-1)
f = torch.zeros_like(f).scatter_(-1, index, varden)
return f
def forward(self, x):
f = self.koda(x)
return f @ self.W_dec.T + self.b_dec, f
def sae_forlust(x, x_hatt, f, lam=5e-4, top_k=False):
rek = (x - x_hatt).pow(2).sum(-1).mean()
if top_k:
return rek, {"rek": float(rek), "L0": float((f > 0).sum(-1).float().mean())}
gleshet = f.abs().sum(-1).mean()
return rek + lam * gleshet, {"rek": float(rek), "L1": float(gleshet),
"L0": float((f > 0).sum(-1).float().mean())}
# Det ärligaste måttet: hur mycket tappar modellen om man ersätter aktiveringen?
@torch.no_grad()
def forlorad_korsentropi(modell, sae, lager, dataloader):
def krok(m, i, o):
return sae(o)[0]
ce_normal = utvardera_ce(modell, dataloader)
h = lager.register_forward_hook(krok)
ce_sae = utvardera_ce(modell, dataloader)
h.remove()
# nollablation som referenspunkt
h0 = lager.register_forward_hook(lambda m, i, o: torch.zeros_like(o))
ce_noll = utvardera_ce(modell, dataloader)
h0.remove()
return {"ce_normal": round(ce_normal, 4), "ce_med_sae": round(ce_sae, 4),
"ce_nollad": round(ce_noll, 4),
"andel_bevarad": round((ce_noll - ce_sae) / (ce_noll - ce_normal), 4)}
# andel_bevarad nära 1.0 → SAE:n fångar nästan allt modellen använder
# Döda features: räkna och återinitiera
@torch.no_grad()
def doda_features(sae, dataloader, steg=200):
aktiv = None
for i, x in enumerate(dataloader):
if i >= steg:
break
f = sae.koda(x)
a = (f > 0).any(0)
aktiv = a if aktiv is None else (aktiv | a)
return {"doda": int((~aktiv).sum()), "av": len(aktiv),
"andel": round(float((~aktiv).float().mean()), 4)}
# Intervention: korrelation räcker inte
@torch.no_grad()
def klampa_feature(modell, sae, lager, feature, varde, indata):
def krok(m, i, o):
f = sae.koda(o)
f[..., feature] = varde
return f @ sae.W_dec.T + sae.b_dec
h = lager.register_forward_hook(krok)
ut = modell(indata)
h.remove()
return ut
# Ändras beteendet när featuren klampas? Annars är den bara korrelerad.
Behärskning innebär
- Förklarar superposition och varför gleshet hjälper
- Tränar en SAE på aktiveringar
- Utvärderar features kritiskt
Logga in för att göra övningarna och bygga upp din behärskning.
Källor
- Anthropic — Towards Monosemanticity (Transformer Circuits) — fri läsning
- arXiv — Scaling and evaluating sparse autoencoders — arXiv (öppen åtkomst; licens per artikel)
- arXiv — Toy Models of Superposition — arXiv (öppen åtkomst; licens per artikel)