All articles AI Safety & Interpretability

Mechanistic Interpretability & Sparse Autoencoders: Deconstructing LLM Polysemanticity

Extracting monosemantic features from LLM activations using dictionary learning, Sparse Autoencoders (SAEs), and circuit probing.

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

  1. Monosemantic Concept Steering: Clamping a specific SAE feature (e.g. "Deception Feature #412") directly controls whether the model generates deceptive responses.
  2. Safety Auditing: Scans intermediate activations for dangerous features (e.g. CBRN weapon synthesis) before token generation completes.
  3. Circuit Probing: Traces causal information flow pathways between attention heads and multi-layer perceptrons.
← Back to all articles