!pip install transformer-lens !pip install sae-lens import torch import torch.nn.functional as F import pickle import os from transformer_lens import HookedTransformer from sae_lens import SAE SAVE_DIR = "/content/" # change to Drive path if mounted # ── Load model and SAE ─────────────────────────────── print("Loading GPT-2 Small...") device = 'cpu' model = HookedTransformer.from_pretrained("gpt2", device=device) model.eval() print("Loading pretrained SAE (layer 8 residual stream)...") sae, cfg_dict, _ = SAE.from_pretrained( release="gpt2-small-res-jb", sae_id="blocks.8.hook_resid_pre", ) sae = sae.to(device) print(f"Model loaded — d_model: {model.cfg.d_model}") print(f"SAE loaded — d_sae: {sae.cfg.d_sae}") # ── Build prompt dataset ──────────────────────────── # Minimal pairs: identical syntactic templates, real vs fictional referents # Same structure eliminates surface-form as a confound 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", ] factual_prompts = [ template.format(country) for template in factual_templates for country in real_countries ] counterfactual_prompts = [ template.format(country) for template in factual_templates for country in fictional_countries ] all_prompts = factual_prompts + counterfactual_prompts print(f"Factual prompts: {len(factual_prompts)}") print(f"Counterfactual prompts: {len(counterfactual_prompts)}") print(f"Total: {len(all_prompts)}") def get_output_entropy(prompts, model, batch_size=4): """ Returns Shannon entropy of the next-token distribution for each prompt. This is the model's behavioral signal. """ entropies = [] for i in range(0, len(prompts), batch_size): batch = prompts[i:i+batch_size] tokens = model.to_tokens(batch, prepend_bos=True) with torch.no_grad(): logits = model(tokens) # [batch, seq_len, vocab] last_logits = logits[:, -1, :] # [batch, vocab] probs = F.softmax(last_logits, dim=-1) entropy = -torch.sum( probs * torch.log(probs + 1e-10), dim=-1 # [batch] ) entropies.extend(entropy.tolist()) if i % 40 == 0: print(f" Entropy computed for {i}/{len(prompts)} prompts") return entropies def get_activations_last_token(prompts, model, sae, layer=8, batch_size=4): """ Extracts: raw_acts [N, d_model] — last-token residual stream activations sae_feats [N, d_sae] — last-token SAE feature activations """ raw_acts_list = [] sae_feats_list = [] hook_name = f"blocks.{layer}.hook_resid_pre" for i in range(0, len(prompts), batch_size): batch = prompts[i:i+batch_size] tokens = model.to_tokens(batch, prepend_bos=True) with torch.no_grad(): _, cache = model.run_with_cache(tokens, names_filter=hook_name) acts = cache[hook_name] # [batch, seq_len, d_model] acts_last = acts[:, -1, :] # [batch, d_model] sae_out = sae.encode(acts_last) # [batch, d_sae] raw_acts_list.append(acts_last) sae_feats_list.append(sae_out) if i % 20 == 0: print(f" Processed {i}/{len(prompts)}") return ( torch.concatenate(raw_acts_list, axis=0), torch.concatenate(sae_feats_list, axis=0), ) def get_sae_entropy(prompts, model, sae): acts, feats = get_activations_last_token(prompts, model, sae) probs = F.softmax(feats, dim=-1) entropies = -torch.sum( probs * torch.log(probs + 1e-10), dim = -1) return entropies entropies = get_output_entropy(all_prompts, model) # Sanity check: do prompt categories align with entropy? factual_entropy = torch.tensor(entropies[:len(factual_prompts)]) counter_entropy = torch.tensor(entropies[len(factual_prompts):]) print(f"\nEntropy by model:") print(f" Factual prompts: {factual_entropy.mean().item()} mean") print(f" Counterfactual prompts: {counter_entropy.mean().item()} mean") sae_entropies = torch.tensor(get_sae_entropy(all_prompts, model, sae)) factual_entropy = torch.tensor(sae_entropies[:len(factual_prompts)]) counter_entropy = torch.tensor(sae_entropies[len(factual_prompts):]) print(f"\nEntropy by sae:") print(f" Factual prompts: {factual_entropy.mean().item()} mean") print(f" Counterfactual prompts: {counter_entropy.mean().item()} mean")