""" 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)