AQIT 0.1.0
Loading...
Searching...
No Matches
interp_score.py
Go to the documentation of this file.
1# Copyright (c) 2025-present Aquin Labs Private Limited. All Rights Reserved.
2# This file is part of the Aquin Engine. Unauthorized copying, modification,
3# distribution, or use of this file, via any medium, is strictly prohibited.
4# Proprietary and confidential. See LICENSE for terms.
5
6"""
7Ingested from inspection-backend/interp_score.py.
8Import adaptation only: sae/feature_analysis/model_config -> aquin.compute.*.
9"""
10from __future__ import annotations
11
12import torch
13import torch.nn.functional as F
14
15from transformer_lens import HookedTransformer
16
17from aquin.compute.sae import SparseAutoencoder
19 get_causal_label,
20 load_sae,
21 load_norm,
22 normalize,
23 _sae_feature_ablate_hook,
24)
25from aquin.compute.model_loader import get_config
26
27from aquin.compute.device import resolve_compute_device, synchronize_device
28
29DEVICE = resolve_compute_device()
30N_SAMPLES = 10
31
32
33def _generate_sentences(label: str, client, n: int = N_SAMPLES) -> dict:
34 from aquin.compute.llm_json import chat_json_completion, sentence_lists_from_llm
35
36 prompt = f"""You are helping evaluate a sparse autoencoder feature with the label: "{label}"
38Generate {n} short sentences (5-15 words each) where this feature SHOULD fire strongly,
39and {n} short sentences where this feature should NOT fire.
40
41Reply ONLY with a JSON object in this exact format, no markdown:
42{{
43 "positive": ["sentence1", "sentence2", ...],
44 "negative": ["sentence1", "sentence2", ...]
45}}"""
46
47 resp = chat_json_completion(
48 client,
49 model="gpt-4o-mini",
50 messages=[{"role": "user", "content": prompt}],
51 max_tokens=800,
52 temperature=0.7,
53 )
54 raw = resp.choices[0].message.content.strip()
55 positive, negative = sentence_lists_from_llm(raw)
56 return {"positive": positive, "negative": negative}
57
58
59def _get_feature_activation(sentence: str, feature_idx: int, model: HookedTransformer, sae: SparseAutoencoder, model_id: str = "llama-3.2-1b", layer: int | None = None) -> float:
60 cfg = get_config(model_id)
61 resolved_layer = layer if layer is not None else cfg["sae_layer"]
62 tokens = model.to_tokens(sentence)
63 with torch.no_grad():
64 _, cache = model.run_with_cache(
65 tokens,
66 names_filter=f"blocks.{resolved_layer}.hook_resid_post",
67 return_type=None,
68 )
69 resid = cache[f"blocks.{resolved_layer}.hook_resid_post"][0]
70 acts = sae.encode(normalize(resid, model_id, resolved_layer))
71 return float(acts[:, feature_idx].max().item())
72
73
74def _cohen_d_score(pos_acts: list[float], neg_acts: list[float]) -> float:
75 if not pos_acts or not neg_acts:
76 return 0.0
77 pos_t = torch.tensor(pos_acts, dtype=torch.float32)
78 neg_t = torch.tensor(neg_acts, dtype=torch.float32)
79 mu_pos = pos_t.mean().item()
80 mu_neg = neg_t.mean().item()
81 pooled_std = torch.cat([pos_t, neg_t]).std().item() + 1e-6
82 raw = (mu_pos - mu_neg) / pooled_std
83 return round(float(max(0.0, min(1.0, raw))), 4)
84
85
86def _feature_purity_score(sentences: list[str], client) -> float | None:
87 if not sentences or client is None:
88 return None
89
90 try:
91 resp = client.embeddings.create(
92 model="text-embedding-3-large",
93 input=sentences[:10],
94 )
95 vecs = torch.tensor(
96 [e.embedding for e in resp.data], dtype=torch.float32
97 )
98 vecs = F.normalize(vecs, dim=-1)
99 sim_matrix = vecs @ vecs.T
100 n = vecs.shape[0]
101 if n < 2:
102 return None
103 upper_mask = torch.ones(n, n, dtype=torch.bool).triu(diagonal=1)
104 mean_sim = sim_matrix[upper_mask].mean().item()
105 purity = (mean_sim + 1.0) / 2.0
106 return round(float(purity), 4)
107 except Exception as e:
108 print(f"[purity] embedding error: {e}", flush=True)
109 return None
110
111
112def _finite(x: float | None, default: float | None = None) -> float | None:
113 import math
114
115 if x is None:
116 return default
117 try:
118 f = float(x)
119 except (TypeError, ValueError):
120 return default
121 if math.isnan(f) or math.isinf(f):
122 return default
123 return f
124
125
126def _kl_div(p: torch.Tensor, q: torch.Tensor) -> float:
127 p = p.clamp(min=1e-10)
128 q = q.clamp(min=1e-10)
129 return float((p * (p / q).log()).sum().item())
131
132def _entropy(p: torch.Tensor) -> float:
133 p = p.clamp(min=1e-10)
134 return float(-(p * p.log()).sum().item())
135
137def run_mui_score(
138 feature_idx: int,
139 prompt: str,
140 model: HookedTransformer,
141 sae: SparseAutoencoder,
142 n_positions: int = 8,
143 model_id: str = "llama-3.2-1b",
144 layer: int | None = None,
145) -> dict:
146 cfg = get_config(model_id)
147 sae_layer = layer if layer is not None else cfg["sae_layer"]
148 tokens = model.to_tokens(prompt)
149
150 with torch.no_grad():
151 baseline_logits, cache = model.run_with_cache(
152 tokens,
153 names_filter=f"blocks.{sae_layer}.hook_resid_post",
154 )
155
156 resid = cache[f"blocks.{sae_layer}.hook_resid_post"][0]
157 acts = sae.encode(normalize(resid, model_id, sae_layer))
158
159 feat_acts = acts[:, feature_idx]
160 top_k = min(n_positions, feat_acts.shape[0])
161 top_positions = feat_acts.topk(top_k).indices.tolist()
162 top_positions = [p for p in top_positions if feat_acts[p].item() > 0.01]
163
164 if not top_positions:
165 return {
166 "feature_idx": feature_idx,
167 "score": 0.0,
168 "per_position": [],
169 "mean_kl": 0.0,
170 "baseline_entropy": None,
171 }
172
173 baseline_probs = torch.softmax(baseline_logits[0, -1], dim=-1)
174 baseline_H = _entropy(baseline_probs)
175
176 per_position = []
177 kl_vals = []
178
179 for pos in top_positions:
180 act_val = feat_acts[pos].item()
181
182 def ablate(value, hook, pos=pos):
183 return _sae_feature_ablate_hook(
184 value, pos, feature_idx, sae, model_id, sae_layer,
185 )
186
187 with torch.no_grad():
188 abl_logits = model.run_with_hooks(
189 tokens,
190 fwd_hooks=[(f"blocks.{sae_layer}.hook_resid_post", ablate)]
191 )
192
193 abl_probs = torch.softmax(abl_logits[0, -1], dim=-1)
194 kl = _finite(_kl_div(baseline_probs, abl_probs), 0.0) or 0.0
195 kl_vals.append(kl)
196
197 context = model.tokenizer.decode(
198 tokens[0, max(0, pos - 3):pos + 4].tolist()
199 )
200 per_position.append({
201 "position": pos,
202 "activation": round(act_val, 3),
203 "kl_divergence": round(kl, 4),
204 "context": context,
205 })
206
207 mean_kl = sum(kl_vals) / len(kl_vals) if kl_vals else 0.0
208 baseline_H = _finite(baseline_H, 0.0) or 0.0
209 mean_kl = _finite(mean_kl, 0.0) or 0.0
210 mui = _finite(min(mean_kl / max(baseline_H, 1e-6), 1.0), 0.0) or 0.0
211 mui = round(float(mui), 4)
212
213 synchronize_device()
214
215 return {
216 "feature_idx": feature_idx,
217 "score": mui,
218 "per_position": sorted(per_position, key=lambda x: x["kl_divergence"], reverse=True),
219 "mean_kl": round(mean_kl, 4),
220 "baseline_entropy": round(baseline_H, 4),
221 }
222
223
225 feature_idx: int,
226 prompt: str,
227 model: HookedTransformer,
228 sae: SparseAutoencoder,
229 client,
230 n_samples: int = N_SAMPLES,
231 model_id: str = "llama-3.2-1b",
232 layer: int | None = None,
233) -> dict:
234 label = get_causal_label(feature_idx, prompt, model, sae, client, model_id=model_id, layer=layer)
235 mui_result = run_mui_score(feature_idx, prompt, model, sae, model_id=model_id, layer=layer)
236
237 if client is None:
238 return {
239 "feature_idx": feature_idx,
240 "label": label,
241 "score": None,
242 "purity_score": None,
243 "mui_score": mui_result["score"],
244 "mui_per_position": mui_result["per_position"],
245 "mui_mean_kl": mui_result["mean_kl"],
246 "baseline_entropy": mui_result["baseline_entropy"],
247 "error": "OpenAI not available. Set OPENAI_API_KEY on this machine.",
248 "positive_examples": [],
249 "negative_examples": [],
250 "positive_mean": None,
251 "negative_mean": None,
252 }
253
254 try:
255 sentences = _generate_sentences(label, client, n=n_samples)
256 except Exception as e:
257 from aquin.compute.llm_json import llm_sentence_generation_error
258 print(f"[interp] sentence generation: {e}", flush=True)
259 return {
260 "feature_idx": feature_idx,
261 "label": label,
262 "score": None,
263 "purity_score": None,
264 "mui_score": mui_result["score"],
265 "mui_per_position": mui_result["per_position"],
266 "mui_mean_kl": mui_result["mean_kl"],
267 "baseline_entropy": mui_result["baseline_entropy"],
268 "error": llm_sentence_generation_error(),
269 "positive_examples": [],
270 "negative_examples": [],
271 "positive_mean": None,
272 "negative_mean": None,
273 }
274
275 positive_results = []
276 for sent in sentences.get("positive", []):
277 try:
278 act = _get_feature_activation(sent, feature_idx, model, sae, model_id=model_id, layer=layer)
279 positive_results.append({"sentence": sent, "activation": round(act, 4)})
280 except Exception as e:
281 print(f"[interp_score] positive sentence failed: {e}", flush=True)
282
283 negative_results = []
284 for sent in sentences.get("negative", []):
285 try:
286 act = _get_feature_activation(sent, feature_idx, model, sae, model_id=model_id, layer=layer)
287 negative_results.append({"sentence": sent, "activation": round(act, 4)})
288 except Exception as e:
289 print(f"[interp_score] negative sentence failed: {e}", flush=True)
290
291 pos_acts = [r["activation"] for r in positive_results]
292 neg_acts = [r["activation"] for r in negative_results]
293
294 score = _cohen_d_score(pos_acts, neg_acts)
295 pos_mean = round(sum(pos_acts) / len(pos_acts), 4) if pos_acts else None
296 neg_mean = round(sum(neg_acts) / len(neg_acts), 4) if neg_acts else None
297
298 purity_score = _feature_purity_score(
299 [r["sentence"] for r in positive_results], client
300 )
301
302 synchronize_device()
303
304 return {
305 "feature_idx": feature_idx,
306 "label": label,
307 "score": score,
308 "purity_score": purity_score,
309 "mui_score": mui_result["score"],
310 "mui_per_position": mui_result["per_position"],
311 "mui_mean_kl": mui_result["mean_kl"],
312 "baseline_entropy": mui_result["baseline_entropy"],
313 "positive_mean": pos_mean,
314 "negative_mean": neg_mean,
315 "positive_examples": sorted(positive_results, key=lambda x: x["activation"], reverse=True),
316 "negative_examples": sorted(negative_results, key=lambda x: x["activation"], reverse=True),
317 }
dict run_mui_score(int feature_idx, str prompt, HookedTransformer model, SparseAutoencoder sae, int n_positions=8, str model_id="llama-3.2-1b", int|None layer=None)
float _entropy(torch.Tensor p)
float|None _feature_purity_score(list[str] sentences, client)
dict _generate_sentences(str label, client, int n=N_SAMPLES)
dict run_interp_score(int feature_idx, str prompt, HookedTransformer model, SparseAutoencoder sae, client, int n_samples=N_SAMPLES, str model_id="llama-3.2-1b", int|None layer=None)
float _cohen_d_score(list[float] pos_acts, list[float] neg_acts)
float _get_feature_activation(str sentence, int feature_idx, HookedTransformer model, SparseAutoencoder sae, str model_id="llama-3.2-1b", int|None layer=None)
float|None _finite(float|None x, float|None default=None)
float _kl_div(torch.Tensor p, torch.Tensor q)