hallucination / experiment /evaluation /bleu_from_cache.py
ToiTenBao's picture
Upload hallucination folder
a2ffd07 verified
Raw
History Blame Contribute Delete
7.92 kB
"""
Compute BLEU-2 (edited vs base) per category from cached captions.json files.
Uses a captions.json that contains base_caption as the reference source,
then joins any other captions.json on (image_id, prompt) to get the
edited captions. No model inference needed.
Usage:
python bleu_from_cache.py \
--base_captions adv_outputs/run_20260421_191124/lora_adapter/captions.json \
--inputs \
"LoRA:lora-baseline/runs/run_20260430_135824/step_500/captions.json" \
"Nullu:Nullu/output/edited_model/LLaVA-7B-top4-0-32-bathroom_toilet/captions.json" \
"EFUF-ep5:EFUF/efuf/checkpoints/.../eval_epoch_005/captions.json"
"""
from __future__ import annotations
import argparse
import json
import math
import sys
from collections import defaultdict
from pathlib import Path
# ---------------------------------------------------------------------------
# BLEU-2 helpers (no external deps; nltk used if available)
# ---------------------------------------------------------------------------
def _make_scorer():
try:
from nltk.translate.bleu_score import sentence_bleu, SmoothingFunction
_smooth = SmoothingFunction().method1
def score(ref: str, hyp: str) -> float:
r, h = ref.lower().split(), hyp.lower().split()
if not r or not h:
return float("nan")
return sentence_bleu([r], h, weights=(0.5, 0.5), smoothing_function=_smooth)
except ImportError:
def score(ref: str, hyp: str) -> float:
# fallback: token-level F1 (unigram)
r, h = set(ref.lower().split()), set(hyp.lower().split())
if not r or not h:
return float("nan")
inter = len(r & h)
p, rec = inter / len(h), inter / len(r)
return 2 * p * rec / (p + rec) if (p + rec) > 0 else 0.0
return score
def _detect_caption_key(record: dict) -> str | None:
for key in ("lora_caption", "edited_caption", "EFUF_caption", "base_caption"):
if key in record and record[key]:
return key
return None
def _category_from_flags(rec: dict) -> str:
"""Derive canonical category name from boolean flags, independent of display naming."""
# Try bathroom/toilet first, then generic scene/object keys
scene_keys = [k for k in rec if k not in ("index", "image_id", "prompt", "category") and
not k.endswith("_caption") and not k.endswith("_mentions_object") and
isinstance(rec[k], int)]
if "bathroom" in rec and "toilet" in rec:
sc, ob = int(rec["bathroom"]), int(rec["toilet"])
elif len(scene_keys) == 2:
sc, ob = int(rec[scene_keys[0]]), int(rec[scene_keys[1]])
else:
return rec.get("category", "unknown")
if sc == 0 and ob == 1:
return "non_scene_with_object"
if sc == 1 and ob == 0:
return "scene_no_object"
if sc == 1 and ob == 1:
return "scene_with_object"
return "neither"
CAT_LABELS = {
"non_scene_with_object": "non_scene_w_object",
"scene_no_object": "scene_no_object (suppression target — lower BLEU expected)",
"scene_with_object": "scene_with_object (key quality metric)",
"neither": "neither",
}
# ---------------------------------------------------------------------------
# Core computation
# ---------------------------------------------------------------------------
def compute_bleu(base_lookup: dict, edited_records: list, scorer) -> dict[str, list[float]]:
"""Return dict of category -> list of BLEU scores."""
by_cat: dict[str, list[float]] = defaultdict(list)
missing = 0
for rec in edited_records:
key = (rec["image_id"], rec["prompt"])
base_rec = base_lookup.get(key)
if base_rec is None:
missing += 1
continue
base_cap = base_rec.get("base_caption", "")
edited_key = _detect_caption_key(rec)
if not edited_key or edited_key == "base_caption":
# skip if only base present
continue
edited_cap = rec[edited_key]
cat = _category_from_flags(rec)
s = scorer(base_cap, edited_cap)
if not math.isnan(s):
by_cat[cat].append(s)
if missing:
print(f" [warn] {missing} records had no matching base caption", file=sys.stderr)
return dict(by_cat)
def mean(vals: list[float]) -> float:
return sum(vals) / len(vals) if vals else float("nan")
# ---------------------------------------------------------------------------
# Main
# ---------------------------------------------------------------------------
def parse_args():
p = argparse.ArgumentParser(description="Compute BLEU-2 from cached captions.json files.")
p.add_argument(
"--base_captions",
default="adv_outputs/run_20260421_191124/lora_adapter/captions.json",
help="captions.json that contains base_caption field.",
)
p.add_argument(
"--inputs", nargs="+", metavar="LABEL:PATH",
help="One or more 'Label:path/to/captions.json' entries.",
)
p.add_argument("--output_json", default=None, help="Optional path to save results as JSON.")
return p.parse_args()
def main():
args = parse_args()
scorer = _make_scorer()
base_path = Path(args.base_captions)
print(f"Loading base captions from: {base_path}")
with open(base_path) as f:
base_data = json.load(f)
base_lookup = {(r["image_id"], r["prompt"]): r for r in base_data}
print(f" {len(base_lookup)} (image_id, prompt) pairs\n")
# Also compute BLEU for the base file's own edited captions (serves as self-check / LoRA result)
inputs: list[tuple[str, Path]] = []
if args.inputs:
for entry in args.inputs:
label, _, path = entry.partition(":")
inputs.append((label.strip(), Path(path.strip())))
# If the base file has an edited caption key too, add it automatically
sample = base_data[0]
base_edited_key = _detect_caption_key({k: v for k, v in sample.items() if k != "base_caption"})
if base_edited_key and base_edited_key != "base_caption":
inputs = [(f"[base file] {base_edited_key}", base_path)] + inputs
all_results: dict[str, dict] = {}
cat_order = ["non_scene_with_object", "scene_no_object", "scene_with_object", "neither"]
for label, path in inputs:
print(f"Processing: {label}")
with open(path) as f:
data = json.load(f)
by_cat = compute_bleu(base_lookup, data, scorer)
row = {cat: mean(by_cat.get(cat, [])) for cat in cat_order}
all_results[label] = row
# Print table
col_w = max(len(l) for l, _ in inputs) + 2 if inputs else 30
col_w = max(col_w, 20)
cats_display = [
("non_scene_with_object", "non_scene_w_obj"),
("scene_no_object", "scene_no_obj "),
("scene_with_object", "scene_with_obj "),
]
header = f"\n{'Method':<{col_w}}" + "".join(f" {c[1]}" for c in cats_display)
print(header)
print("-" * len(header))
for label, row in all_results.items():
line = f"{label:<{col_w}}"
for cat, _ in cats_display:
v = row.get(cat, float("nan"))
line += f" {v:>15.4f}" if not math.isnan(v) else f" {'n/a':>15}"
print(line)
print()
print("Note: scene_no_object = hallucination target (lower BLEU is EXPECTED for good suppression).")
print(" scene_with_object and non_scene_w_obj = quality metrics (higher = less collateral damage).")
if args.output_json:
out = {label: {cat: (None if math.isnan(v) else v) for cat, v in row.items()}
for label, row in all_results.items()}
with open(args.output_json, "w") as f:
json.dump(out, f, indent=2)
print(f"\nSaved to {args.output_json}")
if __name__ == "__main__":
main()