| """ |
| Compare editing-method effects on the toilet readout, per layer. |
| |
| Idea |
| ──── |
| For each model variant (base, lora, efuf, nullu) read the *toilet* attention |
| probe score σ(logit) ∈ [0,1] at every layer, on bathroom-only images, and compare |
| each method to base: |
| |
| cell[method, l] = mean_images( σ_toilet_l(method) − σ_toilet_l(base) ) |
| |
| NO gradient steering — the *edit itself* is the intervention. To isolate the |
| representational effect of the edit (not a change in what the model says), every |
| method reads the SAME fixed context per image: the BASE model's generated caption |
| (or --forced_text). The probe pools over those caption tokens; only the model |
| weights differ between rows. The base row is therefore exactly 0 (reference). |
| |
| Images: bathroom-only (bathroom=1 & toilet=0), split into two populations by the |
| base model's hallucination flag (base_mentions_object on --base_prompt): |
| • NON-hallucinating (base did NOT say toilet) |
| • HALLUCINATING (base DID say toilet) |
| One figure, two panels (non | halluc), each a (methods × layers) heatmap; base − base = 0. |
| |
| Models are loaded ONE AT A TIME (base first, to fix the captions), so we never hold |
| four 7B models on the GPU simultaneously. |
| |
| Saves: heatmap PNG + the raw (methods × layers) matrices as JSON. |
| """ |
|
|
| import argparse |
| import gc |
| import json |
| import os |
| import random |
|
|
| import numpy as np |
| import torch as t |
| from PIL import Image |
| from transformers import AutoConfig, LlavaProcessor |
|
|
| import matplotlib |
| matplotlib.use("Agg") |
| import matplotlib.pyplot as plt |
|
|
| from model.llava.hooked_llava import HookedSAELlavaConditionalGeneration |
| from mechanistic_interp.sequence_probe import sequence_layer_probes_from_checkpoint |
|
|
| _HOOK_SUFFIX = {"pre": "hook_resid_pre", "mid": "hook_resid_mid", "post": "hook_resid_post"} |
|
|
|
|
| |
| |
| |
| _EFUF_PROJ_KEYS = { |
| "model.mm_projector.0.weight": ("linear_1", "weight"), |
| "model.mm_projector.0.bias": ("linear_1", "bias"), |
| "model.mm_projector.2.weight": ("linear_2", "weight"), |
| "model.mm_projector.2.bias": ("linear_2", "bias"), |
| } |
| |
| _DOWN_PROJ_PREFIXES = ("model.layers", "language_model.model.layers", |
| "language_model.layers", "model.language_model.layers") |
| _DOWN_PROJ_SUFFIX = ".mlp.down_proj.weight" |
|
|
|
|
| def apply_efuf_edit(model, ckpt_path): |
| """Overwrite the multimodal projector in-place from an EFUF .pth checkpoint.""" |
| raw = t.load(ckpt_path, map_location="cpu", weights_only=True) |
| missing = [k for k in _EFUF_PROJ_KEYS if k not in raw] |
| if missing: |
| raise KeyError(f"EFUF ckpt {ckpt_path} missing {missing}; has {list(raw.keys())}") |
| proj = model.multi_modal_projector |
| with t.no_grad(): |
| for key, (lin, attr) in _EFUF_PROJ_KEYS.items(): |
| tgt = getattr(getattr(proj, lin), attr) |
| tgt.data.copy_(raw[key].to(device=tgt.device, dtype=tgt.dtype)) |
| return len(_EFUF_PROJ_KEYS) |
|
|
|
|
| def _match_down_proj_layer(key): |
| """Return the layer index if `key` is an mlp.down_proj.weight, else None.""" |
| if not key.endswith(_DOWN_PROJ_SUFFIX): |
| return None |
| for prefix in _DOWN_PROJ_PREFIXES: |
| if key.startswith(prefix + "."): |
| mid = key[len(prefix) + 1: -len(_DOWN_PROJ_SUFFIX)] |
| if mid.isdigit(): |
| return int(mid) |
| return None |
|
|
|
|
| def apply_nullu_edit(model, lm_layers, ckpt_dir, lowest, highest): |
| """Splice Nullu-edited mlp.down_proj.weight into layers [lowest, highest) in-place.""" |
| import glob |
| from safetensors.torch import safe_open |
| wanted = set(range(lowest, highest)) |
| found = {} |
| shards = sorted(glob.glob(os.path.join(ckpt_dir, "*.safetensors"))) |
| if shards: |
| for sh in shards: |
| with safe_open(sh, framework="pt") as f: |
| for key in f.keys(): |
| idx = _match_down_proj_layer(key) |
| if idx is not None and idx in wanted and idx not in found: |
| found[idx] = f.get_tensor(key) |
| if len(found) == len(wanted): |
| break |
| else: |
| for bf in sorted(glob.glob(os.path.join(ckpt_dir, "*.bin"))): |
| sd = t.load(bf, map_location="cpu", weights_only=True) |
| for key, tns in sd.items(): |
| idx = _match_down_proj_layer(key) |
| if idx is not None and idx in wanted and idx not in found: |
| found[idx] = tns |
| if len(found) == len(wanted): |
| break |
| missing = wanted - set(found) |
| if missing: |
| raise KeyError(f"Nullu dir {ckpt_dir} missing down_proj for layers {sorted(missing)}") |
| with t.no_grad(): |
| for idx, w in found.items(): |
| tgt = lm_layers[idx].mlp.down_proj.weight |
| tgt.data.copy_(w.to(device=tgt.device, dtype=tgt.dtype)) |
| return len(found) |
|
|
|
|
| |
| def _decoder_layers(model): |
| lm = getattr(model, "language_model", None) or getattr(getattr(model, "model", None), "language_model", None) |
| inner = getattr(lm, "model", None) |
| if inner is not None and hasattr(inner, "layers"): |
| return inner.layers |
| return lm.layers |
|
|
|
|
| def build_model(variant, args, dtype, device): |
| """Build one model variant. PROBES are loaded separately (base-trained); only the |
| MODEL weights change here. Caller is responsible for freeing the previous model.""" |
| if variant == "lora": |
| from model.llava.hooked_lora_llava import HookedLoRALlava |
| model = HookedLoRALlava.from_pretrained( |
| args.model_name, torch_dtype=dtype, device_map={"": device}).eval() |
| lora_dir = args.lora_path if os.path.isdir(args.lora_path) else os.path.dirname(args.lora_path) |
| model.load_lora_adapter(lora_dir, merge=True) |
| print(f"[compare] model=lora (merged {lora_dir})") |
| else: |
| model = HookedSAELlavaConditionalGeneration.from_pretrained( |
| args.model_name, torch_dtype=dtype, device_map={"": device}).eval() |
| if variant == "efuf": |
| n = apply_efuf_edit(model, args.efuf_path) |
| print(f"[compare] model=efuf ({n} proj tensors, {args.efuf_path})") |
| elif variant == "nullu": |
| n = apply_nullu_edit(model, _decoder_layers(model), args.nullu_path, |
| args.nullu_lowest, args.nullu_highest) |
| print(f"[compare] model=nullu ({n} layers [{args.nullu_lowest},{args.nullu_highest}), {args.nullu_path})") |
| else: |
| print(f"[compare] model=base") |
| return model |
|
|
|
|
| def free_model(model): |
| del model |
| gc.collect() |
| if t.cuda.is_available(): |
| t.cuda.empty_cache() |
|
|
|
|
| def hp_name(layer: int, hook_type: str) -> str: |
| return f"model.language_model.layers.{layer}.{_HOOK_SUFFIX[hook_type]}" |
|
|
|
|
| def caption_slice(attn_mask, asst_text, processor, max_seq_tokens): |
| """(start, end) indices of the caption tokens at the tail of the real region, |
| matching training's a[:max_seq_tokens].""" |
| seq_len = int(attn_mask.sum().item()) |
| cap_len = len(processor.tokenizer(asst_text, add_special_tokens=False)["input_ids"]) |
| if cap_len <= 0: |
| return None |
| start = seq_len - cap_len |
| end = start + min(cap_len, max_seq_tokens) |
| if start < 0: |
| return None |
| return start, end |
|
|
|
|
| @t.no_grad() |
| def generate_caption(model, processor, image, question, device, max_new_tokens): |
| prompt = f"USER: <image>\n{question}\nASSISTANT:" |
| inp = processor(images=[image], text=[prompt], return_tensors="pt").to(device) |
| out = model.generate(**inp, do_sample=False, num_beams=1, use_cache=True, |
| max_new_tokens=max_new_tokens) |
| cap = processor.batch_decode(out, skip_special_tokens=True)[0] |
| return cap.split("ASSISTANT:")[-1].strip() |
|
|
|
|
| def probe_logit(probe_module, layer, seq_feats): |
| """Scalar logit for one image from one layer's attention probe. seq_feats: (1,T,d).""" |
| kpm = t.zeros(seq_feats.shape[:2], dtype=t.bool, device=seq_feats.device) |
| return probe_module.probes[probe_module._idx[layer]](seq_feats, kpm).squeeze(0) |
|
|
|
|
| def probe_score(probe_module, layer, seq_feats): |
| """Bounded probe score σ(logit) = P(concept present) ∈ [0,1].""" |
| return t.sigmoid(probe_logit(probe_module, layer, seq_feats)) |
|
|
|
|
| @t.no_grad() |
| def toilet_scores_per_layer(model, processor, toilet, img, asst, question, |
| device, hps, layers, max_seq_tokens): |
| """Forward `img` with the fixed ASSISTANT answer `asst`, return σ_toilet at every |
| layer (np array of length n_layers), or None if the caption slice is invalid.""" |
| forced = f"USER: <image>\n{question}\nASSISTANT: {asst}" |
| fwd = processor(images=[img], text=[forced], return_tensors="pt").to(device) |
| sl = caption_slice(fwd["attention_mask"], asst, processor, max_seq_tokens) |
| if sl is None: |
| return None |
| s0, s1 = sl |
| acts = {} |
| def cap(name): |
| def _fn(act, hook): |
| acts[name] = act |
| return _fn |
| model.run_with_hooks(fwd, fwd_hooks=[(hp, cap(hp)) for hp in hps]) |
| out = np.empty(len(layers), dtype=np.float64) |
| for l in layers: |
| feats = acts[hps[l]][:, s0:s1].float() |
| out[l] = float(probe_score(toilet, l, feats).item()) |
| return out |
|
|
|
|
| def plot_method_layer(mats, methods, layers, out, title, cbar_label="Δ toilet score σ (method − base)"): |
| """mats: {group_name: (n_methods, n_layers) np array}. One panel per group, shared |
| symmetric diverging scale + single colorbar. methods on y, layers on x.""" |
| groups = [g for g in mats if mats[g] is not None] |
| ncol = len(groups) |
| fig, axes = plt.subplots(1, ncol, figsize=(0.30 * len(layers) + 2.0, 0.55 * len(methods) + 2.0), |
| squeeze=False) |
| allv = np.concatenate([mats[g][np.isfinite(mats[g])].ravel() for g in groups]) |
| vmax = float(np.nanmax(np.abs(allv))) if allv.size else 1.0 |
| vmax = max(vmax, 1e-6) |
| im = None |
| for k, g in enumerate(groups): |
| ax = axes[0][k] |
| im = ax.imshow(mats[g], origin="upper", cmap="RdBu_r", vmin=-vmax, vmax=vmax, aspect="auto") |
| ax.set_title(g) |
| ax.set_xlabel("layer") |
| ax.set_yticks(range(len(methods))) |
| ax.set_yticklabels(methods) |
| ax.set_xticks(range(0, len(layers), 4)) |
| cbar = fig.colorbar(im, ax=axes, fraction=0.025, pad=0.02) |
| cbar.set_label(cbar_label) |
| fig.suptitle(title, fontsize=12) |
| os.makedirs(os.path.dirname(out), exist_ok=True) |
| fig.savefig(out, dpi=150, bbox_inches="tight") |
| plt.close(fig) |
| print(f"[compare] saved {out}") |
|
|
|
|
| def main(): |
| ap = argparse.ArgumentParser(description=__doc__, |
| formatter_class=argparse.RawDescriptionHelpFormatter) |
| 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("--methods", nargs="+", default=["base", "lora", "efuf", "nullu"], |
| choices=["base", "lora", "efuf", "nullu"], |
| help="Variants to compare. 'base' is the reference (its row is all 0); " |
| "it is always loaded first to fix the captions even if omitted here.") |
|
|
| |
| 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("--toilet_probe", default="/data/caotue/latent_probes/seqprobes_toilet/post/seqprobe.pth") |
|
|
| ap.add_argument("--image_folder", default="/data/caotue/CC3M-Dataset/cc3m_images") |
| ap.add_argument("--samples_json", default="mechanistic_interp/toilet_bathroom/samples.json") |
| ap.add_argument("--category", default="bathroom_only", |
| choices=["bathroom_only", "bathroom_toilet", "toilet_only", "others"], |
| help="Dataset to evaluate. bathroom_only = toilet ABSENT (base mention = " |
| "hallucination); bathroom_toilet / toilet_only = toilet PRESENT (Δσ<0 = " |
| "edit suppresses a real toilet); others = unrelated negatives from " |
| "--others_jsonl (single 'unrelated' panel; Δσ≠0 = edit spuriously moves toilet).") |
| ap.add_argument("--others_jsonl", default="mechanistic_interp/neg_cc3m_5k.json", |
| help="For --category others: {train:[ids], validation:[ids]} of CC3M stems.") |
| ap.add_argument("--others_split", default="validation", choices=["validation", "train"], |
| help="Which split of --others_jsonl to use (default validation = held out " |
| "from the 4-variant probe's negatives).") |
| ap.add_argument("--base_prompt", default="Describe this image.", |
| help="prompt_results key whose base_mentions_object flag splits the two populations.") |
| ap.add_argument("--population", default="both", |
| choices=["all", "both", "said", "unsaid", "halluc", "non"], |
| help="How to group the category: 'all' = ONE set, no base-mention split; " |
| "said/halluc = base DID mention toilet; unsaid/non = did not; both = both panels.") |
| ap.add_argument("--id_col", default="image_id") |
| ap.add_argument("--question", default="Describe this image.", |
| help="The text fed to every model (USER turn).") |
| ap.add_argument("--forced_text", default="This image features a bathroom with a", |
| help="Constant ASSISTANT answer used as the fixed context for ALL " |
| "images/methods (the probe reads these tokens). Set to '' to instead " |
| "use the base model's per-image generated caption.") |
| ap.add_argument("--num_images", type=int, default=100, |
| help="Random-subsample cap PER population (0 = ALL). Shuffled by --seed.") |
| ap.add_argument("--seed", type=int, default=0) |
|
|
| ap.add_argument("--hook_type", default="post", choices=["pre", "mid", "post"]) |
| ap.add_argument("--max_new_tokens", type=int, default=64) |
| ap.add_argument("--max_seq_tokens", type=int, default=64) |
|
|
| ap.add_argument("--out", default="mechanistic_interp/graph/compare_baselines_toilet.png") |
| ap.add_argument("--out_json", default="mechanistic_interp/graph/compare_baselines_toilet.json") |
| 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" |
|
|
| n_layers = AutoConfig.from_pretrained(args.model_name).text_config.num_hidden_layers |
| layers = list(range(n_layers)) |
| hps = [hp_name(l, args.hook_type) for l in layers] |
|
|
| |
| methods = ["base"] + [m for m in args.methods if m != "base"] |
| print(f"[compare] device={device} dtype={dtype} layers={n_layers}") |
| print(f"[compare] methods={methods} readout=toilet ({args.toilet_probe})") |
| print(f"[compare] context={'forced: ' + repr(args.forced_text) if args.forced_text else 'base caption per image'}") |
|
|
| |
| stem_to_file = {} |
| for root, _, files in os.walk(args.image_folder): |
| for fn in files: |
| if fn.lower().endswith((".jpg", ".jpeg", ".png", ".webp")): |
| stem_to_file[os.path.splitext(fn)[0]] = os.path.join(root, fn) |
|
|
| if args.category == "others": |
| |
| neg = json.load(open(args.others_jsonl)) |
| raw = neg.get(args.others_split, []) |
| ids = [s for s in (os.path.splitext(os.path.basename(x))[0] for x in raw) if s in stem_to_file] |
| groups = {"unrelated": ids} |
| else: |
| data = json.load(open(args.samples_json)) |
| cat_items = [it for it in data if it.get("category") == args.category] |
| if args.population == "all": |
| |
| ids = [s for it in cat_items |
| if (s := os.path.splitext(os.path.basename(it[args.id_col]))[0]) in stem_to_file] |
| groups = {args.category: ids} |
| else: |
| |
| |
| |
| has_toilet = args.category != "bathroom_only" |
| lbl = ({True: "toilet-said (correct)", False: "toilet-missed"} if has_toilet |
| else {True: "hallucinating", False: "non-hallucinating"}) |
| groups = {lbl[True]: [], lbl[False]: []} |
| for it in cat_items: |
| flag = it["prompt_results"].get(args.base_prompt, {}).get("base_mentions_object") |
| if flag not in (True, False): |
| continue |
| stem = os.path.splitext(os.path.basename(it[args.id_col]))[0] |
| if stem in stem_to_file: |
| groups[lbl[flag]].append(stem) |
| sel = {"both": [True, False], "said": [True], "halluc": [True], |
| "unsaid": [False], "non": [False]}[args.population] |
| groups = {lbl[f]: groups[lbl[f]] for f in sel} |
| for g in groups: |
| if args.num_images and args.num_images > 0: |
| random.Random(args.seed).shuffle(groups[g]) |
| groups[g] = groups[g][:args.num_images] |
| print(f"[compare] {args.category} images: " + ", ".join(f"{g}={len(ids)}" for g, ids in groups.items())) |
|
|
| toilet = sequence_layer_probes_from_checkpoint(args.toilet_probe, device) |
| toilet.eval() |
| processor = LlavaProcessor.from_pretrained(args.model_name) |
|
|
| |
| caption_by_stem = {} |
| img_cache = {} |
| per_img = {m: {} for m in methods} |
|
|
| def load_img(stem): |
| if stem not in img_cache: |
| try: |
| img_cache[stem] = Image.open(stem_to_file[stem]).convert("RGB") |
| except Exception: |
| img_cache[stem] = None |
| return img_cache[stem] |
|
|
| for m in methods: |
| model = build_model(m, args, dtype, device) |
| for g, ids in groups.items(): |
| for i, stem in enumerate(ids): |
| img = load_img(stem) |
| if img is None: |
| continue |
| if m == "base": |
| asst = args.forced_text or generate_caption( |
| model, processor, img, args.question, device, args.max_new_tokens) |
| caption_by_stem[stem] = asst |
| else: |
| asst = caption_by_stem.get(stem) |
| if asst is None: |
| continue |
| scores = toilet_scores_per_layer(model, processor, toilet, img, asst, |
| args.question, device, hps, layers, args.max_seq_tokens) |
| if scores is not None: |
| per_img[m][stem] = scores |
| print(f"[compare] {m}: scored {len([s for s in ids if s in per_img[m]])}/{len(ids)} ({g})") |
| free_model(model) |
|
|
| |
| mats = {} |
| for g, ids in groups.items(): |
| if not ids: |
| mats[g] = None |
| continue |
| M = np.full((len(methods), n_layers), np.nan) |
| for mi, m in enumerate(methods): |
| paired = [per_img[m][s] - per_img["base"][s] |
| for s in ids if s in per_img[m] and s in per_img["base"]] |
| if paired: |
| M[mi] = np.mean(paired, axis=0) |
| mats[g] = M |
|
|
| |
| base_abs = {} |
| for g, ids in groups.items(): |
| arr = [per_img["base"][s] for s in ids if s in per_img["base"]] |
| base_abs[g] = (np.mean(arr, axis=0).tolist() if arr else None) |
|
|
| |
| os.makedirs(os.path.dirname(args.out_json), exist_ok=True) |
| with open(args.out_json, "w") as f: |
| json.dump({"methods": methods, "n_layers": n_layers, "hook_type": args.hook_type, |
| "context": args.forced_text if args.forced_text else "base_caption", |
| "n_images": {g: len(ids) for g, ids in groups.items()}, |
| "base_abs_toilet": base_abs, |
| "delta": {g: (None if mats[g] is None else |
| np.where(np.isnan(mats[g]), None, mats[g]).tolist()) |
| for g in groups}}, f) |
| print(f"[compare] saved {args.out_json}") |
|
|
| |
| ctx = "forced" if args.forced_text else "base caption" |
| plot_method_layer( |
| mats, methods, layers, args.out, |
| title=f"toilet readout: method − base, per layer — {args.category} ({ctx} context, {args.hook_type})") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|