7Ingested from inspection-backend/feature_analysis.py.
8Import adaptation only: model_config -> aquin.compute.model_loader, sae -> aquin.compute.sae.
10from __future__
import annotations
15from transformer_lens
import HookedTransformer
20 get_config
as _get_config,
26DEVICE = resolve_compute_device()
28_sae_cache: dict[tuple, SparseAutoencoder] = {}
29_norm_cache: dict[tuple, dict |
None] = {}
30_session_label_cache = {}
35_kernel_feature_acts =
None
37_kernel_top_features = []
43 return resolve_sae_path(model_id, layer)
47 from pathlib
import Path
51 cfg = get_config(model_id)
52 short = resolve_model_id(model_id)
53 resolved_layer = int(layer
if layer
is not None else cfg[
"sae_layer"])
55 sae_path = resolve_sae_checkpoint_path(short, resolved_layer)
56 if sae_path
is not None:
57 parent = sae_path.parent
59 parent / f
"norm_layer{resolved_layer}.pt",
61 parent / f
"_acts_layer{resolved_layer}" /
"norm.pt",
63 if candidate.is_file():
66 base = Path.home() /
".aquin" /
"sae" / short
67 if "norm_layers" in cfg:
68 rel = cfg[
"norm_layers"].get(resolved_layer)
70 p = base / Path(rel).name
73 p_rel = cfg.get(
"norm_path")
75 p = base / Path(p_rel).name
76 return p
if p.exists()
else None
80def _load_sae_native(model_id: str, layer: int |
None =
None) -> SparseAutoencoder:
82 short = resolve_model_id(model_id)
83 cfg = get_config(short)
84 resolved_layer = int(layer
if layer
is not None else cfg[
"sae_layer"])
85 if not sae_path.exists():
86 raise FileNotFoundError(
87 f
"SAE not found at {sae_path}. Run: aquin load sae {short}-l{resolved_layer}"
91 corrupt = corrupt_sae_checkpoint_hint(short, resolved_layer)
93 raise ValueError(corrupt)
94 return SparseAutoencoder.load(sae_path, device=DEVICE)
97def load_sae(model_id: str =
"llama-3.2-1b", layer: int |
None =
None) -> SparseAutoencoder:
99 short = resolve_model_id(model_id)
100 cfg = get_config(short)
101 resolved_layer = int(layer
if layer
is not None else cfg[
"sae_layer"])
102 key = (short, resolved_layer)
103 if key
not in _sae_cache:
106 _sae_cache[key] = load_sae_from_disk(short, resolved_layer, device=DEVICE)
107 if short ==
"llama-3.2-1b" and layer
is None:
108 _sae = _sae_cache[key]
109 return _sae_cache[key]
112def load_norm(model_id: str =
"llama-3.2-1b", layer: int |
None =
None) -> dict |
None:
114 cfg = get_config(model_id)
115 resolved_layer = layer
if layer
is not None else cfg[
"sae_layer"]
116 key = (model_id, resolved_layer)
117 if key
not in _norm_cache:
119 if norm_path
and norm_path.exists():
122 loaded = load_norm_stats(norm_path, map_location=DEVICE)
123 _norm_cache[key] = loaded
125 _norm_cache[key] =
None
126 if model_id ==
"llama-3.2-1b" and layer
is None:
127 _norm = _norm_cache[key]
128 return _norm_cache[key]
131def normalize(x: torch.Tensor, model_id: str =
"llama-3.2-1b", layer: int |
None =
None) -> torch.Tensor:
135 mean, std = n[
"mean"], n[
"std"]
136 if mean.shape[-1] != x.shape[-1]:
138 return (x - mean) / std
145 sae: SparseAutoencoder,
149 """Zero one SAE feature at a residual position (same logic as label_feature_causally)."""
150 resid_here = value[0, pos]
151 feat_acts = sae.encode(
normalize(resid_here.unsqueeze(0), model_id, sae_layer)).squeeze(0)
152 feat_ablated = feat_acts.clone()
153 feat_ablated[feature_idx] = 0.0
154 delta = sae.decode(feat_ablated) - sae.decode(feat_acts)
156 delta = delta.squeeze(0)
157 value[0, pos] = resid_here + delta
161def label_feature_causally(fi: int, prompt: str, model: HookedTransformer, sae: SparseAutoencoder, client, top_k: int = 5, model_id: str =
"llama-3.2-1b", layer: int |
None =
None) -> str:
162 cfg = get_config(model_id)
163 sae_layer = layer
if layer
is not None else cfg[
"sae_layer"]
165 tokens = model.to_tokens(prompt[:512])
166 if tokens.shape[1] < 4:
167 return f
"feature_{fi}"
169 with torch.no_grad():
170 _, cache = model.run_with_cache(
172 names_filter=f
"blocks.{sae_layer}.hook_resid_post",
175 resid = cache[f
"blocks.{sae_layer}.hook_resid_post"][0]
176 acts = sae.encode(
normalize(resid, model_id, sae_layer))
178 firing_pos = (acts[:, fi] > 0.5).nonzero(as_tuple=
True)[0].tolist()
180 top_pos = int(acts[:, fi].argmax().item())
181 if acts[top_pos, fi].item() < 0.01:
182 return f
"feature_{fi}"
183 firing_pos = [top_pos]
187 for pos
in firing_pos[:3]:
188 act_val = acts[pos, fi].item()
190 with torch.no_grad():
191 baseline_logits = model(tokens)[0, -1]
192 baseline_probs = torch.softmax(baseline_logits, dim=-1)
194 def ablate_feature(value, hook, pos=pos):
197 with torch.no_grad():
198 ablated_logits = model.run_with_hooks(
200 fwd_hooks=[(f
"blocks.{sae_layer}.hook_resid_post", ablate_feature)]
202 ablated_probs = torch.softmax(ablated_logits, dim=-1)
204 delta = baseline_probs - ablated_probs
205 topk_boosted = delta.topk(top_k)
206 topk_suppressed = (-delta).topk(top_k)
209 (model.tokenizer.decode([idx.item()]).strip(), round(val.item(), 4))
210 for idx, val
in zip(topk_boosted.indices, topk_boosted.values)
214 (model.tokenizer.decode([idx.item()]).strip(), round(val.item(), 4))
215 for idx, val
in zip(topk_suppressed.indices, topk_suppressed.values)
219 context = model.tokenizer.decode(tokens[0, max(0, pos - 4):pos + 5].tolist())
220 pivot = model.tokenizer.decode([tokens[0, pos].item()]).strip()
222 causal_examples.append({
223 "activation": round(act_val, 3),
227 "suppresses": suppressed,
230 if not causal_examples:
231 return f
"feature_{fi}"
233 causal_examples.sort(key=
lambda x: x[
"activation"], reverse=
True)
235 formatted =
"\n".join(
236 f
' Context: "...{ex["context"]}..." (token: "{ex["pivot"]}", activation: {ex["activation"]})\n'
237 f
' Causally boosts: {ex["boosts"]}\n'
238 f
' Causally suppresses: {ex["suppresses"]}'
239 for ex
in causal_examples
242 prompt_text = f
"""A sparse autoencoder feature causally influences a language model's predictions as follows:
246Based on what this feature causally promotes and suppresses in this context, give a concise 2-6 word label for its functional role. Good label examples: "promotes plural nouns", "suppresses hedging language", "boosts location names after prepositions".
248Reply with ONLY the label."""
251 return f
"feature_{fi}"
253 resp = client.chat.completions.create(
255 messages=[{
"role":
"user",
"content": prompt_text}],
259 return resp.choices[0].message.content.strip().strip(
'"')
260 except Exception
as e:
261 print(f
"[label] error on feature {fi}: {e}", flush=
True)
262 return f
"feature_{fi}"
265def get_causal_label(fi: int, prompt: str, model: HookedTransformer, sae: SparseAutoencoder, client, model_id: str =
"llama-3.2-1b", layer: int |
None =
None) -> str:
266 cfg = get_config(model_id)
267 resolved_layer = layer
if layer
is not None else cfg[
"sae_layer"]
268 prompt_hash = hashlib.md5(prompt.encode()).hexdigest()[:8]
269 key = (fi, prompt_hash, model_id, resolved_layer)
270 if key
in _session_label_cache:
271 return _session_label_cache[key]
275 _session_label_cache[key] = label
280 """Display form: index plus causal label (never index alone when label is known)."""
282 return f
"{feature_idx} · {label}"
283 return str(feature_idx)
289 mem = ctx.get(
"state", {}).get(
"memory", {})
290 return args.get(
"prompt")
or mem.get(
"lastPrompt")
or "Hello"
297 model: HookedTransformer,
298 sae: SparseAutoencoder,
302 seen: set[int] |
None =
None,
304 fi = int(f[
"feature_idx"])
305 if seen
is not None and fi
in seen
and f.get(
"label"):
308 if not f.get(
"label"):
309 f[
"label"] =
get_causal_label(fi, prompt, model, sae, client, model_id=model_id, layer=layer)
319 model: HookedTransformer,
321 model_id: str =
"llama-3.2-1b",
322 layer: int |
None =
None,
324 """Attach causal labels to inspection feature lists (top + attribution)."""
325 cfg = get_config(model_id)
326 resolved_layer = layer
if layer
is not None else cfg[
"sae_layer"]
327 sae =
load_sae(model_id, resolved_layer)
328 seen: set[int] = set()
330 for f
in feat_result.get(
"top_response_features", []):
332 f, prompt=prompt, model=model, sae=sae, client=client,
333 model_id=model_id, layer=resolved_layer, seen=seen,
335 for attr
in feat_result.get(
"attribution", []):
336 for f
in attr.get(
"driven_by_features", []):
338 f, prompt=prompt, model=model, sae=sae, client=client,
339 model_id=model_id, layer=resolved_layer, seen=seen,
348 model: HookedTransformer,
350 model_id: str =
"llama-3.2-1b",
351 layer: int |
None =
None,
352 label_neighbors: bool =
False,
354 """Add label + feature_ref to a feature-logits or feature-neighbors payload."""
355 cfg = get_config(model_id)
356 resolved_layer = layer
if layer
is not None else cfg[
"sae_layer"]
357 sae =
load_sae(model_id, resolved_layer)
359 fi = result.get(
"feature_idx")
362 int(fi), prompt, model, sae, client, model_id=model_id, layer=resolved_layer,
364 result[
"label"] = label
368 for n
in result.get(
"neighbors", []):
369 nfi = int(n[
"feature_idx"])
371 nfi, prompt, model, sae, client, model_id=model_id, layer=resolved_layer,
383 args: dict |
None =
None,
384 layer: int |
None =
None,
386 """Resolve a causal label using session context (for steer / UI tools)."""
397 or ctx.get(
"state", {}).get(
"activeModelId")
398 or get_active_model_id()
401 model_id = resolve_model_id(model_id)
402 model = get_loaded_model()
404 model = load_model(model_id)
406 cfg = get_config(model_id)
407 resolved_layer = layer
if layer
is not None else cfg[
"sae_layer"]
408 sae =
load_sae(model_id, resolved_layer)
411 int(feature_idx), prompt, model, sae, get_openai_client(ctx),
412 model_id=model_id, layer=resolved_layer,
416def _run_sae_pass(prompt: str, response: str, model: HookedTransformer, top_k: int = TOP_K_FEATURES, model_id: str =
"llama-3.2-1b", layer: int |
None =
None) -> dict:
419 cfg = get_config(model_id)
420 if layer
is not None:
421 sae_layer = require_sae_layer(model_id, int(layer), command=
"trace")
423 sae_layer = int(cfg[
"sae_layer"])
426 full_ctx = f
"{prompt}\n{response}"
427 tokens = model.to_tokens(full_ctx)
429 with torch.no_grad():
430 _, cache = model.run_with_cache(
432 names_filter=f
"blocks.{sae_layer}.hook_resid_post",
435 resid = cache[f
"blocks.{sae_layer}.hook_resid_post"][0]
437 with torch.no_grad():
438 feature_acts = sae.encode(
normalize(resid, model_id, sae_layer))
440 seq_len = resid.shape[0]
441 prompt_ctx = f
"{prompt}\n"
442 prompt_ctx_len = model.to_tokens(prompt_ctx, prepend_bos=
True).shape[1]
444 all_strs = [model.to_string([tokens[0, i].item()])
for i
in range(seq_len)]
445 prompt_strs = all_strs[1:prompt_ctx_len]
446 response_strs = all_strs[prompt_ctx_len:]
447 prompt_idxs = list(range(1, prompt_ctx_len))
448 response_idxs = list(range(prompt_ctx_len, seq_len))
450 def feats_for_unlabeled(positions):
452 for pos
in positions:
453 acts = feature_acts[pos]
454 topk = acts.topk(top_k)
457 "feature_idx": int(i),
458 "activation": round(float(v), 3),
459 "label": f
"feature_{int(i)}",
461 for i, v
in zip(topk.indices, topk.values)
if v > 0.001
465 prompt_features = feats_for_unlabeled(prompt_idxs)
466 response_features = feats_for_unlabeled(response_idxs)
468 resp_acts = feature_acts[prompt_ctx_len:]
469 resp_top = resp_acts.max(0).values.topk(20)
470 top_response_features = []
471 for idx, val
in zip(resp_top.indices.tolist(), resp_top.values.tolist()):
474 best_pos = int(resp_acts[:, idx].argmax().item())
475 top_response_features.append({
477 "activation": round(val, 3),
478 "label": f
"feature_{idx}",
479 "token": response_strs[best_pos].strip()
if best_pos < len(response_strs)
else "",
480 "token_idx": best_pos,
484 for ri, rpos
in enumerate(response_idxs):
485 resp_tok = response_strs[ri].strip()
if ri < len(response_strs)
else ""
486 if not resp_tok
or resp_tok
in (
"the",
"a",
"an",
"is",
"of",
".",
","):
488 resp_feats = feature_acts[rpos]
489 top_resp = resp_feats.topk(top_k)
491 for fidx, fval
in zip(top_resp.indices.tolist(), top_resp.values.tolist()):
494 prompt_feat_acts = feature_acts[prompt_idxs, fidx]
495 active = (prompt_feat_acts > 0.001).nonzero(as_tuple=
True)[0].tolist()
500 "label": f
"feature_{fidx}",
501 "activation": round(fval, 3),
502 "also_in_prompt_positions": active,
503 "also_in_prompt_tokens": [prompt_strs[p].strip()
for p
in active
if p < len(prompt_strs)],
507 "response_token": resp_tok,
509 "driven_by_features": sorted(driven_by, key=
lambda x: x[
"activation"], reverse=
True)[:5],
514 global _kernel_feature_acts, _kernel_resid, _kernel_top_features
515 _kernel_feature_acts = feature_acts.detach().cpu()
516 _kernel_resid = resid.detach().cpu()
517 _kernel_top_features = top_response_features
520 "prompt_tokens": [t.strip()
for t
in prompt_strs],
521 "response_tokens": [t.strip()
for t
in response_strs],
522 "prompt_features": prompt_features,
523 "response_features": response_features,
524 "top_response_features": top_response_features,
525 "attribution": attribution,
526 "sae_layer": sae_layer,
530def run_feature_analysis_unlabeled(prompt: str, response: str, model: HookedTransformer, model_id: str =
"llama-3.2-1b", layer: int |
None =
None) -> dict:
531 return _run_sae_pass(prompt, response, model, model_id=model_id, layer=layer)
534def run_feature_analysis(prompt: str, response: str, model: HookedTransformer, client, top_k: int = TOP_K_FEATURES, model_id: str =
"llama-3.2-1b", layer: int |
None =
None) -> dict:
535 cfg = get_config(model_id)
536 resolved_layer = layer
if layer
is not None else cfg[
"sae_layer"]
537 sae =
load_sae(model_id, resolved_layer)
538 result =
_run_sae_pass(prompt, response, model, top_k, model_id=model_id, layer=resolved_layer)
540 def fill_labels(features_list):
541 for pos_feats
in features_list:
543 fi = int(f[
"feature_idx"])
544 f[
"label"] =
get_causal_label(fi, prompt, model, sae, client, model_id=model_id, layer=resolved_layer)
547 fill_labels(result[
"prompt_features"])
548 fill_labels(result[
"response_features"])
550 for f
in result[
"top_response_features"]:
551 fi = int(f[
"feature_idx"])
552 f[
"label"] =
get_causal_label(fi, prompt, model, sae, client, model_id=model_id, layer=resolved_layer)
555 for attr
in result[
"attribution"]:
556 for f
in attr[
"driven_by_features"]:
557 fi = int(f[
"feature_idx"])
558 f[
"label"] =
get_causal_label(fi, prompt, model, sae, client, model_id=model_id, layer=resolved_layer)
566 model: HookedTransformer,
567 model_id: str =
"llama-3.2-1b",
568 layer: int |
None =
None,
571 """Top vocab tokens boosted/suppressed by an SAE decoder direction (W_dec @ W_U)."""
572 cfg = get_config(model_id)
573 resolved_layer = layer
if layer
is not None else cfg[
"sae_layer"]
574 sae =
load_sae(model_id, resolved_layer)
576 device = model.W_U.device
577 dtype = model.W_U.dtype
578 steer_vec = sae.W_dec[feature_idx].to(device=device, dtype=dtype)
582 with torch.no_grad():
583 logits = project_residual_to_logits(model, steer_vec)
585 top_pos = logits.topk(top_k)
586 top_neg = logits.topk(top_k, largest=
False)
588 def _row(idx: int, val: float) -> dict:
589 token = model.to_string([int(idx)]).strip()
590 return {
"token": token,
"logit": round(float(val), 4)}
592 boosts = [_row(int(i), float(v))
for i, v
in zip(top_pos.indices, top_pos.values)]
593 suppresses = [_row(int(i), float(v))
for i, v
in zip(top_neg.indices, top_neg.values)]
596 "feature_idx": feature_idx,
597 "layer": resolved_layer,
599 "suppresses": suppresses,
601 "bottom": suppresses,
607 model_id: str =
"llama-3.2-1b",
608 layer: int |
None =
None,
611 """Cosine-nearest SAE features in decoder weight space."""
612 cfg = get_config(model_id)
613 resolved_layer = layer
if layer
is not None else cfg[
"sae_layer"]
614 sae =
load_sae(model_id, resolved_layer)
616 with torch.no_grad():
617 W = sae.W_dec.float()
618 W = W / W.norm(dim=-1, keepdim=
True).clamp(min=1e-8)
619 query = W[feature_idx]
621 sims[feature_idx] = -1.0
622 top = sims.topk(top_k)
625 {
"feature_idx": int(i),
"similarity": round(float(s), 4)}
626 for i, s
in zip(top.indices.tolist(), top.values.tolist())
628 return {
"feature_idx": feature_idx,
"layer": resolved_layer,
"neighbors": neighbors}
dict label_inspection_features(dict feat_result, *, str prompt, HookedTransformer model, client, str model_id="llama-3.2-1b", int|None layer=None)
str get_causal_label(int fi, str prompt, HookedTransformer model, SparseAutoencoder sae, client, str model_id="llama-3.2-1b", int|None layer=None)
dict get_feature_logits(int feature_idx, HookedTransformer model, str model_id="llama-3.2-1b", int|None layer=None, int top_k=10)
None _attach_label_to_feature_dict(dict f, *, str prompt, HookedTransformer model, SparseAutoencoder sae, client, str model_id, int|None layer, set[int]|None seen=None)
_get_sae_path_for_layer(str model_id, int|None layer=None)
dict|None load_norm(str model_id="llama-3.2-1b", int|None layer=None)
str label_feature_causally(int fi, str prompt, HookedTransformer model, SparseAutoencoder sae, client, int top_k=5, str model_id="llama-3.2-1b", int|None layer=None)
dict get_feature_neighbors(int feature_idx, str model_id="llama-3.2-1b", int|None layer=None, int top_k=8)
dict run_feature_analysis(str prompt, str response, HookedTransformer model, client, int top_k=TOP_K_FEATURES, str model_id="llama-3.2-1b", int|None layer=None)
dict enrich_feature_tool_result(dict result, *, str prompt, HookedTransformer model, client, str model_id="llama-3.2-1b", int|None layer=None, bool label_neighbors=False)
SparseAutoencoder _load_sae_native(str model_id, int|None layer=None)
torch.Tensor normalize(torch.Tensor x, str model_id="llama-3.2-1b", int|None layer=None)
str format_feature_ref(int feature_idx, str|None label=None)
dict run_feature_analysis_unlabeled(str prompt, str response, HookedTransformer model, str model_id="llama-3.2-1b", int|None layer=None)
torch.Tensor _sae_feature_ablate_hook(torch.Tensor value, int pos, int feature_idx, SparseAutoencoder sae, str model_id, int sae_layer)
_get_norm_path_for_layer(str model_id, int|None layer=None)
SparseAutoencoder load_sae(str model_id="llama-3.2-1b", int|None layer=None)
str resolve_feature_label(int feature_idx, *, dict ctx, dict|None args=None, int|None layer=None)
str prompt_for_labeling(dict|None ctx=None, dict|None args=None)
dict _run_sae_pass(str prompt, str response, HookedTransformer model, int top_k=TOP_K_FEATURES, str model_id="llama-3.2-1b", int|None layer=None)