""" SHARED UTILITIES FOR ADVERSARIAL SAE TRAINING ============================================================================= DO NOT RUN THIS FILE AS A TRAINING SCRIPT. ============================================================================= This module provides shared utilities used by: - finetune_adv_gen.py (generation-time feature suppression training) - finetune_adv_gen_resume.py (resumable version of the above) The prefill-only training loop has been removed. All actual training should use finetune_adv_gen.py or finetune_adv_gen_resume.py. Exported utilities: Classes: - FrozenSAEEncoder: Inference-only SAE encoder (frozen weights) - LayerProbes: Per-layer linear probes for SAE features - HiddenStateCapture: Hook-based activation capture context manager Loss functions: - suppression_loss: Push probe outputs toward target on suppress rows - probe_bce_loss_logits: BCE loss for probe training - probe_present_loss_logits: Encourage probe to fire on present rows - group_lasso_lora_by_layer: L1-over-L2 regularization on LoRA params - activation_retain_mse: MSE between current and base activations - kl_batchmean_masked: KL divergence with row masking - probe_labels: Compute probe targets from scene/object masks Helpers: - load_probe_checkpoint: Load probe weights from various checkpoint formats - count_lm_layers: Count decoder layers in a model - get_sae_features: Extract SAE features from captured activations - parse_args: CLI argument parser for training scripts - _Tee: Stream tee utility for logging - _stable_target_probs: Numerically stable probability normalization - _decoder_layers: Locate decoder layers in various model architectures """ from __future__ import annotations import re import torch import torch.nn as nn import torch.nn.functional as F # --------------------------------------------------------------------------- # Frozen SAE encoder # --------------------------------------------------------------------------- class FrozenSAEEncoder(nn.Module): """Inference-only mirror of BatchTopKSAE.encode(use_threshold=True). Per-row: standardize -> linear(x - pre_encoder_bias) -> JumpReLU(threshold). Returns sparse latents matching what the SAE was trained to produce; gradient flows through the active features only (the boolean gate stops grad on inactive features). """ def __init__(self, d_in: int, d_sae: int, standardize: bool = True): super().__init__() self.encoder = nn.Linear(d_in, d_sae, bias=False) self.pre_encoder_bias = nn.Parameter(torch.zeros(d_in)) self.register_buffer("threshold", torch.zeros(())) self.standardize = standardize for p in self.parameters(): p.requires_grad_(False) def forward(self, x: torch.Tensor) -> torch.Tensor: if self.standardize: mu = x.mean(dim=-1, keepdim=True) x = x - mu std = x.std(dim=-1, keepdim=True) x = x / (std + 1e-5) hidden_pre = self.encoder(x - self.pre_encoder_bias) gate = (F.relu(hidden_pre) > self.threshold).to(hidden_pre.dtype) return hidden_pre * gate @classmethod def from_checkpoint(cls, ckpt_path: str, device: torch.device) -> "FrozenSAEEncoder": ckpt = torch.load(ckpt_path, map_location="cpu", weights_only=False) sd = ckpt["state_dict"] enc_weight = sd["autoencoder.encoder.weight"] d_sae, d_in = enc_weight.shape hp = ckpt.get("hyper_parameters") or ckpt.get("hyperparameters") or {} if isinstance(hp, dict): standardize = bool(hp.get("standardize", True)) else: standardize = bool(getattr(hp, "standardize", True)) enc = cls(d_in, d_sae, standardize=standardize) enc.encoder.weight.data.copy_(enc_weight) enc.pre_encoder_bias.data.copy_(sd["autoencoder.pre_encoder_bias"]) if "autoencoder.threshold" in sd: enc.threshold.data.copy_(sd["autoencoder.threshold"].to(enc.threshold.dtype)) return enc.to(device) # --------------------------------------------------------------------------- # Per-layer linear probes # --------------------------------------------------------------------------- class LayerProbes(nn.Module): """One linear probe per monitored layer: sigma(Linear(f)); train probe with logits + BCEWithLogits.""" def __init__(self, layer_indices: list[int], d_sae: int, spectral_norm: bool = False): super().__init__() self.layer_indices = layer_indices if spectral_norm: self.probes = nn.ModuleList([ nn.utils.spectral_norm(nn.Linear(d_sae, 1)) for _ in layer_indices ]) else: self.probes = nn.ModuleList([nn.Linear(d_sae, 1) for _ in layer_indices]) self._idx = {l: i for i, l in enumerate(layer_indices)} def forward_logits(self, features: dict[int, torch.Tensor]) -> list[torch.Tensor]: return [self.probes[self._idx[l]](features[l]).squeeze(-1) for l in self.layer_indices] def forward(self, features: dict[int, torch.Tensor]) -> list[torch.Tensor]: return [torch.sigmoid(z) for z in self.forward_logits(features)] def layer_probes_from_checkpoint( path: str, layer_indices: list[int], d_sae: int, device: torch.device | str | None = None, ) -> LayerProbes: """Load probes.pt from training. Detects spectral-norm probes (weight_orig/weight_u/weight_v). Always loads weights on CPU then moves the module to device so parameters match the activations' device (avoids CPU weights + CUDA features after load_state_dict). """ sd = torch.load(path, map_location="cpu", weights_only=True) if any(k.startswith("module.") for k in sd): sd = {k.replace("module.", "", 1): v for k, v in sd.items()} use_sn = any("weight_orig" in k for k in sd) probes = LayerProbes(layer_indices, d_sae, spectral_norm=use_sn) probes.load_state_dict(sd, strict=True) if device is not None: dev = device if isinstance(device, torch.device) else torch.device(device) probes = probes.to(dev) return probes def load_probe_checkpoint(probes: LayerProbes, path: str) -> None: """Load weights from LayerProbes state_dict or Train_Probe_SAE LinearProbe (fc.*).""" sd = torch.load(path, map_location="cpu", weights_only=True) if any(k.startswith("module.") for k in sd): sd = {k.replace("module.", "", 1): v for k, v in sd.items()} if any(k.startswith("probes.") for k in sd): probes.load_state_dict(sd, strict=True) return if "fc.weight" in sd: w = sd["fc.weight"] b = sd.get("fc.bias") if b is None: b = torch.zeros(w.shape[0], device=w.device, dtype=w.dtype) for lin in probes.probes: lin.weight.data.copy_(w) lin.bias.data.copy_(b) return raise ValueError( f"Unrecognised probe checkpoint {path}: expected LayerProbes keys or fc.weight/fc.bias" ) # --------------------------------------------------------------------------- # Hook-based hidden state capture # --------------------------------------------------------------------------- def _decoder_layers(raw_model) -> nn.ModuleList: """LLaVA / Llama: locate the LlamaDecoderLayer ModuleList. Transformers has moved the language model around across versions: - older: raw_model.language_model.{model.layers|layers} - newer: raw_model.model.language_model.{layers|model.layers} """ lm = getattr(raw_model, "language_model", None) if lm is None: inner_model = getattr(raw_model, "model", None) if inner_model is not None: lm = getattr(inner_model, "language_model", None) if lm is None: raise AttributeError( "Cannot locate language_model under raw_model " "(checked raw_model.language_model and raw_model.model.language_model)" ) inner = getattr(lm, "model", None) if inner is not None and hasattr(inner, "layers"): return inner.layers if hasattr(lm, "layers"): return lm.layers raise AttributeError( "Cannot find Llama decoder layers: expected language_model.model.layers " "or language_model.layers" ) class HiddenStateCapture: def __init__(self, raw_model, layer_indices: list[int]): self._layers = _decoder_layers(raw_model) self.layer_indices = layer_indices self.hidden_states: dict[int, torch.Tensor] = {} self._hooks: list = [] def __enter__(self): self.hidden_states.clear() for l in self.layer_indices: def _hook(mod, inp, out, idx=l): h = out[0] if isinstance(out, tuple) else out self.hidden_states[idx] = h self._hooks.append(self._layers[l].register_forward_hook(_hook)) return self def __exit__(self, *_): for h in self._hooks: h.remove() self._hooks.clear() # --------------------------------------------------------------------------- # Loss helpers # --------------------------------------------------------------------------- def _stable_target_probs(p: torch.Tensor) -> torch.Tensor: """Clamp + renormalize so kl_div does not get 0 * log 0 -> NaN.""" p = p.float().clamp(min=1e-8) return p / p.sum(dim=-1, keepdim=True) def suppression_loss( probs_list: list[torch.Tensor], suppress_mask: torch.Tensor, mode: str = "union", ) -> torch.Tensor: """Suppression loss on suppress_mask (D_{A,B}), mean over layers. union mode (probe label = 1 for scene-only): probe was trained to fire on the A->B prior. Target = p=0.5 (max entropy): the model erases the prior so the probe is fully uncertain. Loss = -H(p) = p*log(p) + (1-p)*log(1-p), minimised at p=0.5. object_only mode (probe label = 0 for scene-only): probe was trained to detect actual B features. Target = p=0 (BCE toward 0): the model erases B-correlated features so the probe correctly sees no B information. Pushing to 0.5 would make activations ambiguous (worse, not better). Loss = BCE(p, 0) = -log(1-p). """ if not suppress_mask.any(): return probs_list[0].sum() * 0.0 eps = 1e-6 total = 0.0 for probs in probs_list: p = probs[suppress_mask].float().clamp(eps, 1 - eps) if mode == "union": h = -(p * p.log() + (1 - p) * (1 - p).log()) total = total + (-h.mean()) # minimise -H -> push to p=0.5 else: total = total + (-((1 - p).log()).mean()) # BCE toward 0 -> push to p=0 return total / len(probs_list) def probe_bce_loss_logits( logits_list: list[torch.Tensor], labels: torch.Tensor, label_smoothing: float = 0.0, ) -> torch.Tensor: labels = labels.float() if label_smoothing > 0.0: labels = labels * (1.0 - label_smoothing) + 0.5 * label_smoothing return sum( F.binary_cross_entropy_with_logits(z.float(), labels, reduction="mean") for z in logits_list ) / len(logits_list) def probe_present_loss_logits( logits_list: list[torch.Tensor], present_mask: torch.Tensor, target: float = 1.0, ) -> torch.Tensor: """Encourage probe to predict 'present' on D_B (has_object==1).""" if not present_mask.any(): return logits_list[0].sum() * 0.0 y = logits_list[0].new_full((int(present_mask.sum().item()),), float(target)) return sum( F.binary_cross_entropy_with_logits(z[present_mask].float(), y, reduction="mean") for z in logits_list ) / len(logits_list) _LAYER_RE = re.compile(r"\.layers\.(\d+)\.") def group_lasso_lora_by_layer(trainable_named_params) -> torch.Tensor: """L1(L2(DeltaW)) across layers: sum_l sqrt(sum_{p in layer l} ||p||_2^2). Uses parameter-name heuristic to map LoRA params to decoder layers. """ per_layer_sq = {} device = None for name, p in trainable_named_params: if not p.requires_grad: continue m = _LAYER_RE.search(name) if m is None: continue l = int(m.group(1)) if device is None: device = p.device per_layer_sq[l] = per_layer_sq.get(l, 0.0) + (p.float().pow(2).sum()) if not per_layer_sq: return torch.zeros((), device=device if device is not None else None) return torch.stack([(v + 1e-12).sqrt() for _, v in sorted(per_layer_sq.items())]).sum() def activation_retain_mse( h_current: dict[int, torch.Tensor], h_base: dict[int, torch.Tensor], retain_mask: torch.Tensor, ) -> torch.Tensor: if not retain_mask.any(): any_h = next(iter(h_current.values())) return any_h.sum() * 0.0 total = 0.0 n = 0 for l, hc in h_current.items(): hb = h_base.get(l) if hb is None: continue # Pool over tokens; compare per-sample activations. dc = hc.mean(dim=1)[retain_mask].float() db = hb.mean(dim=1)[retain_mask].float() total = total + F.mse_loss(dc, db, reduction="mean") n += 1 if n == 0: any_h = next(iter(h_current.values())) return any_h.sum() * 0.0 return total / n def kl_batchmean_masked( current_logits: torch.Tensor, p_base: torch.Tensor, sample_mask: torch.Tensor, ) -> torch.Tensor: """Same KL as L_task but averaged only over batch rows with sample_mask.""" if not sample_mask.any(): return current_logits.sum() * 0.0 log_q = F.log_softmax(current_logits[sample_mask].float(), dim=-1) pb = _stable_target_probs(p_base[sample_mask]) return F.kl_div(log_q, pb, reduction="batchmean") def get_sae_features( capture: HiddenStateCapture, sae: FrozenSAEEncoder, ) -> dict[int, torch.Tensor]: return {l: sae(h.mean(dim=1)) for l, h in capture.hidden_states.items()} def count_lm_layers(model) -> int: try: return len(_decoder_layers(model)) except AttributeError: pass for _, module in model.named_modules(): if hasattr(module, "layers") and isinstance(module.layers, nn.ModuleList): return len(module.layers) return 32 def probe_labels( is_scene: torch.Tensor, has_object: torch.Tensor, mode: str, ) -> torch.Tensor: if mode == "union": return ((is_scene + has_object) > 0).float() if mode == "object_only": return has_object.float() if mode == "scene_only": return is_scene.float() raise ValueError(f"Unknown probe_label_mode: {mode}") # --------------------------------------------------------------------------- # Misc utilities # --------------------------------------------------------------------------- class _Tee: """Stream tee utility for logging to both stdout and file.""" def __init__(self, *streams): self._streams = streams def write(self, data): for s in self._streams: s.write(data) s.flush() def flush(self): for s in self._streams: s.flush() def isatty(self) -> bool: return any(getattr(s, "isatty", lambda: False)() for s in self._streams) def fileno(self): return self._streams[0].fileno() def parse_args(): """CLI argument parser for training scripts.""" import argparse parser = argparse.ArgumentParser(description="Adversarial SAE representation suppression") parser.add_argument("--config", type=str, required=True) parser.add_argument("--relation", type=str, default=None) args, unknown = parser.parse_known_args() overrides = {} i = 0 while i < len(unknown): key = unknown[i].lstrip("-") if i + 1 < len(unknown) and not unknown[i + 1].startswith("--"): overrides[key] = unknown[i + 1] i += 2 else: overrides[key] = "true" i += 1 return args, overrides