import torch import torch as t from jaxtyping import Float, Int from functools import partial from torch.utils.data import Dataset, DataLoader, Subset from typing import Callable, List, Dict, Any, Optional, Literal from transformer_lens.hook_points import HookPoint from .SAE_Tools import * from .SAE_Trainer import DataConfig from PIL import Image from pathlib import Path from transformers import BlipProcessor, LlavaProcessor import torch.nn.functional as F from tqdm import tqdm from torchvision.utils import make_grid from sae.Load_Data import * import numpy as np import matplotlib.pyplot as plt from contextlib import contextmanager # UTILS FOR VISUALIZING BLIP VISION SAEs def fetch_topk_activating_crops_lvlm( feat_idx: int, values: Tensor, indices: Tensor, dataset: Dataset, top_k: int, output_dir: str, verbose: bool = True, ): mask = (values > 1e-3) * (indices == feat_idx) if verbose: print("Feature ID:", feat_idx, "Density:", mask.sum().item() / mask.numel(), "Max act:", (values * mask).max().item()) if mask.sum() == 0: if verbose: print("No activation!") return filtered_values = values * mask.to(values.dtype) # (b, activated_dim) sumed_filtered_values = filtered_values.sum(dim=1) # (b) top_vals, top_indices = sumed_filtered_values.topk(k=top_k) # (topk) top_acts = filtered_values[top_indices, :] # (topk, activated_dim) visualize_crops( top_vals = top_vals, top_indices = top_indices, top_acts = top_acts, feat_idx = feat_idx, multi_crop_dataset = dataset, output_dir = output_dir, ) def denorm(img_tensor): mean = torch.tensor([0.485, 0.456, 0.406]).view(3,1,1) std = torch.tensor([0.229, 0.224, 0.225]).view(3,1,1) return (img_tensor * std + mean).clamp(0, 1) def visualize_crops( top_vals: torch.Tensor, top_indices: torch.Tensor, top_acts: torch.Tensor, feat_idx: list[int], multi_crop_dataset: Dataset, output_dir: str, ): """ For each feature in feature_list, visualize and save the top-k crops. """ crops = [] # crops from dataset for each in range(top_indices.shape[0]): patch = multi_crop_dataset[top_indices[each].item()] crops.append(patch['pixel_values']) crops_tensor = torch.stack(crops, dim=0) # [top_k, 3, 384, 384] # Save images save_dir = Path(f'{output_dir}/feature_{feat_idx}_top_crops') save_dir.mkdir(parents=True, exist_ok=True) max_crop = denorm(crops_tensor[0]) # highest activation max_act_val = top_vals[0].item() max_pil = Image.fromarray((max_crop.permute(1,2,0).cpu().numpy() * 255).astype(np.uint8)) grid = make_grid(denorm(crops_tensor), nrow=5, padding=12, pad_value=1.0) # white padding grid = (grid.clamp(0,1) * 255).byte().cpu().permute(1,2,0).numpy() grid_pil = Image.fromarray(grid) max_pil.save(save_dir / f'max_activation_act_{max_act_val:.4f}.jpg') grid_pil.save(save_dir / "top_grid.jpg") # UTILS FOR VISUALIZING BLIP TEXT SAEs # --- Updated Data Extraction & Visualization --- def fetch_feature_activation_intervals_blip( processor: BlipProcessor | LlavaProcessor, feat_idx: int, values: t.Tensor, indices: t.Tensor, toks: t.Tensor, n_intervals: int = 5, k_per_interval: int = 5, buffer: int = 5, ) -> Dict[str, Any]: """ Extracts examples for specific feature activation intervals. """ # 1. Reconstruct dense activations feat_acts = t.zeros_like(toks, dtype=t.float) mask = (indices == feat_idx) feat_acts = (values * mask.float()).sum(dim=-1) max_act = feat_acts.max().item() # Debug output if max_act <= 0: return {"feature_id": feat_idx, "intervals": [], "max_act": 0} boundaries = t.linspace(1e-3, max_act, n_intervals + 1) interval_data = [] total_tokens = feat_acts.numel() seq_len = toks.shape[1] for i in range(n_intervals, 0, -1): lower_bound = boundaries[i-1].item() upper_bound = boundaries[i].item() # Inclusive upper bound for the top interval if i == n_intervals: interval_mask = (feat_acts >= lower_bound) & (feat_acts <= upper_bound + 1e-4) else: interval_mask = (feat_acts >= lower_bound) & (feat_acts < upper_bound) count = interval_mask.sum().item() density = count / total_tokens if count == 0: continue # Apply mask to activations so get_k_largest only sees valid interval data masked_acts = feat_acts * interval_mask.float() # Call the FIXED selector # We pass the buffer here so it knows how wide to mask neighbors, # but the function now safely handles edges. selected_indices = get_k_largest_indices(masked_acts, k=k_per_interval, buffer=buffer, no_overlap=True) if selected_indices.shape[0] == 0: continue examples = [] for j in range(selected_indices.shape[0]): row_idx = selected_indices[j, 0].item() col_idx = selected_indices[j, 1].item() actual_val = feat_acts[row_idx, col_idx].item() # type: ignore # --- Dynamic Window Slicing --- # Calculate start/end based on buffer, clamped to sequence limits. # This removes the need for padding tokens. start_idx = max(0, col_idx - buffer) end_idx = min(seq_len, col_idx + buffer + 1) # 1. Short View (Windowed) raw_ids_short = toks[row_idx, start_idx:end_idx].tolist() # type: ignore str_tokens_short = [] for tok_id in raw_ids_short: str_tokens_short.append(" " + processor.tokenizer.decode(t.tensor(tok_id))) acts_short = feat_acts[row_idx, start_idx:end_idx].tolist() # type: ignore # Calculate target index relative to the short slice # e.g., if slice starts at 5 and target is at 7, relative is 2 relative_target_idx = col_idx - start_idx # 2. Full View raw_ids_full = toks[row_idx].tolist() # type: ignore acts_full = feat_acts[row_idx].tolist() # type: ignore str_tokens_full = [] for tok_id in raw_ids_full: str_tokens_full.append(" " + processor.tokenizer.decode(t.tensor(tok_id))) examples.append({ "short": { "tokens": str_tokens_short, "activations": acts_short, "target_idx": relative_target_idx }, "full": { "tokens": str_tokens_full, "activations": acts_full, "target_idx": col_idx }, "target_act": actual_val, }) if len(examples) > 0: interval_data.append({ "min": lower_bound, "max": upper_bound, "density": density, "examples": examples }) return { "feature_id": feat_idx, "max_act": max_act, "intervals": interval_data } def pad_and_concat(tensors, dim=0, padding_value=0): """ Pad tensors to same size along all dimensions except 'dim' and concatenate """ # Find max size for each dimension max_sizes = [max(tens.size(d) for tens in tensors) for d in range(tensors[0].dim())] padded_tensors = [] for tens in tqdm(tensors, desc = "Pad and concatenate tensors"): # Calculate padding for each dimension padding = [] for i, (current, target) in enumerate(zip(tens.shape, max_sizes)): if i != dim: # Don't pad the concatenation dimension padding = [0, target - current] + padding # Prepend for reverse order if padding: padded = F.pad(tens, padding, value=padding_value) padded_tensors.append(padded) else: padded_tensors.append(tens) return torch.cat(padded_tensors, dim=dim) @t.no_grad() def cache_sae_lvlm( saes: List[Any], model, dataloader: DataLoader, device: str, filter_seq_length: int | None = None, # filter too long seq return_toks: bool = True, stop_at_batch: int | None = None, ): act_dict = {sae.cfg.hook_name: [] for sae in saes} def caching_hook(act: Tensor, hook: HookPoint): act_dict[".".join(hook.name.split(".")[:-1])].append(act.detach().cpu()) data_toks = [] with model.saes(saes, use_error_term=True): with model.hooks(fwd_hooks=[(lambda name: "acts_post" in name, caching_hook)]): for batch_idx_iter, batch in enumerate(tqdm(dataloader, desc="Caching SAE acts")): if stop_at_batch is not None: if batch_idx_iter > stop_at_batch: break inputs = { "pixel_values": batch["pixel_values"].to(device), "input_ids": batch["input_ids"].to(device), "attention_mask": batch["attention_mask"].to(device), } if filter_seq_length is not None and batch["input_ids"].shape[1] > filter_seq_length: continue if return_toks: data_toks.append(batch["input_ids"].cpu()) model(inputs) if return_toks: data_toks = pad_and_concat(data_toks, dim=0, padding_value=0) # [PAD] token for key, act_list in act_dict.items(): acts = pad_and_concat(act_list, dim=0, padding_value=0) # (b, s, d_sae), pad 0 activation value max_k = int((acts > 1e-3).sum(dim=-1).max().item()) values, indices = acts.topk(k=max_k, dim=-1) act_dict[key] = (values, indices) return act_dict, data_toks def cache_vision_sae_lvlm( saes: List[Any], model, dataloader: DataLoader, device: str, mode: str | Literal["acts_post", "acts_pre"] = "acts_post", filter_seq_length: int | None = None, # filter too long seq return_toks: bool = True, stop_at_batch: int | None = None, ): act_dict = {sae.cfg.hook_name: [] for sae in saes} def caching_hook(act: Tensor, hook: HookPoint): # print(act.shape) act_dict[".".join(hook.name.split(".")[:-1])].append(act.squeeze(1).detach().cpu()) @contextmanager def _hook_vision_sae(): pass_through_vision_sae_cache = [t.tensor(0)] # placeholder def hook_fn(act: Tensor, hook: HookPoint, sae_name: str): if sae_name + ".hook_sae_input" == hook.name: pass_through_vision_sae_cache[0] = act return act.mean(dim=1, keepdim=True) elif sae_name + ".hook_sae_output" == hook.name: return pass_through_vision_sae_cache[0] elif sae_name + ".hook_sae_error" == hook.name: act = t.zeros_like(act).mean(dim=1, keepdim=True) return act try: for vision_sae in saes: vision_sae.add_hook( lambda name: True, partial(hook_fn, sae_name=vision_sae.cfg.hook_name), dir="fwd", ) yield finally: for vision_sae in saes: vision_sae.reset_hooks() data_toks = [] with _hook_vision_sae(): with model.saes(saes, use_error_term=True): with model.hooks(fwd_hooks=[(lambda name: mode in name, caching_hook)]): for batch_idx_iter, batch in enumerate(tqdm(dataloader, desc="Caching SAE acts")): if stop_at_batch is not None: if batch_idx_iter > stop_at_batch: break inputs = { "pixel_values": batch["pixel_values"].to(device), "input_ids": batch["input_ids"].to(device), "attention_mask": batch["attention_mask"].to(device), } if filter_seq_length is not None and batch["input_ids"].shape[1] > filter_seq_length: continue if return_toks: data_toks.append(batch["input_ids"].cpu()) model(inputs) if return_toks: data_toks = pad_and_concat(data_toks, dim=0, padding_value=0) # [PAD] token for key, act_list in act_dict.items(): acts = pad_and_concat(act_list, dim=0, padding_value=0) # (b, d_sae), pad 0 activation value in batch # print(acts.shape) max_k = int((acts > 1e-3).sum(dim=-1).max().item()) values, indices = acts.topk(k=max_k, dim=-1) act_dict[key] = (values, indices) return act_dict, data_toks def ig_sae_act_to_inputs( feature_id: int, model, inputs: dict[str, torch.Tensor], steps: int, sae, device: str, ): """ Computes Integrated Gradients from SAE feature activation back to input pixels and text embeddings. Fully batched: works with inputs of shape [B, ...] Returns: ig_image: [B, 3, H, W] ig_text: [B, seq_len, d_model] # attribution to token embeddings """ baseline_pixels = torch.zeros_like(inputs["pixel_values"]) baseline_input_ids = torch.zeros_like(inputs["input_ids"]) baseline_embeds = model.text_decoder.bert.embeddings(input_ids=baseline_input_ids) real_embeds = model.text_decoder.bert.embeddings(input_ids=inputs["input_ids"]) total_grad_pix = torch.zeros_like(inputs["pixel_values"]) total_grad_txt = torch.zeros_like(real_embeds) # [B, seq, d_model] def make_embed_hook(interp_embeds): def hook_fn(old_embeds, hook): return interp_embeds return hook_fn captured_acts = None def capture_sae_acts(acts, hook): nonlocal captured_acts captured_acts = acts # [B, seq, n_features] return acts for alpha in torch.linspace(0, 1, steps, device=device): interp_pixels = baseline_pixels + alpha * (inputs["pixel_values"] - baseline_pixels) interp_embeds = baseline_embeds + alpha * (real_embeds - baseline_embeds) interp_pixels.requires_grad_(True) interp_embeds.requires_grad_(True) interp_inputs = { "pixel_values": interp_pixels, "input_ids": inputs["input_ids"], "attention_mask": inputs["attention_mask"], } _ = model.run_with_hooks_with_saes( inputs=interp_inputs, fwd_hooks=[ ("text_decoder.bert.hook_text_embeddings", make_embed_hook(interp_embeds)), (sae.cfg.hook_name + '.hook_sae_acts_post', capture_sae_acts), ], saes=[sae], ) feature_acts = captured_acts[..., feature_id] # [B, seq] loss = feature_acts.sum(dim=-1) # -> [B] loss = loss.sum() grad_pix, grad_txt = torch.autograd.grad( loss, [interp_pixels, interp_embeds], ) total_grad_pix += grad_pix.detach() total_grad_txt += grad_txt.detach() ig_image = (inputs["pixel_values"] - baseline_pixels) * (total_grad_pix / steps) ig_text = (real_embeds - baseline_embeds) * (total_grad_txt / steps) return ig_image, ig_text def vis_text_ig( ig_text: Float[Tensor, "1 seq d_model"], inputs: dict, processor, feature_id: int, ): print(f"Feature {feature_id} - Text IG Attribution:") ig_text = ig_text.squeeze(0) # [seq, d_model] sign = ig_text.sum(dim=-1).sign() # [seq] magnitude = ig_text.abs().sum(dim=-1) # [seq] token_importance = sign * magnitude tokens = processor.tokenizer.convert_ids_to_tokens(inputs["input_ids"][0]) for tok, imp in zip(tokens, token_importance.tolist()): print(f"{tok:15} → {imp: .4f}")