Sparse autoencoders for features
Be able to train an SAE on activations and interpret the features.
Prerequisites
- EAutoencodersrequired
- FActivations and linear probesrequired
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:
Three implementation details that decide whether it works:
- Normalise the decoder's columns to unit norm. Otherwise the network can get around the L1 penalty by making small and large.
- Subtract before the encoder. That centres the input and improves the reconstruction noticeably.
- 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 largest activations and sets the rest to zero. That gives direct control over the sparsity and removes the 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:
| Metric | Measures |
|---|---|
| The reconstruction loss | how much information is preserved |
| L0 (the number active) | the actual sparsity, typically 20–100 |
| The lost cross-entropy | replace the activation with the reconstruction in the model and measure how much worse it gets |
| Interpretability | can 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:
| Objection | The implication |
|---|---|
| No ground truth | we do not know which «real» features exist |
| Feature splitting | one concept is split into several features as the width increases |
| The reconstruction is incomplete | part of the model's behaviour is not captured |
| Correlation, not cause | that 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
- Anthropic — Towards Monosemanticity (Transformer Circuits) — free to read
- arXiv — Scaling and evaluating sparse autoencoders — arXiv (open access; licence per article)
- arXiv — Toy Models of Superposition — arXiv (open access; licence per article)