Looking Inside a Trained Model
Sparse autoencoders for feature discovery
A sparse autoencoder unpacks a model's crowded, overlapping neurons into a wider set of directions, each closer to one clean, single concept.
- 10 min read
- 3 reading levels
- Published
Read these first
On this page 5
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
A sparse autoencoder sorts a model's crowded numbers into a wider set of slots. Each slot holds closer to one clean idea.
Picture a junk drawer so full several unrelated items share one compartment. A charger, a rubber band and spare keys, all tangled together. You cannot grab the keys alone.
Now picture emptying that drawer into a bigger organiser, one item type per slot. A sparse autoencoder does this to a model's crowded neurons, spreading tangled concepts into cleaner, separate directions.
Why it exists
The previous lesson showed why neurons get crowded. A model has far more concepts than neurons, so it shares space between concepts that rarely occur together. That crowding is called superposition.
Crowded neurons are hard to interpret directly. A single neuron firing tells you little. It might be tracking several unrelated things at once, mixed together in one number.
A sparse autoencoder was built to undo that mixing, at least approximately. Given only a model's raw numbers, no labels, it learns a wider set of directions. Ideally, each one tracks a single, cleaner idea.
How it works
Model's crowded numbers: [ 0.4, -0.2, 0.9, 0.1, -0.5 ]
(5 numbers, several ideas mixed together)
Sparse autoencoder expands this into many more numbers,
most of them exactly zero for any given input:
[ 0, 0, 0.8, 0, 0, 0, 0, 0.3, 0, 0, 0, ... ]
(a few active slots, each closer to one idea)Only a few slots switch on for any single input, hence "sparse". The hope is that each slot, when it does switch on, corresponds to something a person could recognise and name.
Where you have already seen it
- AI company research blogs. A model "having a feature" for a specific concept, like a city or a coding pattern, usually describes an SAE-discovered direction.
- Model steering tools. Once a clean feature direction is found, some products let users push it up or down to change behaviour.
- AI safety monitoring. Watching for a specific concerning feature, like deception-related directions, to switch on during generation.
Remember this
- A sparse autoencoder expands a model's crowded numbers into a wider, mostly-off set of directions.
- The goal is untangling mixed-up concepts into cleaner, more separable ones.
- It needs no labelled data, only the model's own raw activity to train on.
What to learn next
- Superposition and polysemantic neurons — the crowding problem this technique tries to undo.
- Steering a model with activation vectors — using a discovered feature direction on purpose.
- Where a model stores a fact — a specific, well-studied case of what these directions can represent.
Developer — Code and libraries.
This builds a small sparse autoencoder on top of the toy superposed model from the previous lesson, then checks how well its learned directions line up with the true underlying features.
Setup
pip install torchTraining a tiny sparse autoencoder
import torch
torch.manual_seed(0)
n_features, n_hidden, sparsity = 20, 5, 0.95
# Step 1: rebuild the superposed toy model from the previous lesson. It
# compresses 20 sparse features into 5 numbers, mixing several into each.
W = torch.randn(n_hidden, n_features, requires_grad=True)
opt = torch.optim.Adam([W], lr=0.02)
def make_batch(n=512):
x = torch.rand(n, n_features)
mask = (torch.rand(n, n_features) > sparsity).float()
return x * mask
for _ in range(2000):
x = make_batch()
hidden = x @ W.T
x_hat = torch.relu(hidden @ W)
loss = ((x - x_hat) ** 2).mean()
opt.zero_grad(); loss.backward(); opt.step()
W = W.detach()
# Step 2: a sparse autoencoder that sees only the 5-number compressed
# activations, and tries to re-expand them into a wider, sparse code.
latent_dim = 20 # deliberately as wide as the true number of features
enc = torch.nn.Linear(n_hidden, latent_dim)
dec = torch.nn.Linear(latent_dim, n_hidden, bias=False)
sae_opt = torch.optim.Adam(list(enc.parameters()) + list(dec.parameters()), lr=0.01)
l1_weight = 0.01 # how strongly to push the code toward sparsity
for step in range(3000):
x = make_batch()
with torch.no_grad():
hidden = x @ W.T # the superposed activation the SAE sees
latent = torch.relu(enc(hidden)) # wide, sparse code
recon = dec(latent)
loss = ((hidden - recon) ** 2).mean() + l1_weight * latent.abs().mean()
sae_opt.zero_grad(); loss.backward(); sae_opt.step()
# Step 3: does each learned latent line up with one true feature?
with torch.no_grad():
x_big = make_batch(4000)
hidden_big = x_big @ W.T
latent_big = torch.relu(enc(hidden_big))
print(f"fraction of latents active per example: {(latent_big > 0.01).float().mean(1).mean():.2f}")
best_corr = []
for f in range(n_features):
corrs = [torch.corrcoef(torch.stack([x_big[:, f], latent_big[:, j]]))[0, 1].nan_to_num()
for j in range(latent_dim)]
best_corr.append(max(corrs).item())
print(f"mean best correlation (true feature <-> a learned latent): {sum(best_corr)/len(best_corr):.3f}")fraction of latents active per example: 0.11 mean best correlation (true feature <-> a learned latent): 0.507
Roughly 11-12% of latents switch on per example, close to the true 5% activity rate the underlying features were generated with. Small floating-point differences between CPU runs mean this figure can shift a point or two, even with the seed fixed. Treat the general shape as the result, not the third decimal. The mean correlation, 0.507, is real but far from perfect. This is an honest result: this small, quickly trained SAE partly recovers the true features, not completely.
Line by line
latent_dim = 20 deliberately matches the true number of underlying features. Real sparse autoencoders often use a latent dimension many times larger than the input. The true number of features in a real model is unknown, and usually assumed to far exceed the neuron count.
The loss has two parts: reconstruction error, and 0.01 * latent.abs().mean(), the sparsity penalty. Without the second term, the autoencoder would happily use every latent for every input, defeating the entire purpose.
The correlation loop checks each true feature against every learned latent, keeping the single best match. A high average means the SAE's latents line up well with genuine, independent underlying features.
Common mistakes
Setting the sparsity penalty too low or too high. Too low, and latents stop being sparse, recreating the original crowding problem. Too high, and the SAE stops reconstructing anything useful at all, sacrificing accuracy for emptiness.
Assuming a high correlation score means a perfect feature match. As shown above, 0.507 is meaningfully above zero and still leaves real interference. Sparse autoencoders are an approximation, not a guaranteed exact decomposition.
Training on too little data. The SAE here trains on freshly sampled batches every step, which is generous. A real SAE trained on limited, non-refreshing activation data can overfit to quirks of that specific dataset instead of the true feature structure.
Try it yourself
Change l1_weight from 0.01 to 0.1, a much stronger sparsity penalty, and rerun. Compare the fraction of active latents and the mean correlation against the original run.
Expect fewer active latents per example, and likely a change, in either direction, to the correlation score. Finding the right sparsity strength is one of the real practical challenges in training a useful sparse autoencoder.
What to learn next
- Steering a model with activation vectors — putting a discovered feature direction to active use.
- Superposition and polysemantic neurons — the crowding phenomenon motivating this entire technique.
- Where a model stores a fact — a concrete example of the kind of feature an SAE might isolate.
Researcher — Mathematics and papers.
Architecture
A sparse autoencoder trained on activations a in R^d from some model layer learns:
z = ReLU(W_enc @ a + b_enc)
a_hat = W_dec @ z + b_dec
Loss = || a - a_hat ||^2 + lambda * || z ||_1z in R^kis the sparse latent code, typically withk >> d(an "overcomplete" dictionary), unlike the developer block's illustrativek = d * 4toy setting.W_dec's columns are the learned feature directions, or dictionary elements, each intended to correspond to one interpretable concept.lambdacontrols the sparsity-reconstruction trade-off, directly analogous to thel1_weightin the developer block.
This is dictionary learning, a decades-old technique from signal processing, applied to neural network activations. The novelty in recent interpretability work is scale: training these on activations from billions of tokens through frontier models.
The dead-latent and feature-splitting problems
Dead latents: a latent that never activates across the training distribution provides no information and wastes dictionary capacity. Standard mitigations include periodic resampling of dead latents' weights, and auxiliary loss terms that specifically penalise latents going permanently silent.
Feature splitting: increasing k does not only add new, unrelated features. It often splits one coarse feature into several more specific ones instead, for example a general "programming language" feature splitting into separate Python and JavaScript features. This makes "the correct dictionary size" an ill-defined question, since the right granularity depends on the intended use.
Evaluating a sparse autoencoder
Unlike supervised learning, there is no ground-truth feature set to check against in a real model. The toy setting in the developer block is the exception, where the true features are known by construction. Standard evaluation instead uses proxy metrics:
- Reconstruction loss and sparsity (L0), traded off directly against each other, plotted as a Pareto frontier across different
lambdavalues. - Interpretability of top-activating examples: for a given latent, do the inputs that activate it most strongly share a clear, nameable, human-recognisable theme?
- Downstream causal relevance: does patching or ablating a specific latent, in the manner of activation patching, produce a predictable, interpretable change in model behaviour?
Key references
- Sharkey, L., Braun, D. & Millidge, B. (2022). Taking Features Out of Superposition with Sparse Autoencoders. Alignment Forum — the initial proposal connecting SAEs to interpreting superposition.
- Bricken, T. et al. (2023). Towards Monosemanticity: Decomposing Language Models With Dictionary Learning. Anthropic. transformer-circuits.pub
- Cunningham, H. et al. (2023). Sparse Autoencoders Find Highly Interpretable Features in Language Models. arXiv:2309.08600
- Templeton, A. et al. (2024). Scaling Monosemanticity. Anthropic — SAEs trained at production model scale (Claude 3 Sonnet).
Current state and open problems
Sparse autoencoders are currently the leading practical technique for extracting human-interpretable features from superposed activations at real model scale, and have found genuinely novel, previously undocumented model behaviours this way.
Open problems include the dead-latent and feature-splitting issues above, and the lack of a principled way to choose dictionary size. A deeper question remains unsettled too: whether a model's "true" feature basis is even a well-defined, discoverable object, or whether interpretability here is inherently an approximation, with no single correct answer to converge toward.
What to learn next
- Steering a model with activation vectors — acting on a discovered feature direction, not only observing it.
- Activation patching — the causal-evaluation technique used to test whether a found latent actually matters.
- Where a model stores a fact — a well-studied target for exactly this kind of feature-level analysis.