hallucination / sae /SAE_Blip_Explaining_Utils.py
ToiTenBao's picture
Upload hallucination folder
a2ffd07 verified
Raw
History Blame Contribute Delete
16.3 kB
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}")