| 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 |
|
|
|
|
|
|
| |
|
|
| 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) |
| |
| sumed_filtered_values = filtered_values.sum(dim=1) |
| top_vals, top_indices = sumed_filtered_values.topk(k=top_k) |
| top_acts = filtered_values[top_indices, :] |
| |
| |
| 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 = [] |
| 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) |
| |
| |
| 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]) |
| 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) |
| 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") |
|
|
|
|
|
|
|
|
| |
|
|
| |
|
|
| 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. |
| """ |
| |
| |
| 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() |
| |
| |
| 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() |
| |
| |
| 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 |
|
|
| |
| masked_acts = feat_acts * interval_mask.float() |
| |
| |
| |
| |
| 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() |
|
|
| |
| |
| |
| start_idx = max(0, col_idx - buffer) |
| end_idx = min(seq_len, col_idx + buffer + 1) |
| |
| |
| raw_ids_short = toks[row_idx, start_idx:end_idx].tolist() |
| 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() |
| |
| |
| |
| relative_target_idx = col_idx - start_idx |
| |
| |
| raw_ids_full = toks[row_idx].tolist() |
| acts_full = feat_acts[row_idx].tolist() |
| 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 |
| """ |
| |
| 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"): |
| |
| padding = [] |
| for i, (current, target) in enumerate(zip(tens.shape, max_sizes)): |
| if i != dim: |
| padding = [0, target - current] + padding |
| |
| 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, |
| 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) |
| |
| for key, act_list in act_dict.items(): |
| acts = pad_and_concat(act_list, dim=0, padding_value=0) |
| 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, |
| 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.squeeze(1).detach().cpu()) |
| |
| @contextmanager |
| def _hook_vision_sae(): |
| pass_through_vision_sae_cache = [t.tensor(0)] |
| 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) |
| |
| for key, act_list in act_dict.items(): |
| acts = pad_and_concat(act_list, dim=0, padding_value=0) |
| |
| 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) |
|
|
| 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 |
| 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] |
| loss = feature_acts.sum(dim=-1) |
| 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) |
| |
| sign = ig_text.sum(dim=-1).sign() |
| magnitude = ig_text.abs().sum(dim=-1) |
| |
| 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}") |
|
|
|
|
|
|