# ============================================================ # EXPERIMENT 1: Quantifying Information Loss # ============================================================ !pip install transformer_lens sae_lens scikit-learn -q import torch import torch.nn.functional as F import numpy as np from transformer_lens import HookedTransformer from sae_lens import SAE from sklearn.linear_model import LogisticRegression from sklearn.model_selection import train_test_split from sklearn.preprocessing import StandardScaler from sklearn.metrics import accuracy_score device = "cuda" if torch.cuda.is_available() else "cpu" # ── Load model and SAE ─────────────────────────────── model = HookedTransformer.from_pretrained("gpt2", device=device) model.eval() sae, cfg_dict, _ = SAE.from_pretrained( release="gpt2-small-res-jb", sae_id="blocks.8.hook_resid_pre", ) sae = sae.to(device) # ── Prompts ────────────────────────────────────────── factual_templates = [ "The capital city of {} is", "The currency of {} is the", "The official language of {} is", "The largest city in {} is", ] real_countries = [ "France", "Germany", "Japan", "Brazil", "Australia", "Canada", "India", "Mexico", "Italy", "Spain", "Portugal", "Argentina", "Egypt", "Nigeria", "Kenya", "Sweden", "Norway", "Denmark", "Finland", "Poland", "Greece", "Turkey", "Iran", "Pakistan", "Thailand", "Vietnam", "South Korea", "Indonesia", "Philippines", "Malaysia", ] fictional_countries = [ "Valdoria", "Zephyria", "New Lemuria", "Thessmark", "Arcturis", "Hegemoria", "Veranthos", "Caldoria", "Myranthis", "Delvara", "Quorrath", "Syntropia", "Velundris", "Thaloria", "Zordania", "Ketharia", "Omnivast", "Selphronia", "Draventis", "Nullhaven", "Primoria", "Vexalund", "Threnody", "Galvantis", "Mortherion", "Sundropia", "Westhallow", "Colendris", "Ashenvale", "Ironmere", ] all_prompts = ( [t.format(c) for t in factual_templates for c in real_countries] + [t.format(c) for t in factual_templates for c in fictional_countries] ) # ── Entropy labels ─────────────────────────────────── def get_entropy_labels(prompts, model, batch_size=4): entropies = [] for i in range(0, len(prompts), batch_size): tokens = model.to_tokens(prompts[i:i+batch_size], prepend_bos=True) with torch.no_grad(): logits = model(tokens) probs = F.softmax(logits[:, -1, :], dim=-1) h = -torch.sum(probs * torch.log(probs + 1e-10), dim=-1) entropies.extend(h.cpu().numpy()) entropies = np.array(entropies) return (entropies > np.median(entropies)).astype(int) print("Computing entropy labels...") labels = get_entropy_labels(all_prompts, model) # ── Extract activations ────────────────────────────── def get_activations(prompts, model, sae, batch_size=4): raw_list, sae_list = [], [] for i in range(0, len(prompts), batch_size): tokens = model.to_tokens(prompts[i:i+batch_size], prepend_bos=True) with torch.no_grad(): _, cache = model.run_with_cache( tokens, names_filter="blocks.8.hook_resid_pre" ) acts = cache["blocks.8.hook_resid_pre"][:, -1, :] raw_list.append(acts.cpu().float().numpy()) sae_list.append(sae.encode(acts).cpu().float().numpy()) return np.concatenate(raw_list), np.concatenate(sae_list) print("Extracting activations...") raw_acts, sae_feats = get_activations(all_prompts, model, sae) # ── Classify and report ────────────────────────────── def holdout_accuracy(X, y): X_tr, X_te, y_tr, y_te = train_test_split( X, y, test_size=0.3, random_state=42, stratify=y ) sc = StandardScaler() clf = LogisticRegression(max_iter=1000, random_state=42) clf.fit(sc.fit_transform(X_tr), y_tr) return accuracy_score(y_te, clf.predict(sc.transform(X_te))) acc_baseline = 0.500 acc_sae = holdout_accuracy(sae_feats, labels) acc_raw = holdout_accuracy(raw_acts, labels) print() print(f"Baseline: {acc_baseline:.3f}") print(f"SAE features (text): {acc_sae:.3f}") print(f"Raw activations: {acc_raw:.3f}") print(f"Information lost by SAE: {acc_raw - acc_sae:.3f}")