AQIT 0.1.0
Loading...
Searching...
No Matches
feature_analysis.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/feature_analysis.py.
8Import adaptation only: model_config -> aquin.compute.model_loader, sae -> aquin.compute.sae.
9"""
10from __future__ import annotations
11
12import hashlib
13
14import torch
15from transformer_lens import HookedTransformer
16
17from aquin.compute.sae import SparseAutoencoder
19 get_config,
20 get_config as _get_config,
21 resolve_model_id,
22)
23from aquin.compute.device import resolve_compute_device, synchronize_device
24
25TOP_K_FEATURES = 10
26DEVICE = resolve_compute_device()
27
28_sae_cache: dict[tuple, SparseAutoencoder] = {}
29_norm_cache: dict[tuple, dict | None] = {}
30_session_label_cache = {}
31
32_sae = None
33_norm = None
35_kernel_feature_acts = None
36_kernel_resid = None
37_kernel_top_features = []
38
40def _get_sae_path_for_layer(model_id: str, layer: int | None = None):
41 from aquin.compute.model_loader import resolve_sae_path
42
43 return resolve_sae_path(model_id, layer)
45
46def _get_norm_path_for_layer(model_id: str, layer: int | None = None):
47 from pathlib import Path
48
49 from aquin.compute.model_loader import resolve_model_id, resolve_sae_checkpoint_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"])
54
55 sae_path = resolve_sae_checkpoint_path(short, resolved_layer)
56 if sae_path is not None:
57 parent = sae_path.parent
58 for candidate in (
59 parent / f"norm_layer{resolved_layer}.pt",
60 parent / "norm.pt",
61 parent / f"_acts_layer{resolved_layer}" / "norm.pt",
62 ):
63 if candidate.is_file():
64 return candidate
65
66 base = Path.home() / ".aquin" / "sae" / short
67 if "norm_layers" in cfg:
68 rel = cfg["norm_layers"].get(resolved_layer)
69 if rel:
70 p = base / Path(rel).name
71 if p.exists():
72 return p
73 p_rel = cfg.get("norm_path")
74 if p_rel:
75 p = base / Path(p_rel).name
76 return p if p.exists() else None
77 return None
78
79
80def _load_sae_native(model_id: str, layer: int | None = None) -> SparseAutoencoder:
81 sae_path = _get_sae_path_for_layer(model_id, layer)
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}"
88 )
89 from aquin.compute.model_loader import corrupt_sae_checkpoint_hint
90
91 corrupt = corrupt_sae_checkpoint_hint(short, resolved_layer)
92 if corrupt:
93 raise ValueError(corrupt)
94 return SparseAutoencoder.load(sae_path, device=DEVICE)
95
96
97def load_sae(model_id: str = "llama-3.2-1b", layer: int | None = None) -> SparseAutoencoder:
98 global _sae
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:
104 from aquin.compute.model_loader import load_sae_from_disk
105
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]
110
111
112def load_norm(model_id: str = "llama-3.2-1b", layer: int | None = None) -> dict | None:
113 global _norm
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:
118 norm_path = _get_norm_path_for_layer(model_id, resolved_layer)
119 if norm_path and norm_path.exists():
120 from aquin.compute.torch_io import load_norm_stats
121
122 loaded = load_norm_stats(norm_path, map_location=DEVICE)
123 _norm_cache[key] = loaded
124 else:
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]
129
130
131def normalize(x: torch.Tensor, model_id: str = "llama-3.2-1b", layer: int | None = None) -> torch.Tensor:
132 n = load_norm(model_id, layer)
133 if n is None:
134 return x
135 mean, std = n["mean"], n["std"]
136 if mean.shape[-1] != x.shape[-1]:
137 return x
138 return (x - mean) / std
139
140
142 value: torch.Tensor,
143 pos: int,
144 feature_idx: int,
145 sae: SparseAutoencoder,
146 model_id: str,
147 sae_layer: int,
148) -> torch.Tensor:
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)
155 if delta.dim() > 1:
156 delta = delta.squeeze(0)
157 value[0, pos] = resid_here + delta
158 return value
159
160
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"]
164
165 tokens = model.to_tokens(prompt[:512])
166 if tokens.shape[1] < 4:
167 return f"feature_{fi}"
168
169 with torch.no_grad():
170 _, cache = model.run_with_cache(
171 tokens,
172 names_filter=f"blocks.{sae_layer}.hook_resid_post",
173 return_type=None,
174 )
175 resid = cache[f"blocks.{sae_layer}.hook_resid_post"][0]
176 acts = sae.encode(normalize(resid, model_id, sae_layer))
177
178 firing_pos = (acts[:, fi] > 0.5).nonzero(as_tuple=True)[0].tolist()
179 if not firing_pos:
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]
184
185 causal_examples = []
186
187 for pos in firing_pos[:3]:
188 act_val = acts[pos, fi].item()
189
190 with torch.no_grad():
191 baseline_logits = model(tokens)[0, -1]
192 baseline_probs = torch.softmax(baseline_logits, dim=-1)
193
194 def ablate_feature(value, hook, pos=pos):
195 return _sae_feature_ablate_hook(value, pos, fi, sae, model_id, sae_layer)
196
197 with torch.no_grad():
198 ablated_logits = model.run_with_hooks(
199 tokens,
200 fwd_hooks=[(f"blocks.{sae_layer}.hook_resid_post", ablate_feature)]
201 )[0, -1]
202 ablated_probs = torch.softmax(ablated_logits, dim=-1)
203
204 delta = baseline_probs - ablated_probs
205 topk_boosted = delta.topk(top_k)
206 topk_suppressed = (-delta).topk(top_k)
207
208 boosted = [
209 (model.tokenizer.decode([idx.item()]).strip(), round(val.item(), 4))
210 for idx, val in zip(topk_boosted.indices, topk_boosted.values)
211 if val > 0.001
212 ]
213 suppressed = [
214 (model.tokenizer.decode([idx.item()]).strip(), round(val.item(), 4))
215 for idx, val in zip(topk_suppressed.indices, topk_suppressed.values)
216 if val > 0.001
217 ]
218
219 context = model.tokenizer.decode(tokens[0, max(0, pos - 4):pos + 5].tolist())
220 pivot = model.tokenizer.decode([tokens[0, pos].item()]).strip()
221
222 causal_examples.append({
223 "activation": round(act_val, 3),
224 "context": context,
225 "pivot": pivot,
226 "boosts": boosted,
227 "suppresses": suppressed,
228 })
229
230 if not causal_examples:
231 return f"feature_{fi}"
232
233 causal_examples.sort(key=lambda x: x["activation"], reverse=True)
234
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
240 )
241
242 prompt_text = f"""A sparse autoencoder feature causally influences a language model's predictions as follows:
243
244{formatted}
245
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".
247
248Reply with ONLY the label."""
249
250 if client is None:
251 return f"feature_{fi}"
252 try:
253 resp = client.chat.completions.create(
254 model="gpt-4o-mini",
255 messages=[{"role": "user", "content": prompt_text}],
256 max_tokens=20,
257 temperature=0.0,
258 )
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}"
263
264
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]
272
273 label = label_feature_causally(fi, prompt, model, sae, client, model_id=model_id, layer=resolved_layer)
274
275 _session_label_cache[key] = label
276 return label
277
278
279def format_feature_ref(feature_idx: int, label: str | None = None) -> str:
280 """Display form: index plus causal label (never index alone when label is known)."""
281 if label:
282 return f"{feature_idx} · {label}"
283 return str(feature_idx)
284
285
286def prompt_for_labeling(ctx: dict | None = None, args: dict | None = None) -> str:
287 args = args or {}
288 ctx = ctx or {}
289 mem = ctx.get("state", {}).get("memory", {})
290 return args.get("prompt") or mem.get("lastPrompt") or "Hello"
291
292
294 f: dict,
295 *,
296 prompt: str,
297 model: HookedTransformer,
298 sae: SparseAutoencoder,
299 client,
300 model_id: str,
301 layer: int | None,
302 seen: set[int] | None = None,
303) -> None:
304 fi = int(f["feature_idx"])
305 if seen is not None and fi in seen and f.get("label"):
306 f["feature_ref"] = format_feature_ref(fi, f["label"])
307 return
308 if not f.get("label"):
309 f["label"] = get_causal_label(fi, prompt, model, sae, client, model_id=model_id, layer=layer)
310 f["feature_ref"] = format_feature_ref(fi, f["label"])
311 if seen is not None:
312 seen.add(fi)
313
314
316 feat_result: dict,
317 *,
318 prompt: str,
319 model: HookedTransformer,
320 client,
321 model_id: str = "llama-3.2-1b",
322 layer: int | None = None,
323) -> dict:
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()
329
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,
334 )
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,
340 )
341 return feat_result
342
343
345 result: dict,
346 *,
347 prompt: str,
348 model: HookedTransformer,
349 client,
350 model_id: str = "llama-3.2-1b",
351 layer: int | None = None,
352 label_neighbors: bool = False,
353) -> dict:
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)
358
359 fi = result.get("feature_idx")
360 if fi is not None:
361 label = get_causal_label(
362 int(fi), prompt, model, sae, client, model_id=model_id, layer=resolved_layer,
363 )
364 result["label"] = label
365 result["feature_ref"] = format_feature_ref(int(fi), label)
366
367 if label_neighbors:
368 for n in result.get("neighbors", []):
369 nfi = int(n["feature_idx"])
370 nlabel = get_causal_label(
371 nfi, prompt, model, sae, client, model_id=model_id, layer=resolved_layer,
372 )
373 n["label"] = nlabel
374 n["feature_ref"] = format_feature_ref(nfi, nlabel)
375
376 return result
377
378
380 feature_idx: int,
381 *,
382 ctx: dict,
383 args: dict | None = None,
384 layer: int | None = None,
385) -> str:
386 """Resolve a causal label using session context (for steer / UI tools)."""
387 from aquin.compute.model_loader import (
388 get_active_model_id,
389 get_loaded_model,
390 load_model,
391 resolve_model_id,
392 )
393
394 args = args or {}
395 model_id = (
396 args.get("model_id")
397 or ctx.get("state", {}).get("activeModelId")
398 or get_active_model_id()
399 or "llama-3.2-1b"
400 )
401 model_id = resolve_model_id(model_id)
402 model = get_loaded_model()
403 if model is None:
404 model = load_model(model_id)
405 prompt = prompt_for_labeling(ctx, args)
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)
409 from aquin.compute.openai_client import get_openai_client
410 return get_causal_label(
411 int(feature_idx), prompt, model, sae, get_openai_client(ctx),
412 model_id=model_id, layer=resolved_layer,
413 )
414
415
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:
417 from aquin.compute.model_loader import require_sae_layer
418
419 cfg = get_config(model_id)
420 if layer is not None:
421 sae_layer = require_sae_layer(model_id, int(layer), command="trace")
422 else:
423 sae_layer = int(cfg["sae_layer"])
424 sae = load_sae(model_id, sae_layer)
425
426 full_ctx = f"{prompt}\n{response}"
427 tokens = model.to_tokens(full_ctx)
428
429 with torch.no_grad():
430 _, cache = model.run_with_cache(
431 tokens,
432 names_filter=f"blocks.{sae_layer}.hook_resid_post",
433 return_type=None,
434 )
435 resid = cache[f"blocks.{sae_layer}.hook_resid_post"][0]
436
437 with torch.no_grad():
438 feature_acts = sae.encode(normalize(resid, model_id, sae_layer))
439
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]
443
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))
449
450 def feats_for_unlabeled(positions):
451 out = []
452 for pos in positions:
453 acts = feature_acts[pos]
454 topk = acts.topk(top_k)
455 out.append([
456 {
457 "feature_idx": int(i),
458 "activation": round(float(v), 3),
459 "label": f"feature_{int(i)}",
460 }
461 for i, v in zip(topk.indices, topk.values) if v > 0.001
462 ])
463 return out
464
465 prompt_features = feats_for_unlabeled(prompt_idxs)
466 response_features = feats_for_unlabeled(response_idxs)
467
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()):
472 if val < 0.001:
473 continue
474 best_pos = int(resp_acts[:, idx].argmax().item())
475 top_response_features.append({
476 "feature_idx": idx,
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,
481 })
482
483 attribution = []
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", ".", ","):
487 continue
488 resp_feats = feature_acts[rpos]
489 top_resp = resp_feats.topk(top_k)
490 driven_by = []
491 for fidx, fval in zip(top_resp.indices.tolist(), top_resp.values.tolist()):
492 if fval < 0.001:
493 continue
494 prompt_feat_acts = feature_acts[prompt_idxs, fidx]
495 active = (prompt_feat_acts > 0.001).nonzero(as_tuple=True)[0].tolist()
496 if not active:
497 continue
498 driven_by.append({
499 "feature_idx": fidx,
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)],
504 })
505 if driven_by:
506 attribution.append({
507 "response_token": resp_tok,
508 "response_ti": ri,
509 "driven_by_features": sorted(driven_by, key=lambda x: x["activation"], reverse=True)[:5],
510 })
511
512 synchronize_device()
513
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
518
519 return {
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,
527 }
528
529
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)
532
533
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)
539
540 def fill_labels(features_list):
541 for pos_feats in features_list:
542 for f in pos_feats:
543 fi = int(f["feature_idx"])
544 f["label"] = get_causal_label(fi, prompt, model, sae, client, model_id=model_id, layer=resolved_layer)
545 f["feature_ref"] = format_feature_ref(fi, f["label"])
546
547 fill_labels(result["prompt_features"])
548 fill_labels(result["response_features"])
549
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)
553 f["feature_ref"] = format_feature_ref(fi, f["label"])
554
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)
559 f["feature_ref"] = format_feature_ref(fi, f["label"])
560
561 return result
562
563
565 feature_idx: int,
566 model: HookedTransformer,
567 model_id: str = "llama-3.2-1b",
568 layer: int | None = None,
569 top_k: int = 10,
570) -> dict:
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)
575
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)
579
580 from aquin.compute.hf_llm_shim import project_residual_to_logits
581
582 with torch.no_grad():
583 logits = project_residual_to_logits(model, steer_vec)
584
585 top_pos = logits.topk(top_k)
586 top_neg = logits.topk(top_k, largest=False)
587
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)}
591
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)]
594
595 return {
596 "feature_idx": feature_idx,
597 "layer": resolved_layer,
598 "boosts": boosts,
599 "suppresses": suppresses,
600 "top": boosts,
601 "bottom": suppresses,
602 }
603
604
606 feature_idx: int,
607 model_id: str = "llama-3.2-1b",
608 layer: int | None = None,
609 top_k: int = 8,
610) -> dict:
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)
615
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]
620 sims = W @ query
621 sims[feature_idx] = -1.0
622 top = sims.topk(top_k)
623
624 neighbors = [
625 {"feature_idx": int(i), "similarity": round(float(s), 4)}
626 for i, s in zip(top.indices.tolist(), top.values.tolist())
627 ]
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)