File size: 20,321 Bytes
a2ffd07 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 | """
Per-layer top-k SAE feature selection.
Two modes
---------
probe — load `probes_all_layers.pt` (state-dict with keys
`probes.{L}.weight` shape (1, d_sae) and `probes.{L}.bias`).
For each layer, take the top-k features by *positive* probe weight.
f1 — score each SAE feature individually as a 1-D classifier on the
labelled dataset used to train the probes. Sweeps thresholds per
feature and keeps the best F1 (across both directions:
`act > τ → pos` and `act < τ → pos`).
Always runs LLaVA forward + SAE encode + max-pool over tokens
inline (no disk cache). Pooling is always over the assistant
caption only — the last ``caption_text_len`` residual positions
of the teacher-forced forward.
Output JSON (`--out`) format
---------------------------
{
"layer_0": {"features": [int, ...]},
...
"layer_31": {"features": [int, ...]},
"_meta": {mode, top_k, ...}
}
The JSON consumed by `delta_w_feature_trace.py`.
"""
from __future__ import annotations
import argparse
import json
import os
import re
import sys
from typing import Dict, List
# Make `hallucination.*` (the repo treated as a package) importable so
# downstream modules like `sae.SAE_Tools` can resolve their own imports.
_PARENT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
if _PARENT not in sys.path:
sys.path.insert(0, _PARENT)
_REPO = os.path.dirname(os.path.abspath(__file__))
_REPO = os.path.dirname(_REPO)
if _REPO not in sys.path:
sys.path.insert(0, _REPO)
import numpy as np
import torch as t
from tqdm import tqdm
_PROBE_KEY_RE = re.compile(r"probes\.(\d+)\.weight$")
def _make_hook(act_buf: dict, name: str):
"""Factory for hook function that captures activation to act_buf."""
def _fn(act, hook):
act_buf[name] = act.detach()
return _fn
def _parse_probe_layers(state_dict: dict) -> Dict[int, t.Tensor]:
"""{layer_idx: weight tensor of shape (d_sae,)} from a state-dict."""
out = {}
for k, v in state_dict.items():
m = _PROBE_KEY_RE.match(k)
if m is None:
continue
layer = int(m.group(1))
# Squeeze the leading 1 of a binary-probe weight.
out[layer] = v.squeeze(0).float().cpu()
return out
# ── Mode: probe ─────────────────────────────────────────────────────────────
def select_by_probe(probes_path: str, top_k: int) -> Dict[str, dict]:
"""
Loads probe weights and selects top-k features.
If probes_path is a directory, it expects one file per layer (probe_{hook}_{L}.pth).
If it's a file, it expects a multi-layer state-dict.
"""
weights = {}
if os.path.isdir(probes_path):
# Folder mode: Expects files like probe_0.pth, probe_1.pth ...
files = sorted(
[f for f in os.listdir(probes_path) if f.endswith(".pth")],
key=lambda x: int(re.search(r"(\d+)\.pth$", x).group(1)) if re.search(r"(\d+)\.pth$", x) else 0
)
for f in files:
path = os.path.join(probes_path, f)
sd = t.load(path, map_location="cpu", weights_only=False)
# If the file is a state_dict from DDP or a custom module, it might be nested
if "module" in sd:
sd = sd["module"]
# Find the weight tensor. Probe could be saved as:
# * nn.Linear -> keys: 'weight', 'bias'
# * nn.Sequential(fc) -> keys: 'fc.weight', 'fc.bias'
w = None
if "weight" in sd:
w = sd["weight"]
elif "fc.weight" in sd:
w = sd["fc.weight"]
elif len(sd) == 1:
w = list(sd.values())[0]
if w is not None:
# Extract layer index from filename probe_{hook}_{L}.pth
m = re.search(r"(\d+)\.pth$", f)
if m:
layer = int(m.group(1))
weights[layer] = w.squeeze().float().cpu()
else:
# File mode: multi-layer state-dict
sd = t.load(probes_path, map_location="cpu", weights_only=False)
weights = _parse_probe_layers(sd)
if not weights:
raise RuntimeError(f"No `probes.{{L}}.weight` keys found in {probes_path}")
out = {}
for layer, w in sorted(weights.items()):
# Most positive entries -> strongest features
topk = t.topk(w, k=min(top_k, w.numel()), largest=True).indices.tolist()
out[f"layer_{layer}"] = {"features": [int(i) for i in topk]}
return out
# ── Mode: f1 ────────────────────────────────────────────────────────────────
def _per_feature_best_f1(acts: t.Tensor, labels: t.Tensor) -> t.Tensor:
"""For each column j of `acts` (N, F), compute the best F1 over all
thresholds on `acts[:, j]` using only the high→pos direction
(act > τ predicts positive). Returns a length-F float tensor."""
acts_np = acts.cpu().numpy().astype(np.float32)
y = labels.cpu().numpy().astype(np.int64).reshape(-1)
N, F = acts_np.shape
pos = int(y.sum())
if pos == 0 or pos == N:
return t.zeros(F)
out = np.zeros(F, dtype=np.float32)
eps = 1e-12
for j in tqdm(range(F), desc="F1 sweep", leave=False):
a = acts_np[:, j]
order = np.argsort(-a) # descending: highest activation first
sl = y[order]
cum_tp = np.cumsum(sl)
cum_pred_pos = np.arange(1, N + 1)
prec = cum_tp / (cum_pred_pos + eps)
rec = cum_tp / (pos + eps)
f1 = 2 * prec * rec / (prec + rec + eps)
out[j] = float(f1.max())
return t.from_numpy(out)
def _compute_features_inline(
args, hook_points: List[str], img_ids_labels: List[tuple],
) -> Dict[str, tuple]:
"""Run LLaVA forward + SAE encode + max-pool over tokens directly.
Pipeline per image (mirrors ``Train_Probe_SAE.phase1_cache``):
* Greedy-decode a caption from the LLaVA base model.
* Teacher-force ``USER: <image>\nDescribe this image. \nASSISTANT: {caption}``
with hooks at every requested layer, capturing residuals.
* SAE-encode each layer's residual and max-pool over tokens
(uses ``Train_Probe_SAE.sae_encode_and_pool``).
"""
from PIL import Image
from datasets import load_dataset
from transformers import LlavaProcessor
from model.llava.hooked_llava import HookedSAELlavaConditionalGeneration
from sae.SAE_Tools import load_sae_model
from training.Train_Probe_SAE import sae_encode_and_pool
dtype = {"float32": t.float32, "float16": t.float16, "bfloat16": t.bfloat16}[args.dtype]
device = args.device
print("Loading model …")
processor = LlavaProcessor.from_pretrained(args.model_name)
model = HookedSAELlavaConditionalGeneration.from_pretrained(
args.model_name, attn_implementation="eager",
).to(device, dtype=dtype).eval()
print("Loading SAE …")
sae = load_sae_model(
args.sae_ckpt, model_type="llava", hook_type="text", device=device,
)
# Image index: HF dataset images (positives) keyed by image_id, plus
# fall-back lookup in the CC3M folders for negatives.
hf_index: Dict[str, "Image.Image"] = {}
for split in ("train", "validation"):
ds = load_dataset(args.hf_dataset, split=split)
for row in tqdm(ds, desc=f"indexing HF:{split}", leave=False):
img = row["image"]
if not isinstance(img, Image.Image):
img = Image.open(img)
hf_index[str(row[args.id_col])] = img.convert("RGB")
def _resolve_image(sid: str):
img = hf_index.get(sid)
if img is not None:
return img
for folder in (args.neg_train_dir, args.neg_val_dir):
if not folder:
continue
p = os.path.join(folder, f"{sid}.jpg")
if os.path.exists(p):
try:
return Image.open(p).convert("RGB")
except Exception:
return None
return None
feats: Dict[str, list] = {hp: [] for hp in hook_points}
labels: Dict[str, list] = {hp: [] for hp in hook_points}
n_resolved = n_missing = 0
prompt_q = "USER: <image>\nDescribe this image. \nASSISTANT:"
# Prompt text token count (constant across rows). With right-padding on
# the teacher-forced forward, per-row caption tokens occupy positions
# [cap_start, cap_end) where:
# cap_start = T_act_max - text_max + K (constant in the batch)
# cap_end_b = T_act_max - (text_max - text_len_b)
K = len(processor.tokenizer(prompt_q)["input_ids"])
img_batch = max(1, int(args.img_batch))
orig_padding_side = processor.tokenizer.padding_side
for start in tqdm(range(0, len(img_ids_labels), img_batch), desc="forward+SAE"):
chunk = img_ids_labels[start : start + img_batch]
sids, labs, imgs = [], [], []
for img_id, label in chunk:
sid = str(img_id)
img = _resolve_image(sid)
if img is None:
n_missing += 1
continue
n_resolved += 1
sids.append(sid); labs.append(label); imgs.append(img)
if not sids:
continue
B = len(sids)
# Batched caption generation — left-pad for HF generate.
processor.tokenizer.padding_side = "left"
prompt_inputs = processor(
images=imgs, text=[prompt_q] * B,
return_tensors="pt", padding=True,
).to(device)
with t.no_grad():
outputs = model.generate(
**prompt_inputs, do_sample=False, num_beams=1,
use_cache=True, max_new_tokens=args.max_new_tokens,
)
decoded = processor.batch_decode(outputs, skip_special_tokens=True)
asst_list = [c.split("ASSISTANT:")[-1].strip() for c in decoded]
keep = [i for i, a in enumerate(asst_list) if a]
if not keep:
continue
forced_texts = [
f"USER: <image>\nDescribe this image. \nASSISTANT: {asst_list[i]}"
for i in keep
]
f_imgs = [imgs[i] for i in keep]
f_labs = [labs[i] for i in keep]
# Batched teacher-forced forward — right-pad so caption positions
# are derivable from the attention_mask alone.
processor.tokenizer.padding_side = "right"
forced_inputs = processor(
images=f_imgs, text=forced_texts,
return_tensors="pt", padding=True,
).to(device)
text_lens = forced_inputs.attention_mask.sum(dim=1).tolist()
text_max = max(text_lens)
keep2 = [i for i, tl in enumerate(text_lens) if tl - K > 0]
if not keep2:
continue
act_buf: dict = {}
with t.no_grad():
model.run_with_hooks(
forced_inputs,
fwd_hooks=[(hp, _make_hook(act_buf, hp)) for hp in hook_points],
)
present = [hp for hp in hook_points if hp in act_buf]
if not present:
continue
T_act_max = act_buf[present[0]].shape[1]
cap_start = T_act_max - text_max + K # constant across the batch
for b in keep2:
cap_end = T_act_max - (text_max - text_lens[b])
acts_list_b = [act_buf[hp][b][cap_start:cap_end] for hp in present]
pooled_list = sae_encode_and_pool(
acts_list_b, sae, args.sae_batch, device,
)
for hp, pooled in zip(present, pooled_list):
feats[hp].append(pooled)
labels[hp].append(f_labs[b])
processor.tokenizer.padding_side = orig_padding_side
print(f"inline forward done: resolved={n_resolved}, missing={n_missing}")
if n_missing > 0 and n_resolved == 0:
raise RuntimeError("All images missing — check neg_train_dir/neg_val_dir paths")
out: Dict[str, tuple] = {}
for hp in hook_points:
if not feats[hp]:
continue
out[hp] = (t.stack(feats[hp]), t.tensor(labels[hp]))
return out
def select_by_f1(args, layers: List[int]) -> Dict[str, dict]:
"""Score each SAE feature by best-threshold F1 against the
``--probe_type`` column of ``--hf_dataset`` (binary 0/1 label).
Works for any object — ``toilet``, ``bathroom``, or whatever 0/1 column
the HF dataset exposes. Always runs LLaVA forward + SAE encode inline
(see ``_compute_features_inline``)."""
from datasets import load_dataset
from training.Train_Probe_SAE import build_hook_points
# Match Train_Probe_SAE.py's split logic.
train_ds = load_dataset(args.hf_dataset, split="train")
val_ds = load_dataset(args.hf_dataset, split="validation")
# Positives: probe_type=1 regardless of other_object value.
label_col = args.probe_type
train_pos = [row[args.id_col] for row in train_ds if row[label_col] == 1]
val_pos = [row[args.id_col] for row in val_ds if row[label_col] == 1]
# HF negatives: all probe_type=0 rows, or contrastive subset if other_object given.
if args.all_negatives:
train_hf_neg = [row[args.id_col] for row in train_ds if row[label_col] == 0]
val_hf_neg = [row[args.id_col] for row in val_ds if row[label_col] == 0]
elif args.other_object:
train_hf_neg = [row[args.id_col] for row in train_ds
if row[args.other_object] == 1 and row[label_col] == 0]
val_hf_neg = [row[args.id_col] for row in val_ds
if row[args.other_object] == 1 and row[label_col] == 0]
else:
train_hf_neg, val_hf_neg = [], []
# JSON random negatives.
train_json_neg, val_json_neg = [], []
if args.neg_jsonl:
with open(args.neg_jsonl) as f:
neg = json.load(f)
train_json_neg = neg.get("train", [])
val_json_neg = neg.get("validation", [])
train_neg = train_hf_neg + train_json_neg
val_neg = val_hf_neg + val_json_neg
# Combine train+val so we score features across as much labelled data as
# we have; the F1 here is a feature-quality proxy, not a held-out metric.
img_ids_labels = (
[(i, 1) for i in train_pos + val_pos]
+ [(i, 0) for i in train_neg + val_neg]
)
print(
f"F1 dataset: pos={len(train_pos) + len(val_pos)} "
f"neg={len(train_neg) + len(val_neg)} "
f"(hf_contrastive={len(train_hf_neg) + len(val_hf_neg)}, "
f"json={len(train_json_neg) + len(val_json_neg)}) "
f"total={len(img_ids_labels)}"
)
hook_points = build_hook_points(layers, args.hook_type)
feats = _compute_features_inline(args, hook_points, img_ids_labels)
out = {}
for hp, layer in zip(hook_points, layers):
if hp not in feats:
print(f" layer {layer}: no activations for {hp}; skipping")
continue
data, lbl = feats[hp] # (N, d_sae), (N,)
f1 = _per_feature_best_f1(data, lbl)
topk = t.topk(f1, k=min(args.top_k, f1.numel()), largest=True).indices.tolist()
out[f"layer_{layer}"] = {"features": [int(i) for i in topk]}
print(f" layer {layer}: top F1 = {f1.max().item():.4f} | top features = {topk[:5]}...")
return out
# ── Main ────────────────────────────────────────────────────────────────────
def main():
p = argparse.ArgumentParser()
p.add_argument("--mode", choices=["probe", "f1"], default="probe")
p.add_argument("--top_k", type=int, default=20)
p.add_argument(
"--probes_path",
default="mechanistic_interp/probes/probes_all_layers.pt",
help="(probe mode) Path to the multi-layer probe state-dict.",
)
p.add_argument("--out", required=True, help="Output JSON path.")
p.add_argument(
"--label", default=None,
help="Object/concept this feature set describes (e.g. 'toilet', "
"'bathroom', 'microwave'). Recorded in the output JSON under "
"`_meta.object` so downstream consumers can route feature sets "
"per object/relation. Defaults to --probe_type in f1 mode; "
"required (or 'unknown') in probe mode if you want the field "
"populated.",
)
# F1-mode-only options
p.add_argument("--sae_ckpt", default=None, help="(f1) SAE checkpoint.")
p.add_argument("--hf_dataset", default="pbcong/bathroom-toilet")
p.add_argument("--id_col", default="image_id")
p.add_argument("--probe_type", default="toilet")
p.add_argument("--neg_jsonl", default=None, help="(f1) {'train':[...], 'validation':[...]} JSON.")
p.add_argument("--other_object", default=None,
help="(f1) Contrastive column: negatives = other_object=1 & probe_type=0. "
"E.g. 'bathroom' when --probe_type toilet.")
p.add_argument("--all_negatives", action="store_true",
help="(f1) Use ALL probe_type=0 rows as negatives (overrides --other_object).")
p.add_argument("--hook_type", default="post", choices=["pre", "mid", "post"])
p.add_argument("--sae_batch", type=int, default=2048)
p.add_argument("--device", default="cuda:0")
p.add_argument(
"--layers", type=int, nargs="+", default=None,
help="(f1) Layers to score. Defaults to 0..31 if omitted.",
)
# F1-mode inline-forward options.
p.add_argument("--neg_train_dir", default="CC3M-Dataset/cc3m_images/train",
help="(f1 inline) Folder for train negatives (image_id.jpg).")
p.add_argument("--neg_val_dir", default="CC3M-Dataset/cc3m_images/val",
help="(f1 inline) Folder for validation negatives.")
p.add_argument("--model_name", default="llava-hf/llava-1.5-7b-hf",
help="(f1 inline) LLaVA checkpoint to forward.")
p.add_argument("--dtype", default="bfloat16",
choices=["float32", "float16", "bfloat16"],
help="(f1 inline) Model dtype.")
p.add_argument("--max_new_tokens", type=int, default=200,
help="(f1 inline) Caption budget for the base decode.")
p.add_argument("--img_batch", type=int, default=1,
help="(f1 inline) Images per LLaVA forward batch. "
"Generate uses left-padding; teacher-forced forward "
"uses right-padding with attention-mask-derived "
"caption windows.")
args = p.parse_args()
if args.mode == "probe":
result = select_by_probe(args.probes_path, args.top_k)
else:
if args.sae_ckpt is None:
p.error("--mode f1 requires --sae_ckpt")
if not args.all_negatives and not args.other_object and not args.neg_jsonl:
p.error("--mode f1 requires --all_negatives, --neg_jsonl, and/or --other_object")
layers = args.layers or list(range(32))
result = select_by_f1(args, layers)
# Default the recorded label to --probe_type in f1 mode (it IS the
# object column); leave it explicit-or-unknown in probe mode since the
# state-dict is opaque about its target.
object_label = args.label or (args.probe_type if args.mode == "f1" else "unknown")
result["_meta"] = {
"mode": args.mode,
"top_k": args.top_k,
"object": object_label,
"probes_path": args.probes_path if args.mode == "probe" else None,
"hf_dataset": args.hf_dataset if args.mode == "f1" else None,
"hook_type": args.hook_type,
"sae_ckpt": args.sae_ckpt,
}
os.makedirs(os.path.dirname(os.path.abspath(args.out)) or ".", exist_ok=True)
with open(args.out, "w") as f:
json.dump(result, f, indent=2)
n_layers = sum(1 for k in result if k.startswith("layer_"))
print(f"Wrote {n_layers} layers x top-{args.top_k} features → {args.out}")
if __name__ == "__main__":
main()
|