Hoppa till innehållet
AI-grafen
G· Frontier Labinterpretability· ca 120 min· volatil — kontrolleras ofta· verifierad 2026-09-21

Sparse autoencoders för features

Kunna träna en SAE på aktiveringar och tolka features.

Förkunskaper

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:

f=ReLU(We(x−bd)+be),x^=Wdf+bdf = \mathrm{ReLU}(W_e(x - b_d) + b_e), \qquad \hat{x} = W_d f + b_d

L=∥x−x^∥22⏟rekonstruktion+λ∥f∥1⏟gleshet\mathcal{L} = \underbrace{\|x - \hat{x}\|_2^2}_{\text{rekonstruktion}} + \lambda \underbrace{\|f\|_1}_{\text{gleshet}}

Tre implementationsdetaljer som avgör om det fungerar:

  1. Normalisera decoderns kolumner till enhetsnorm. Annars kan nätet kringgå L1-straffet genom att göra ff litet och WdW_d stort.
  2. Subtrahera bdb_d före encodern. Det centrerar indata och förbättrar rekonstruktionen märkbart.
  3. 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 kk största aktiveringarna och sätter resten till noll. Det ger direkt kontroll över glesheten och tar bort λ\lambda-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åttMäter
Rekonstruktionsförlusthur mycket information som bevaras
L0 (antal aktiva)faktisk gleshet, typiskt 20–100
Förlorad korsentropiersätt aktiveringen med rekonstruktionen i modellen och mät hur mycket sämre den blir
Tolkningsbarhetkan 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ändningInnebörd
Ingen grundsanningvi vet inte vilka «verkliga» features som finns
Featuresplitingett begrepp delas i flera features när bredden ökar
Rekonstruktionen är ofullständigen del av modellens beteende fångas inte
Korrelation, inte orsakatt 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

Alla källor och licenser