38Generate {n} short sentences (5-15 words each) where this feature SHOULD fire strongly,
39and {n} short sentences where this feature should NOT fire.
41Reply ONLY with a JSON object in this exact format, no markdown:
43 "positive": ["sentence1", "sentence2", ...],
44 "negative": ["sentence1", "sentence2", ...]
47 resp = chat_json_completion(
50 messages=[{
"role":
"user",
"content": prompt}],
54 raw = resp.choices[0].message.content.strip()
55 positive, negative = sentence_lists_from_llm(raw)
56 return {
"positive": positive,
"negative": negative}
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)
64 _, cache = model.run_with_cache(
66 names_filter=f
"blocks.{resolved_layer}.hook_resid_post",
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())
74def _cohen_d_score(pos_acts: list[float], neg_acts: list[float]) -> float:
75 if not pos_acts
or not neg_acts:
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)
87 if not sentences
or client
is None:
91 resp = client.embeddings.create(
92 model=
"text-embedding-3-large",
96 [e.embedding
for e
in resp.data], dtype=torch.float32
98 vecs = F.normalize(vecs, dim=-1)
99 sim_matrix = vecs @ vecs.T
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)
112def _finite(x: float |
None, default: float |
None =
None) -> float |
None:
141 sae: SparseAutoencoder,
142 n_positions: int = 8,
143 model_id: str =
"llama-3.2-1b",
144 layer: int |
None =
None,
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)
150 with torch.no_grad():
151 baseline_logits, cache = model.run_with_cache(
153 names_filter=f
"blocks.{sae_layer}.hook_resid_post",
156 resid = cache[f
"blocks.{sae_layer}.hook_resid_post"][0]
157 acts = sae.encode(normalize(resid, model_id, sae_layer))
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]
164 if not top_positions:
166 "feature_idx": feature_idx,
170 "baseline_entropy":
None,
173 baseline_probs = torch.softmax(baseline_logits[0, -1], dim=-1)
174 baseline_H =
_entropy(baseline_probs)
179 for pos
in top_positions:
180 act_val = feat_acts[pos].item()
182 def ablate(value, hook, pos=pos):
183 return _sae_feature_ablate_hook(
184 value, pos, feature_idx, sae, model_id, sae_layer,
187 with torch.no_grad():
188 abl_logits = model.run_with_hooks(
190 fwd_hooks=[(f
"blocks.{sae_layer}.hook_resid_post", ablate)]
193 abl_probs = torch.softmax(abl_logits[0, -1], dim=-1)
197 context = model.tokenizer.decode(
198 tokens[0, max(0, pos - 3):pos + 4].tolist()
200 per_position.append({
202 "activation": round(act_val, 3),
203 "kl_divergence": round(kl, 4),
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)
216 "feature_idx": feature_idx,
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),
227 model: HookedTransformer,
228 sae: SparseAutoencoder,
230 n_samples: int = N_SAMPLES,
231 model_id: str =
"llama-3.2-1b",
232 layer: int |
None =
None,
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)
239 "feature_idx": feature_idx,
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,
256 except Exception
as e:
258 print(f
"[interp] sentence generation: {e}", flush=
True)
260 "feature_idx": feature_idx,
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,
275 positive_results = []
276 for sent
in sentences.get(
"positive", []):
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)
283 negative_results = []
284 for sent
in sentences.get(
"negative", []):
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)
291 pos_acts = [r[
"activation"]
for r
in positive_results]
292 neg_acts = [r[
"activation"]
for r
in negative_results]
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
299 [r[
"sentence"]
for r
in positive_results], client
305 "feature_idx": feature_idx,
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),
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)
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)