| """ |
| Per-layer validation BCE for a trained SequenceLayerProbes checkpoint. |
| |
| Training logged BCE averaged over all layers; this recomputes per-layer |
| BCEWithLogits on the validation split so each layer's loss is available. |
| |
| Val recipe mirrors train_probe_latent.py: |
| toilet : pos = toilet==1 (HF val); neg = neg_cc3m_5k.json validation |
| bathroom: pos = bathroom==1 (HF val); neg = toilet-only (HF val) + JSON validation |
| """ |
| import argparse |
| import json |
| import os |
|
|
| import numpy as np |
| import torch as t |
| import torch.nn.functional as F |
| from PIL import Image |
| from datasets import load_dataset |
| from sklearn.metrics import ( |
| roc_auc_score, f1_score, precision_score, recall_score, confusion_matrix, |
| ) |
|
|
| from transformers import LlavaProcessor |
| from mechanistic_interp.sequence_probe import sequence_layer_probes_from_checkpoint |
| from mechanistic_interp.gradient_ascent import ( |
| hp_name, caption_slice, generate_caption, probe_logit, load_variant_model, |
| ) |
|
|
|
|
| def main(): |
| ap = argparse.ArgumentParser() |
| ap.add_argument("--probe", required=True, help="seqprobe.pth checkpoint") |
| ap.add_argument("--object", required=True, choices=["toilet", "bathroom"]) |
| ap.add_argument("--model_name", default="llava-hf/llava-1.5-7b-hf") |
| ap.add_argument("--device_id", type=int, default=0) |
| ap.add_argument("--dtype", default="bfloat16", choices=["float32", "float16", "bfloat16"]) |
| |
| |
| |
| ap.add_argument("--variant", default="base", choices=["base", "lora", "nullu", "efuf"]) |
| ap.add_argument("--lora_path", default="/data/caotue/multilayer-sae/adv_gen_outputs/run_bathroom_toilet_v2/lora_adapter") |
| ap.add_argument("--efuf_path", default="/data/caotue/multilayer-sae/EFUF/efuf/checkpoints/llava_vicuna_7b/bathroom_toilet_paper_10ep/epoch_002.pth") |
| ap.add_argument("--nullu_path", default="/data/caotue/nullu/edited_models/LLaVA-7B-top4-0-32-bathroom_toilet") |
| ap.add_argument("--nullu_lowest", type=int, default=8) |
| ap.add_argument("--nullu_highest", type=int, default=32) |
| ap.add_argument("--hf_dataset", default="pbcong/bathroom-toilet") |
| ap.add_argument("--neg_jsonl", default="mechanistic_interp/neg_cc3m_5k.json") |
| ap.add_argument("--samples_json", default=None, |
| help="If set, validate against samples.json with label = base_mentions_object " |
| "(did the BASE model mention the object?) for --base_prompt, instead of the " |
| "HF ground-truth pos/neg split.") |
| ap.add_argument("--base_prompt", default="Describe this image.", |
| help="prompt_results key whose base_mentions_object is the label (samples.json mode).") |
| ap.add_argument("--image_folder", default="/data/caotue/CC3M-Dataset/cc3m_images") |
| ap.add_argument("--question", default="Describe this image.") |
| ap.add_argument("--max_new_tokens", type=int, default=256) |
| ap.add_argument("--max_seq_tokens", type=int, default=64) |
| ap.add_argument("--hook_type", default="post", choices=["pre", "mid", "post"]) |
| ap.add_argument("--max_val", type=int, default=0, |
| help="Cap val images (random, balanced shuffle); 0 = all.") |
| ap.add_argument("--seed", type=int, default=0) |
| args = ap.parse_args() |
|
|
| dtype = {"float32": t.float32, "float16": t.float16, "bfloat16": t.bfloat16}[args.dtype] |
| device = f"cuda:{args.device_id}" if t.cuda.is_available() else "cpu" |
|
|
| stem_to_file = {} |
| for r, _, fs in os.walk(args.image_folder): |
| for fn in fs: |
| if fn.lower().endswith((".jpg", ".jpeg", ".png", ".webp")): |
| stem_to_file[os.path.splitext(fn)[0]] = os.path.join(r, fn) |
|
|
| def inf(ids): |
| return [s for s in (os.path.splitext(os.path.basename(i))[0] for i in ids) if s in stem_to_file] |
|
|
| if args.samples_json: |
| |
| data = json.load(open(args.samples_json)) |
| ids_labels = [] |
| for it in data: |
| pr = it.get("prompt_results", {}).get(args.base_prompt, {}) |
| bmo = pr.get("base_mentions_object") |
| if bmo is None: |
| continue |
| stem = os.path.splitext(os.path.basename(it["image_id"]))[0] |
| if stem in stem_to_file: |
| ids_labels.append((stem, int(bool(bmo)))) |
| src = f"samples.json[base_mentions_object @ '{args.base_prompt}']" |
| else: |
| val = load_dataset(args.hf_dataset, split="validation") |
| other = "bathroom" if args.object == "toilet" else "toilet" |
| pos = inf([row["image_id"] for row in val if row[args.object] == 1]) |
| jneg = inf(json.load(open(args.neg_jsonl)).get("validation", [])) |
| if args.object == "bathroom": |
| toilet_only = inf([row["image_id"] for row in val if row[other] == 1 and row[args.object] == 0]) |
| neg = toilet_only + jneg |
| else: |
| neg = jneg |
| ids_labels = [(s, 1) for s in pos] + [(s, 0) for s in neg] |
| src = f"HF {args.hf_dataset}[validation] ground-truth {args.object}" |
| if args.max_val and len(ids_labels) > args.max_val: |
| import random |
| random.Random(args.seed).shuffle(ids_labels) |
| ids_labels = ids_labels[: args.max_val] |
| npos = sum(1 for _, y in ids_labels if y == 1) |
| print(f"[bce] {args.object} [{src}]: val {npos} pos + {len(ids_labels)-npos} neg = {len(ids_labels)}") |
|
|
| model = load_variant_model(args, dtype, device) |
| processor = LlavaProcessor.from_pretrained(args.model_name) |
| probe = sequence_layer_probes_from_checkpoint(args.probe, device) |
| probe.eval() |
| layers = probe.layer_indices |
| hps = [hp_name(l, args.hook_type) for l in layers] |
|
|
| logits = {l: [] for l in layers} |
| ys = [] |
| for k, (stem, y) in enumerate(ids_labels): |
| try: |
| img = Image.open(stem_to_file[stem]).convert("RGB") |
| except Exception: |
| continue |
| asst = generate_caption(model, processor, img, args.question, device, args.max_new_tokens) |
| forced = f"USER: <image>\n{args.question}\nASSISTANT: {asst}" |
| fwd = processor(images=[img], text=[forced], return_tensors="pt").to(device) |
| sl = caption_slice(fwd["attention_mask"], asst, processor, args.max_seq_tokens) |
| if sl is None: |
| continue |
| s0, s1 = sl |
| acts = {} |
| with t.no_grad(): |
| model.run_with_hooks(fwd, fwd_hooks=[(hp, (lambda a, hook, n=hp: acts.__setitem__(n, a))) for hp in hps]) |
| with t.no_grad(): |
| for l in layers: |
| logits[l].append(float(probe_logit(probe, l, acts[hps[l]][:, s0:s1].float()).item())) |
| ys.append(y) |
| if (k + 1) % 100 == 0: |
| print(f"[bce] {k+1}/{len(ids_labels)}") |
|
|
| y = t.tensor(ys, dtype=t.float32) |
| y_np = y.numpy() |
| print(f"\nlayer Acc AUC F1 Prec Recall BCE") |
| out = {} |
| for l in layers: |
| z = t.tensor(logits[l]) |
| bce = float(F.binary_cross_entropy_with_logits(z, y).item()) |
| pred = (z > 0).float().numpy() |
| acc = float((pred == y_np).mean()) |
| try: |
| auc = roc_auc_score(y_np, z.numpy()) |
| except ValueError: |
| auc = float("nan") |
| f1 = float(f1_score(y_np, pred, zero_division=0)) |
| prec = float(precision_score(y_np, pred, zero_division=0)) |
| rec = float(recall_score(y_np, pred, zero_division=0)) |
| out[l] = {"accuracy": acc, "auc": auc, "f1": f1, |
| "precision": prec, "recall": rec, "bce": bce} |
| print(f"{l:5d} {acc:.4f} {auc:.4f} {f1:.4f} {prec:.4f} {rec:.4f} {bce:.4f}") |
| tag = "_basemention_metrics" if args.samples_json else "_perlayer_metrics" |
| op = os.path.splitext(args.probe)[0] + tag + ".json" |
| json.dump({"object": args.object, "variant": args.variant, "n": len(ys), |
| "label_source": src, "n_pos": int(y.sum().item()), "per_layer": out}, |
| open(op, "w"), indent=2) |
| print(f"saved {op}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|