Large Language Model hidden states suffer from polysemanticity: individual neurons do not represent single concepts; instead, a single neuron fires for completely unrelated concepts (e.g. both "Golden Gate Bridge" and "Python programming"). This occurs because neural networks pack millions of concepts into lower-dimensional vector spaces via superposition. Sparse Autoencoders (SAEs) unpack these dense activations into millions of monosemantic, human-interpretable features.
Individual neurons are polysemantic mixtures; Sparse Autoencoders disentangle superposed activations into clean, monosemantic concepts.
The superposition hypothesis & polysemanticity
Anthropic's Mechanistic Interpretability team (Elhage et al., 2022) formalized the Superposition Hypothesis: when the number of features in the real world exceeds the dimension of a hidden layer d_model, neural networks represent features as non-orthogonal directions in high-dimensional space.
Because features are not aligned with standard coordinate axes, reading raw neuron activations directly provides little insight into what the model is actually computing.
Sparse Autoencoder (SAE) architecture
A Sparse Autoencoder takes an intermediate activation vector x of dimension d_model and projects it into a sparse, high-dimensional latent space f of dimension d_sae (where d_sae = 32 * d_model or 64 * d_model):
x (d_model = 4096) ---> Encoder (W_enc) ---> f (d_sae = 131,072) [Sparse Latent]
|
x_hat (Reconstructed) <--- Decoder (W_dec) <--------------+
The mathematical formulation enforces extreme sparsity via an L1 penalty or Top-K activation function so that only a tiny fraction (e.g. 50 out of 131,072 features) are active for any given token.
PyTorch Sparse Autoencoder (SAE) implementation
import torch
import torch.nn as nn
import torch.nn.functional as F
class SparseAutoencoder(nn.Module):
"""
Top-K Sparse Autoencoder (SAE) for LLM activation dictionary learning.
"""
def __init__(self, d_model: int = 4096, expansion_factor: int = 32, k_sparse: int = 64):
super().__init__()
self.d_sae = d_model * expansion_factor
self.k_sparse = k_sparse
# Encoder & Decoder matrices
self.W_enc = nn.Parameter(torch.randn(d_model, self.d_sae) * 0.02)
self.b_enc = nn.Parameter(torch.zeros(self.d_sae))
self.W_dec = nn.Parameter(torch.randn(self.d_sae, d_model) * 0.02)
self.b_dec = nn.Parameter(torch.zeros(d_model))
# Normalize decoder weights to unit norm
self.normalize_decoder()
@torch.no_grad()
def normalize_decoder(self):
self.W_dec.data = F.normalize(self.W_dec.data, dim=1)
def forward(self, x: torch.Tensor):
# x: (batch_size, d_model)
x_centered = x - self.b_dec
hidden = F.relu(x_centered @ self.W_enc + self.b_enc)
# Enforce Top-K sparsity constraint
topk_vals, topk_indices = torch.topk(hidden, self.k_sparse, dim=-1)
sparse_latents = torch.zeros_like(hidden).scatter_(-1, topk_indices, topk_vals)
# Reconstruction pass
x_reconstructed = sparse_latents @ self.W_dec + self.b_dec
# Loss: Reconstruction MSE
reconstruction_loss = F.mse_loss(x_reconstructed, x)
return x_reconstructed, sparse_latents, reconstruction_loss
# Example activation pass
sae = SparseAutoencoder(d_model=4096, expansion_factor=32, k_sparse=64)
raw_activations = torch.randn(4, 4096)
x_hat, latents, mse_loss = sae(raw_activations)
print(f"Reconstruction MSE: {mse_loss.item():.4f}")
print(f"Active Latents per Token: {(latents > 0).sum(dim=-1).tolist()}")
Applications for safety & alignment auditability
- Monosemantic Concept Steering: Clamping a specific SAE feature (e.g. "Deception Feature #412") directly controls whether the model generates deceptive responses.
- Safety Auditing: Scans intermediate activations for dangerous features (e.g. CBRN weapon synthesis) before token generation completes.
- Circuit Probing: Traces causal information flow pathways between attention heads and multi-layer perceptrons.