psikosen's picture
Update to v5: Prefix Sliding KV-cache, SMELT scaling, CMA checkpoint, sPTC tool caller
2976bd9 verified
Raw History Blame Contribute Delete
6.26 kB
"""
Canopy-R3 Sparse Autoencoder (SAE) Engine for Mechanistic Interpretability.
Directly adapted from "Dissecting Hierarchical Reasoning Models: A Mechanistic Study" (MBZUAI, 2026).
Implements:
1. Top-K and L1 Sparse Autoencoders with empirical mean-centering (arXiv:2605.31518).
2. Decoder column unit-norm normalization.
3. Activation harvesting hooks for the Thought Bus and recurrent MoE layers.
4. Causal feature ablation tools to test whether internal reasoning features are localized or distributed.
"""
from typing import Dict, Tuple, Optional, List, Union
import torch
import torch.nn as nn
import torch.nn.functional as F
class CanopySparseAutoencoder(nn.Module):
"""
Sparse Autoencoder (SAE) for discovering interpretable latent features in Canopy-R3.
Supports Top-K exact structural sparsity and L1 regularization with mean-centering.
"""
def __init__(
self,
input_dim: int = 768,
dict_size: int = 2048,
mode: str = "topk",
top_k: int = 32,
l1_coeff: float = 0.01,
):
super().__init__()
self.input_dim = input_dim
self.dict_size = dict_size
self.mode = mode.lower()
self.top_k = top_k
self.l1_coeff = l1_coeff
if self.mode not in ("topk", "l1"):
raise ValueError(f"Unknown mode '{mode}'. Choose 'topk' or 'l1'.")
self.encoder = nn.Linear(input_dim, dict_size)
self.decoder = nn.Linear(dict_size, input_dim, bias=False)
self.pre_bias = nn.Parameter(torch.zeros(input_dim))
self.register_buffer("act_mean", torch.zeros(input_dim))
# Initializations
nn.init.kaiming_uniform_(self.encoder.weight, nonlinearity="relu")
nn.init.zeros_(self.encoder.bias)
nn.init.orthogonal_(self.decoder.weight)
self.normalize_decoder()
@torch.no_grad()
def normalize_decoder(self):
"""Projects decoder weight columns to unit norm."""
norms = self.decoder.weight.data.norm(p=2, dim=0, keepdim=True).clamp(min=1e-8)
self.decoder.weight.data.div_(norms)
def set_mean(self, mean: torch.Tensor):
"""Sets the empirical activation mean to eliminate dimension-level shift."""
self.act_mean.copy_(mean.squeeze().detach())
def encode(self, x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
"""Encodes activations into overcomplete sparse feature representations."""
x_centered = x - self.act_mean - self.pre_bias
pre_act = self.encoder(x_centered)
if self.mode == "topk":
# Exact top-K structural sparsity
k = min(self.top_k, self.dict_size)
topk_vals, topk_idx = torch.topk(F.relu(pre_act), k=k, dim=-1)
sparse_act = torch.zeros_like(pre_act).scatter_(-1, topk_idx, topk_vals)
else:
# L1 mode
sparse_act = F.relu(pre_act)
return sparse_act, pre_act
def decode(self, h: torch.Tensor) -> torch.Tensor:
"""Decodes sparse features back to original activation space."""
return self.decoder(h) + self.pre_bias + self.act_mean
def forward(self, x: torch.Tensor) -> Dict[str, torch.Tensor]:
orig_shape = x.shape
x_flat = x.view(-1, self.input_dim)
sparse_act, pre_act = self.encode(x_flat)
x_reconstructed = self.decode(sparse_act)
mse_loss = F.mse_loss(x_reconstructed, x_flat)
if self.mode == "l1":
l1_loss = self.l1_coeff * sparse_act.abs().mean()
total_loss = mse_loss + l1_loss
else:
l1_loss = torch.tensor(0.0, device=x.device)
total_loss = mse_loss
# Reconstruction metrics
with torch.no_grad():
var_x = (x_flat - x_flat.mean(dim=0)).pow(2).sum()
fvuda = (x_flat - x_reconstructed).pow(2).sum() / (var_x.clamp(min=1e-8))
explained_variance = (1.0 - fvuda).clamp(0.0, 1.0)
l0_norm = (sparse_act > 0).float().sum(dim=-1).mean()
return {
"loss": total_loss,
"mse_loss": mse_loss,
"l1_loss": l1_loss,
"explained_variance": explained_variance,
"l0_norm": l0_norm,
"reconstructed": x_reconstructed.view(orig_shape),
"sparse_act": sparse_act.view(*orig_shape[:-1], self.dict_size),
}
@torch.no_grad()
def ablate_features(
self,
x: torch.Tensor,
features_to_zero: Union[List[int], torch.Tensor],
) -> torch.Tensor:
"""
Causally ablates a set of dictionary features by zeroing them out during reconstruction.
"""
orig_shape = x.shape
x_flat = x.view(-1, self.input_dim)
sparse_act, _ = self.encode(x_flat)
if isinstance(features_to_zero, list):
features_to_zero = torch.tensor(features_to_zero, device=x.device)
sparse_act.index_fill_(-1, features_to_zero, 0.0)
ablated_recon = self.decode(sparse_act)
return ablated_recon.view(orig_shape)
class SAEActivationRecorder:
"""Context manager and hook utility to harvest activations from Canopy layers."""
def __init__(self, target_module: nn.Module):
self.target_module = target_module
self.activations: List[torch.Tensor] = []
self._hook_handle = None
def _hook(self, module, inputs, outputs):
if isinstance(outputs, tuple):
act = outputs[0]
elif isinstance(outputs, dict):
act = outputs.get("hidden_states", outputs.get("x", None))
else:
act = outputs
if act is not None and isinstance(act, torch.Tensor):
self.activations.append(act.detach().cpu())
def __enter__(self):
self.activations.clear()
self._hook_handle = self.target_module.register_forward_hook(self._hook)
return self
def __exit__(self, exc_type, exc_val, exc_tb):
if self._hook_handle is not None:
self._hook_handle.remove()
self._hook_handle = None
def get_concatenated(self) -> Optional[torch.Tensor]:
if not self.activations:
return None
return torch.cat([a.view(-1, a.shape[-1]) for a in self.activations], dim=0)