Skip to content
AI-grafen
GFrontier LabInterpretability· about 120 min· fast-moving, sources checked often· verified 2026-09-21· EN

Sparse autoencoders for features

Be able to train an SAE on activations and interpret the features.

Prerequisites

Intuition

A model's neurons are polysemantic: the same neuron responds to cat faces, car fronts and the letter A. That makes them nearly impossible to interpret.

The explanation is superposition: the model needs to represent more concepts than it has dimensions, and therefore packs several into each direction. That works because most concepts are rare and rarely occur at the same time.

Sparse autoencoders turn the bottleneck around. An ordinary autoencoder compresses to fewer dimensions. An SAE expands to more — 8 to 64 times as many — but forces the representation to be sparse: only a handful are active at a time.

activation (512) ──encoder──→ features (16 384, ~20 active) ──decoder──→ reconstruction (512)

The idea is that if the model internally represents thousands of concepts in superposition, then a sufficiently wide and sparse basis can take them apart — and then every direction becomes monosemantic.

Formal

The architecture and the loss. The encoder is a linear layer with a ReLU, the decoder a linear layer without a non-linearity:

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⏟reconstruction+λ∥f∥1⏟sparsity\mathcal{L} = \underbrace{\|x - \hat{x}\|_2^2}_{\text{reconstruction}} + \lambda \underbrace{\|f\|_1}_{\text{sparsity}}

Three implementation details that decide whether it works:

  1. Normalise the decoder's columns to unit norm. Otherwise the network can get around the L1 penalty by making ff small and WdW_d large.
  2. Subtract bdb_d before the encoder. That centres the input and improves the reconstruction noticeably.
  3. Handle dead features. Many features stop activating entirely. The remedy: «resampling» — reinitialise the dead features towards examples with a high reconstruction loss.

Newer variants replace L1 with a hard cut-off: a top-k SAE keeps only the kk largest activations and sets the rest to zero. That gives direct control over the sparsity and removes the λ\lambda setting, which is otherwise sensitive. JumpReLU and gated SAEs are other variants that reduce the systematic underestimation L1 causes.

Evaluation — four metrics, and none of them is enough on its own:

MetricMeasures
The reconstruction losshow much information is preserved
L0 (the number active)the actual sparsity, typically 20–100
The lost cross-entropyreplace the activation with the reconstruction in the model and measure how much worse it gets
Interpretabilitycan an explanation of the feature predict when it activates?

The third metric is the most honest: an SAE with a low reconstruction loss but a large CE loss has missed precisely what the model actually uses.

The critical stance. SAEs give beautiful, apparently interpretable features — and there are reasons for caution:

ObjectionThe implication
No ground truthwe do not know which «real» features exist
Feature splittingone concept is split into several features as the width increases
The reconstruction is incompletepart of the model's behaviour is not captured
Correlation, not causethat a feature activates does not mean it is used

The last requires the same remedy as always in interpretability: intervention. Clamp the feature to zero or to a high value and measure whether the model's behaviour changes. Without that you have described a pattern, not a mechanism.

What SAEs have given anyway: they have made it possible to find and name tens of thousands of features in production models, and to steer behaviour by amplifying or damping individual directions. That is a real advance — but the field is young, and the negative results are as important to read as the positive ones.

Code

import torch, torch.nn as nn, torch.nn.functional as F

class SAE(nn.Module):
    def __init__(self, d_model=512, expansion=32, top_k=None):
        super().__init__()
        d_sae = d_model * expansion
        self.top_k = top_k
        self.b_dec = nn.Parameter(torch.zeros(d_model))
        self.W_enc = nn.Parameter(torch.empty(d_sae, d_model))
        self.b_enc = nn.Parameter(torch.zeros(d_sae))
        self.W_dec = nn.Parameter(torch.empty(d_model, d_sae))
        nn.init.kaiming_uniform_(self.W_enc)
        self.W_dec.data = self.W_enc.data.T.clone()
        self.normalise_decoder()

    @torch.no_grad()
    def normalise_decoder(self):
        self.W_dec.data /= self.W_dec.data.norm(dim=0, keepdim=True) + 1e-8

    def encode(self, x):
        f = F.relu((x - self.b_dec) @ self.W_enc.T + self.b_enc)
        if self.top_k:                          # top-k: hard sparsity, no lambda
            values, index = f.topk(self.top_k, dim=-1)
            f = torch.zeros_like(f).scatter_(-1, index, values)
        return f

    def forward(self, x):
        f = self.encode(x)
        return f @ self.W_dec.T + self.b_dec, f

def sae_loss(x, x_hat, f, lam=5e-4, top_k=False):
    rec = (x - x_hat).pow(2).sum(-1).mean()
    if top_k:
        return rec, {"rec": float(rec), "L0": float((f > 0).sum(-1).float().mean())}
    sparsity = f.abs().sum(-1).mean()
    return rec + lam * sparsity, {"rec": float(rec), "L1": float(sparsity),
                                  "L0": float((f > 0).sum(-1).float().mean())}

# The most honest metric: how much does the model lose if the activation is replaced?
@torch.no_grad()
def lost_cross_entropy(model, sae, layer, dataloader):
    def hook(m, i, o):
        return sae(o)[0]
    ce_normal = evaluate_ce(model, dataloader)
    h = layer.register_forward_hook(hook)
    ce_sae = evaluate_ce(model, dataloader)
    h.remove()
    # zero ablation as a reference point
    h0 = layer.register_forward_hook(lambda m, i, o: torch.zeros_like(o))
    ce_zeroed = evaluate_ce(model, dataloader)
    h0.remove()
    return {"ce_normal": round(ce_normal, 4), "ce_with_sae": round(ce_sae, 4),
            "ce_zeroed": round(ce_zeroed, 4),
            "share_preserved": round((ce_zeroed - ce_sae) / (ce_zeroed - ce_normal), 4)}
# share_preserved near 1.0 → the SAE captures nearly everything the model uses

# Dead features: count them and reinitialise
@torch.no_grad()
def dead_features(sae, dataloader, steps=200):
    active = None
    for i, x in enumerate(dataloader):
        if i >= steps:
            break
        f = sae.encode(x)
        a = (f > 0).any(0)
        active = a if active is None else (active | a)
    return {"dead": int((~active).sum()), "out_of": len(active),
            "share": round(float((~active).float().mean()), 4)}

# An intervention: correlation is not enough
@torch.no_grad()
def clamp_feature(model, sae, layer, feature, value, inputs):
    def hook(m, i, o):
        f = sae.encode(o)
        f[..., feature] = value
        return f @ sae.W_dec.T + sae.b_dec
    h = layer.register_forward_hook(hook)
    out = model(inputs)
    h.remove()
    return out
# Does the behaviour change when the feature is clamped? Otherwise it is merely correlated.

Mastery means

  • Explains superposition and why sparsity helps
  • Trains an SAE on activations
  • Evaluates the features critically

Sign in to do the exercises and build your mastery up.

Sources

All the sources and licences