| from collections.abc import Generator, Sequence |
| from typing import TypeVar, overload |
|
|
| import torch |
| from tqdm.autonotebook import tqdm |
| from transformer_lens.hook_points import HookedRootModule |
| from torch.utils.data import TensorDataset, DataLoader |
| from sae_lens import SAE |
| from typing import Dict, List, Tuple |
| from sae.SAE_Trainer import DataConfig |
| from sae.Load_Data import load_lvlm_data |
| from sae.SAE_Tools import * |
| from IPython.display import HTML, display |
|
|
| T = TypeVar("T") |
| K = TypeVar("K") |
|
|
|
|
| DEFAULT_DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu") |
|
|
|
|
| @overload |
| def batchify( |
| data: Sequence[T], batch_size: int, show_progress: bool = False |
| ) -> Generator[Sequence[T], None, None]: ... |
|
|
|
|
| @overload |
| def batchify( |
| data: torch.Tensor, batch_size: int, show_progress: bool = False |
| ) -> Generator[torch.Tensor, None, None]: ... |
|
|
|
|
| def batchify( |
| data: Sequence[T] | torch.Tensor, batch_size: int, show_progress: bool = False |
| ) -> Generator[Sequence[T] | torch.Tensor, None, None]: |
| """Generate batches from data. If show_progress is True, display a progress bar.""" |
|
|
| for i in tqdm( |
| range(0, len(data), batch_size), |
| total=(len(data) // batch_size + (len(data) % batch_size != 0)), |
| disable=not show_progress, |
| ): |
| yield data[i : i + batch_size] |
|
|
|
|
| def flip_dict(d: dict[T, T]) -> dict[T, T]: |
| """Flip a dictionary, i.e. {a: b} -> {b: a}""" |
| return {v: k for k, v in d.items()} |
|
|
|
|
| def listify(item: T | list[T]) -> list[T]: |
| """Convert an item or list of items to a list.""" |
| if isinstance(item, list): |
| return item |
| return [item] |
|
|
|
|
| def dict_zip(*dicts: dict[T, K]) -> Generator[tuple[T, tuple[K, ...]], None, None]: |
| """Zip together multiple dictionaries, iterating their common keys and a tuple of values.""" |
| if not dicts: |
| return |
| keys = set(dicts[0]).intersection(*dicts[1:]) |
| for key in keys: |
| yield key, tuple(d[key] for d in dicts) |
| |
| def from_top_k_to_tensor( |
| d_sae: int, |
| indices: torch.Tensor, |
| values: torch.Tensor, |
| device: torch.device | str = "cpu", |
| ): |
| if len(indices.shape) == 2: |
| b, _ = indices.shape |
| latents = torch.zeros((b, d_sae), device=device, dtype=values.dtype) |
| elif len(indices.shape) == 3: |
| b, s, _ = indices.shape |
| latents = torch.zeros((b, s, d_sae), device=device, dtype=values.dtype) |
| |
| latents = latents.scatter_( |
| -1, indices.to(device), values.to(device) |
| ) |
| return latents |
| |
| def get_sae_acts( |
| input_activations: torch.Tensor, |
| sae: HookedRootModule, |
| batch_size: int = 4096, |
| device: torch.device | str = "cpu", |
| convert_to_cpu: bool = False, |
| verbose: bool = True, |
| ) -> torch.Tensor | Tuple[torch.Tensor, torch.Tensor]: |
| indices, values = get_sae_activations( |
| sae, |
| DataLoader(TensorDataset(input_activations), batch_size=batch_size, shuffle=False), |
| device, |
| no_tqdm=not verbose, |
| ) |
| |
| return from_top_k_to_tensor( |
| sae.cfg.d_sae, |
| indices, |
| values, |
| device=device if not convert_to_cpu else "cpu", |
| ) |
| |
| def load_sae( |
| model: HookedRootModule, list_release: List[str], list_sae_id: List[str], device: torch.device | str |
| ) -> List[HookedRootModule]: |
| list_saes = [] |
| for release, sae_id in zip(list_release, list_sae_id): |
| if release == "local": |
| list_saes.append( |
| load_sae_model( |
| sae_id, |
| model, |
| device=str(device), |
| ) |
| ) |
| else: |
| sae, _, _ = SAE.from_pretrained( |
| release=release, |
| sae_id=sae_id, |
| device=str(device), |
| ) |
| list_saes.append(sae) |
| |
| return list_saes |
|
|
| def load_data_toks(data_length: int, tok_name: str) -> Tensor: |
| num_workers=4 |
| hf_dataset="yerevann/coco-karpathy" |
| local_train_path="./COCO-Dataset/train_rest" |
| local_val_path="./COCO-Dataset/val" |
| tok_name="Salesforce/blip-image-captioning-base" |
| batch_size=16 |
| max_length=512 |
| data_config = DataConfig( |
| batch_size=batch_size, |
| hf_dataset=hf_dataset, |
| local_train_path=local_train_path, |
| local_val_path=local_val_path, |
| num_workers=num_workers, |
| max_length=max_length, |
| processor = tok_name, |
| ) |
| _, data_loader = load_lvlm_data(data_config) |
| data_toks = extract_data(data_loader, num_batches=data_length) |
| return data_toks |
|
|
| def cache_activation_model( |
| hook_name: str, |
| model: HookedTransformer, |
| x: Tensor, |
| batch_size_model: int, |
| verbose: bool = True |
| ) -> Tensor: |
| target_cache = [] |
| def hook_fn(tens: Tensor, hook: HookPoint): |
| batch, seq = tens.shape[0], tens.shape[1] |
| target_cache.append(tens.reshape(batch * seq, -1).cpu().detach()) |
| |
| with t.no_grad(): |
| with model.hooks( |
| fwd_hooks=[ |
| (hook_name, hook_fn) |
| ] |
| ): |
| for toks in batchify(x, batch_size_model, show_progress=verbose): |
| model(toks) |
| |
| acts = t.cat(target_cache) |
| |
| return acts |
|
|
| @t.inference_mode() |
| def subsample_tensor(tensor: torch.Tensor, max_samples: int) -> torch.Tensor: |
| """ |
| Subsample a 2D tensor if the number of samples exceeds the specified maximum. |
| |
| Args: |
| tensor (torch.Tensor): Input 2D tensor of shape (n_sample, f). |
| max_samples (int): Maximum number of samples to retain. |
| |
| Returns: |
| torch.Tensor: Subsampled tensor of shape (min(n_sample, max_samples), f). |
| """ |
| n_sample, f = tensor.shape |
| if n_sample > max_samples: |
| indices = torch.randperm(n_sample)[:max_samples] |
| return tensor[indices] |
| return tensor |
|
|
|
|
| @t.no_grad() |
| def select_feature_from_probe( |
| probe_weight: Tensor, |
| W_dec: Tensor, |
| sae_acts: Tensor, |
| labels: Tensor, |
| ): |
| mask = t.where(labels, t.ones_like(labels).float(), -t.ones_like(labels).float()).unsqueeze(-1) |
| positive_label_acts = (sae_acts * mask).mean(0).clamp(min=0) |
| positive_label_directions = positive_label_acts.unsqueeze(-1) * W_dec |
| |
| def normalize(tens: Tensor): |
| return tens / tens.norm(2, dim=1).max() |
| |
| scores = normalize(positive_label_directions) @ normalize(probe_weight).T |
| |
| return scores |
|
|
|
|
| @t.inference_mode() |
| def compute_f1( |
| masks: t.Tensor, |
| indices: t.Tensor, |
| target_idx: int, |
| device: t.device, |
| pad_value: int = -1, |
| feature_batch_size: int = 64, |
| target_batch_size: int = 16, |
| compute_dtype: t.dtype = t.float32, |
| other_feat_idx: t.Tensor | None = None, |
| ): |
| """ |
| Compute maximum F1 scores for a batch of target feature activations vs. other features, |
| and return both the F1 scores and the indices of the features that achieved them. |
| |
| Args: |
| masks: Bool or 0/1 tensor of shape (num_targets, n_sample). masks[i, n] == 1 if target i is active on sample n. |
| indices: Int tensor of shape (n_sample, k). Each row holds up to k activated feature indices for that sample. |
| target_idx: The feature index to exclude from the candidates (e.g., the "self" feature). |
| device: Torch device to run on. |
| pad_value: Padding value in `indices` rows. |
| feature_batch_size: Number of candidate features per batch. |
| target_batch_size: Number of target rows per batch. |
| compute_dtype: Accumulation dtype (float32 by default). |
| other_feat_idx: (n_index) If we only compute F1 among a certain feature indices. |
| |
| Returns: |
| Tuple of (f1_scores, feature_indices): |
| - f1_scores: 1D tensor of shape (num_targets,) with the max F1 over all other features for each target. |
| - feature_indices: 1D tensor of shape (num_targets,) with the index of the feature that achieved the max F1. |
| """ |
| masks = masks.to(device=device).bool() |
| indices = indices.to(device=device) |
|
|
| n_sample = indices.shape[0] |
| num_targets = masks.shape[0] |
|
|
| |
| if other_feat_idx is None: |
| valid_mask = indices.ne(pad_value) |
| if valid_mask.any(): |
| flat_indices = indices[valid_mask] |
| |
| unique_indices = t.unique(flat_indices, sorted=False) |
| |
| other_feat_idx = unique_indices[unique_indices.ne(t.as_tensor(target_idx, device=unique_indices.device))] |
| else: |
| other_feat_idx = t.empty(0, dtype=indices.dtype, device=indices.device) |
|
|
| if other_feat_idx.numel() == 0: |
| zeros = t.zeros(num_targets, dtype=compute_dtype, device=device) |
| neg_ones = t.full((num_targets,), -1, dtype=indices.dtype, device=device) |
| return zeros, neg_ones |
|
|
| |
| num_target_batches = (num_targets + target_batch_size - 1) // target_batch_size |
| all_max_f1_scores = [] |
| all_best_indices = [] |
|
|
| for target_batch_idx in range(num_target_batches): |
| st = target_batch_idx * target_batch_size |
| en = min(st + target_batch_size, num_targets) |
|
|
| |
| T = masks[st:en] |
| Tb = T.shape[0] |
|
|
| |
| A = T.sum(dim=1, dtype=compute_dtype) |
|
|
| |
| best_f1 = t.zeros(Tb, dtype=compute_dtype, device=device) |
| best_indices = t.full((Tb,), -1, dtype=other_feat_idx.dtype, device=device) |
|
|
| |
| num_feature_batches = (other_feat_idx.numel() + feature_batch_size - 1) // feature_batch_size |
| for j in range(num_feature_batches): |
| fs = j * feature_batch_size |
| fe = min(fs + feature_batch_size, other_feat_idx.numel()) |
| feature_batch = other_feat_idx[fs:fe] |
|
|
| |
| |
| |
| F = (indices.unsqueeze(-1) == feature_batch.view(1, 1, -1)).any(dim=1) |
|
|
| |
| B = F.sum(dim=0, dtype=compute_dtype) |
|
|
| |
| TP = t.matmul(T.to(dtype=compute_dtype), F.to(dtype=compute_dtype)) |
|
|
| |
| denom = A.unsqueeze(1) + B.unsqueeze(0) |
| f1 = t.where(denom > 0, (2.0 * TP) / denom, t.zeros((), dtype=compute_dtype, device=device)) |
|
|
| |
| batch_best_f1, batch_best_idx = f1.max(dim=1) |
| batch_best_idx = batch_best_idx + fs |
| |
| |
| update_mask = batch_best_f1 > best_f1 |
| best_f1 = t.where(update_mask, batch_best_f1, best_f1) |
| best_indices = t.where( |
| update_mask, |
| other_feat_idx[batch_best_idx], |
| best_indices |
| ) |
|
|
| all_max_f1_scores.append(best_f1) |
| all_best_indices.append(best_indices) |
|
|
| return t.cat(all_max_f1_scores, dim=0), t.cat(all_best_indices, dim=0) |
|
|
|
|
| def extract_context(data_tensor: Tensor, index_tensor: Tensor, context_size=15): |
| """ |
| Extract context windows of ±context_size around each index. |
| |
| Args: |
| data_tensor: 2D tensor of shape (batch, seq_len) |
| index_tensor: 1D tensor of shape (batch,) containing indices |
| context_size: Size of context window on each side (default: 15) |
| |
| Returns: |
| 2D tensor of shape (batch, 2 * context_size + 1) with context windows |
| """ |
| batch_size, seq_len = data_tensor.shape |
| window_size = 2 * context_size + 1 |
| |
| |
| indices = t.arange(-context_size, context_size + 1, device=data_tensor.device) |
| indices = indices.view(1, -1).expand(index_tensor.shape[0], window_size) |
| |
| |
| indices = indices + index_tensor |
| |
| |
| indices = t.clamp(indices, 0, data_tensor.flatten().shape[0]-1) |
| |
| context_windows = data_tensor.flatten()[indices] |
| return context_windows |
|
|
| def highlight_html(strings: List[str], highlight_index: int): |
| """ |
| For Jupyter notebooks - uses HTML formatting |
| """ |
| html_str = "" |
| for i, s in enumerate(strings): |
| s = s.replace("�", "").replace("\n", "↵") |
| if i == highlight_index: |
| html_str += f'<span style="color: red; font-weight: bold">{s}</span>' |
| else: |
| html_str += f'{s}' |
| display(HTML(html_str)) |
| |
| def show_activation( |
| model: HookedTransformer, |
| data_toks: Tensor, |
| index_tensor: Tensor, |
| num_examples: int = 50, |
| context_size: int = 15, |
| ): |
| contexts = extract_context(data_toks, index_tensor, context_size=context_size) |
| for context in contexts[:num_examples]: |
| highlight_html(model.to_str_tokens(context), context_size) |