| |
|
|
| import torch |
| from huggingface_hub import PyTorchModelHubMixin |
| from torch.nn import Linear, Module, Parameter |
| from torch import Tensor |
|
|
| from .decoder import decode |
| from .types import CrosscoderOut, Stats, TopK, SAEOut |
| from .Standard import SAE_Template, SAEConfig |
| from .SAE_Wrapper import _disable_hooks |
| from typing import Any, Tuple, Dict, List, Optional |
|
|
| from .Utils import standardize, unit_norm_decoder |
| from transformer_lens.hook_points import HookedRootModule |
| from transformer_lens.hook_points import HookPoint |
|
|
| class TopKSAE(SAE_Template): |
| last_nonzero: torch.Tensor |
| """The number of steps since the latents have activated.""" |
|
|
| def __init__( |
| self, |
| d_in: int, |
| d_sae: int, |
| hook_names: list[str], |
| k: int, |
| dead_steps_threshold: int, |
| dead_threshold: float = 1e-3, |
| |
| auxk: int | None = 256, |
| standardize: bool = True, |
| **kwargs, |
| ) -> None: |
| """ |
| Args: |
| d_in (int): The number of inputs. |
| |
| d_sae (int): The number of latents. |
| |
| k (int): The number of largest latents to keep. |
| |
| dead_steps_threshold (int): The number of steps after which a latent is |
| flagged as dead during training. |
| |
| dead_threshold (float): The threshold for a latent to be considered |
| activated. Defaults to 1e-3. |
| |
| auxk (int | None): The number of dead latents with which to model the |
| reconstruction error. Defaults to 256. |
| |
| standardize (bool): Whether to standardize the inputs. Defaults to True. |
| """ |
|
|
| super().__init__( |
| d_in=d_in, |
| d_sae=d_sae, |
| hook_names=hook_names, |
| auxk=auxk, |
| dead_steps_threshold=dead_steps_threshold, |
| dead_threshold=dead_threshold, |
| standardize=standardize, |
| ) |
|
|
| self.k = k |
|
|
| def encode( |
| self, inputs: torch.Tensor |
| ) -> tuple[Tensor, TopK, TopK | None, Stats | None, torch.Tensor]: |
| inputs = self.hook_sae_input(inputs) |
| |
| stats = None |
| if self.cfg.standardize: |
| inputs, stats = standardize(inputs) |
|
|
| |
| hidden_pre = self.hook_sae_acts_pre(self.encoder.forward(inputs - self.pre_encoder_bias)) |
|
|
| |
| values, indices = torch.topk( |
| hidden_pre, |
| k=self.k, |
| sorted=False |
| ) |
| topk = TopK(torch.relu(values), indices) |
| |
| latents = torch.zeros_like(hidden_pre) |
| latents = self.hook_sae_acts_post(latents.scatter_(-1, topk.indices, topk.values)) |
|
|
| self.update_last_nonzero(topk, inputs.device) |
| dead, auxk = self.compute_dead_latents_and_auxk(hidden_pre) |
|
|
| return latents, topk, auxk, stats, dead |
|
|
| def decode(self, latents: Tensor, stats: Stats | None = None) -> torch.Tensor: |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| recons = (latents @ self.decoder.weight.T) + self.pre_encoder_bias |
| |
| if stats is not None: |
| recons = recons * stats.std + stats.mean |
| return self.hook_sae_recons(recons) |
|
|
| def forward_training(self, inputs: torch.Tensor) -> SAEOut: |
| latents, topk, auxk, stats, dead = self.encode(inputs) |
|
|
| recons = self.decode(latents, stats) |
|
|
| auxk_recons = None |
| if auxk is not None: |
| auxk_latents = torch.zeros_like(latents) |
| auxk_latents.scatter_( |
| -1, |
| auxk.indices, |
| torch.relu(auxk.values), |
| ) |
| auxk_recons = self.decode(auxk_latents) |
|
|
| |
| return SAEOut(topk, recons, auxk, auxk_recons, dead, 0) |
| |
| class TopKTranscoder(TopKSAE): |
| def __init__( |
| self, |
| d_in: int, |
| d_sae: int, |
| hook_names: list[str], |
| k: int, |
| dead_steps_threshold: int, |
| auxk: int | None = 256, |
| dead_threshold: float = 1e-3, |
| standardize: bool = True, |
| **kwargs |
| ): |
| assert len(hook_names) == 2, "TopKTranscoder requires exactly two hook names." |
| output_hook, input_hook = hook_names |
| super().__init__( |
| d_in=d_in, |
| d_sae=d_sae, |
| hook_names=[output_hook, input_hook], |
| k=k, |
| auxk=auxk, |
| dead_steps_threshold=dead_steps_threshold, |
| dead_threshold=dead_threshold, |
| standardize=standardize, |
| ) |
| self.input_hook = input_hook |
| self.output_hook = output_hook |
| |
| self.decoder_bias = Parameter(torch.zeros(d_in)) |
| |
| def decode(self, latents: Tensor, stats: Stats | None = None) -> torch.Tensor: |
| recons = (latents @ self.decoder.weight.T) + self.decoder_bias |
| return self.hook_sae_recons(recons) |
| |
| def forward( |
| self, |
| x: Any, |
| ) -> torch.Tensor: |
| ''' |
| Modify the forward pass to allow gradient flows through the error term. |
| ''' |
| if isinstance(x, torch.Tensor): |
| inputs = x |
| outputs = None |
| elif isinstance(x, Tuple): |
| inputs, outputs = x |
| assert isinstance(inputs, torch.Tensor), "Inputs must be a torch.Tensor." |
| assert isinstance(outputs, torch.Tensor), "Outputs must be a torch.Tensor." |
| else: |
| raise TypeError("Input must be a torch.Tensor or a tuple of (inputs, outputs).") |
| |
| latents, _, _, stats, _ = self.encode(inputs) |
| sae_out = self.decode(latents, stats) |
|
|
| if self.use_error_term and outputs is not None: |
| with torch.no_grad() if self.detach_error_term else torch.enable_grad(): |
| with _disable_hooks(self): |
| clead_sae, _, _, clean_stats, _ = self.encode(inputs) |
| input_reconstruct_clean = self.decode(clead_sae, clean_stats) |
| |
| if self.disable_error_grad: |
| |
| |
| with torch.no_grad(): |
| sae_error = self.hook_sae_error(outputs - input_reconstruct_clean) |
| else: |
| |
| |
| with torch.no_grad() if self.detach_error_term else torch.enable_grad(): |
| temp_error = (outputs - input_reconstruct_clean) |
| temp_error.requires_grad_() |
| |
| sae_error = self.hook_sae_error(temp_error) |
| sae_out = sae_out + sae_error |
| |
| return self.hook_sae_output(sae_out) |
| |
| @property |
| def b_dec(self) -> torch.Tensor: |
| """ |
| Returns the decoder bias. |
| """ |
| return self.decoder_bias |
| |
| class TopKCrosscoder(TopKSAE): |
| def __init__( |
| self, |
| d_in: int, |
| d_sae: int, |
| input_hook: str, |
| output_hooks: list[str], |
| k: int, |
| dead_steps_threshold: int, |
| auxk: int | None = 256, |
| dead_threshold: float = 1e-3, |
| standardize: bool = True, |
| **kwargs |
| ): |
| self.input_hook = input_hook |
| self.output_hooks = output_hooks |
| |
| super().__init__( |
| d_in=d_in, |
| d_sae=d_sae, |
| hook_names=output_hooks + [input_hook], |
| k=k, |
| auxk=auxk, |
| dead_steps_threshold=dead_steps_threshold, |
| dead_threshold=dead_threshold, |
| standardize=standardize, |
| ) |
| |
| self.crosscoder_decoders = torch.nn.ModuleList( |
| [Linear(d_sae, d_in, bias=True) for _ in range(len(output_hooks)-1)] |
| ) |
| |
| self.decoder_bias = Parameter(torch.zeros(d_in)) |
| |
| def decode(self, latents: Tensor, stats: Stats | None = None) -> torch.Tensor: |
| recons = (latents @ self.decoder.weight.T) + self.decoder_bias |
| return self.hook_sae_recons(recons) |
| |
| def crosscoder_decode( |
| self, |
| latents: Tensor, |
| ) -> List[torch.Tensor]: |
| """ |
| Decode the latents into a list of tensors, one for each output hook. |
| """ |
| recons = [] |
| for decoder in self.crosscoder_decoders: |
| recons.append((decoder(latents))) |
| |
| return recons |
| |
| def _forward( |
| self, |
| inputs: torch.Tensor, |
| outputs: torch.Tensor | None, |
| pv_cons: torch.Tensor | float, |
| sae_out: torch.Tensor, |
| ) -> torch.Tensor: |
| if self.use_error_term and outputs is not None: |
| with torch.no_grad() if self.detach_error_term else torch.enable_grad(): |
| with _disable_hooks(self): |
| clead_sae, _, _, clean_stats, _ = self.encode(inputs) |
| input_reconstruct_clean = self.decode(clead_sae, clean_stats) + pv_cons |
| |
| if self.disable_error_grad: |
| |
| |
| with torch.no_grad(): |
| sae_error = self.hook_sae_error(outputs - input_reconstruct_clean) |
| else: |
| |
| |
| with torch.no_grad() if self.detach_error_term else torch.enable_grad(): |
| temp_error = (outputs - input_reconstruct_clean) |
| temp_error.requires_grad_() |
| |
| sae_error = self.hook_sae_error(temp_error) |
| sae_out = sae_out + sae_error |
| |
| return self.hook_sae_output(sae_out) |
| |
| def forward( |
| self, |
| x: Any, |
| ) -> torch.Tensor: |
| inputs, outputs, pv_cons = self.check_input(x) |
| |
| latents, _, _, stats, _ = self.encode(inputs) |
| sae_out = self.decode(latents, stats) + pv_cons |
|
|
| return self._forward(inputs, outputs, pv_cons, sae_out) |
| |
| def forward_crosscoder( |
| self, |
| x: Any, |
| ) -> Tuple[Tensor, List[Tensor]]: |
| inputs, outputs, pv_cons = self.check_input(x) |
| |
| latents, _, _, stats, _ = self.encode(inputs) |
| sae_out = self.decode(latents, stats) + pv_cons |
|
|
| output = self._forward(inputs, outputs, pv_cons, sae_out) |
| cross_outputs = self.crosscoder_decode(latents) |
| |
| return output, cross_outputs |
| |
| def forward_training(self, inputs: Any) -> CrosscoderOut: |
| if isinstance(inputs, Tuple): |
| inputs, pv_cons = inputs |
| assert isinstance(inputs, torch.Tensor), "Inputs must be a torch.Tensor." |
| assert isinstance(pv_cons, torch.Tensor) or isinstance(pv_cons, float), "Previous layer construction must be a torch.Tensor." |
| else: |
| raise TypeError("Input must be a tuple of (inputs, previous_layer_construction).") |
| latents, topk, auxk, stats, dead = self.encode(inputs) |
|
|
| recons = self.decode(latents, stats) + pv_cons |
| |
| cross_recons = self.crosscoder_decode(latents) |
|
|
| auxk_recons = None |
| if auxk is not None: |
| auxk_latents = torch.zeros_like(latents) |
| auxk_latents.scatter_( |
| -1, |
| auxk.indices, |
| torch.relu(auxk.values), |
| ) |
| auxk_recons = self.decode(auxk_latents) |
|
|
| |
| return CrosscoderOut(topk, recons, cross_recons, auxk, auxk_recons, dead, 0) |
| |
| def check_input( |
| self, |
| x: Any, |
| ) -> Tuple[torch.Tensor, torch.Tensor | None, torch.Tensor | float]: |
| """ |
| Check the input type and return the inputs, outputs, and previous layer construction. |
| """ |
| if isinstance(x, Tuple): |
| if len(x) == 2: |
| inputs, pv_cons = x |
| outputs = None |
| elif len(x) == 3: |
| inputs, outputs, pv_cons = x |
| else: |
| raise ValueError("Input tuple must be of length 2 or 3.") |
| assert isinstance(inputs, torch.Tensor), "Inputs must be a torch.Tensor." |
| assert isinstance(outputs, torch.Tensor) or outputs is None, "Outputs must be a torch.Tensor or None." |
| assert isinstance(pv_cons, torch.Tensor) or isinstance(pv_cons, float), "Previous layer construction must be a torch.Tensor." |
| return inputs, outputs, pv_cons |
| else: |
| raise TypeError("Input must be a tuple of (inputs, outputs, previous_layer_construction).") |