""" 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"]) # Model variant the activations come from — MUST match how the probe was trained # (e.g. seqprobes_512_*_lora ⇒ --variant lora). HF ground-truth labels (samples.json # base_mentions is base-only and invalid for edited variants). 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: # Behavioral label: did the BASE model mention the object? (base_mentions_object) 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: \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()