ToiTenBao's picture
Upload hallucination folder
a2ffd07 verified
Raw
History Blame Contribute Delete
32.8 kB
"""
Build the knowledge-editing dataset from scene/object relation data.
Generates an edit_set.json that all KME methods (EasyEdit baselines + ours)
can consume. The file is method-agnostic — each method reads the parts it
needs.
Supports any relation defined in experiment/config/relations.json.
Use --relation to select (default: bathroom_toilet).
The PRIMARY edit framing is captioning:
Image + "Describe this image."
old: "A bathroom with a toilet, sink, and mirror" (hallucinated)
new: "A bathroom with a sink and mirror" (object removed)
The target for each edit instance is the original model's own caption with
object mentions surgically removed. This is a minimal edit — the model's
style, vocabulary, and all correct content are preserved.
Pipeline:
Step 1 (no GPU): Build structure from CSV or HuggingFace
Step 2 (GPU): Generate original captions → clean them → fill targets
Usage:
# HuggingFace dataset (default)
python -m experiment.knowledge_editing.build_edit_set \
--relation bathroom_toilet \
--output experiment/knowledge_editing/edit_set.json
# With a different relation
python -m experiment.knowledge_editing.build_edit_set \
--relation kitchen_microwave \
--output experiment/knowledge_editing/edit_set_kitchen_microwave.json
# Legacy CSV path
python -m experiment.knowledge_editing.build_edit_set \
--csv CC3M-Dataset/bathroom_filter/bathroom_toilet_labels.csv \
--image_dir CC3M-Dataset/cc3m_images/train \
--output experiment/knowledge_editing/edit_set.json
"""
import argparse
import csv
import json
import os
import random
import re
import sys
from typing import Optional
from experiment.config.relation_config import get_relation_config, RelationConfig
from experiment.data.hf_loader import HF_DATASET_ID, hf_rows as _hf_rows
# ---------------------------------------------------------------------------
# Constants — shared with training and evaluation pipelines
# ---------------------------------------------------------------------------
# Legacy CSV-only constants (kept for backward compatibility)
SPLIT_SEED = 42
SPLIT_TEST_SIZE = 0.2
# Default cap for evaluation set (per category, from HF val split)
DEFAULT_EVAL_PER_CATEGORY = 50
# Legacy defaults (bathroom_toilet). Overridden by RelationConfig at runtime.
CAPTION_PROMPT = "In this bathroom there is"
TRAIN_PROMPTS = [
"In this bathroom there is",
]
GENERALITY_PROMPTS = [
"This bathroom contains",
"In this bathroom I can see",
"The objects in this bathroom are",
]
TOILET_KEYWORDS = [
"toilet", "toilets", "Toilet", "Toilets",
"commode", "lavatory", "latrine",
]
def _build_object_re(keywords: list[str]) -> re.Pattern:
"""Build a regex that matches any of the given keywords (case-insensitive)."""
return re.compile(
r'\b(?:' + '|'.join(re.escape(k) for k in keywords) + r')s?\b',
re.IGNORECASE,
)
# Default regex (backward compat)
_TOILET_RE = _build_object_re(TOILET_KEYWORDS)
# ---------------------------------------------------------------------------
# Caption cleaning
# ---------------------------------------------------------------------------
def clean_object_mentions(text: str, object_re: re.Pattern = None) -> str:
"""Remove object mentions from a caption, cleaning up grammar artifacts.
Args:
text: Caption text to clean.
object_re: Compiled regex matching the object keywords.
Defaults to _TOILET_RE for backward compat.
"""
if object_re is None:
object_re = _TOILET_RE
cleaned = object_re.sub("", text)
# Fix grammar artifacts from removal
cleaned = re.sub(r'\ba\s+,', ',', cleaned) # "a , sink" → ", sink"
cleaned = re.sub(r',\s*,', ',', cleaned) # ",, sink" → ", sink"
cleaned = re.sub(r',\s*and\s*,', ',', cleaned) # ", and ," → ","
cleaned = re.sub(r',\s*\.', '.', cleaned) # ",." → "."
cleaned = re.sub(r'\.\s*\.', '.', cleaned) # ".." → "."
cleaned = re.sub(r'\bwith\s*,', 'with', cleaned) # "with , sink" → "with sink"
cleaned = re.sub(r'\bwith\s+and\b', 'with', cleaned) # "with and sink" → "with sink"
cleaned = re.sub(r'\band\s+and\b', 'and', cleaned) # "and and" → "and"
cleaned = re.sub(r'\ba\s+and\b', 'a', cleaned) # "a and sink" → "a sink"
cleaned = re.sub(r',\s+and\s*$', '', cleaned) # trailing ", and"
cleaned = re.sub(r',\s*$', '.', cleaned) # trailing comma
cleaned = re.sub(r'\s{2,}', ' ', cleaned) # double spaces
cleaned = cleaned.strip().strip(',').strip()
return cleaned
def has_substance(text: str, min_words: int = 4) -> bool:
"""Check if a cleaned caption still has enough content to be useful."""
words = text.split()
return len(words) >= min_words
# ---------------------------------------------------------------------------
# CSV loading + category splitting
# ---------------------------------------------------------------------------
def load_csv(csv_path: str, image_dir: str) -> list[dict]:
"""Load dataset rows from a CSV + image directory (legacy path)."""
rows = []
missing = 0
with open(csv_path, "r") as f:
reader = csv.DictReader(f)
for row in reader:
image_path = os.path.join(image_dir, f"{row['image_id']}.jpg")
if not os.path.exists(image_path):
missing += 1
continue
rows.append({
"image_id": row["image_id"],
"bathroom": int(row.get("bathroom", 0)),
"toilet": int(row.get("toilet", 0)),
"image_path": image_path,
})
print(f"Loaded {len(rows)} rows from CSV ({missing} images missing)")
return rows
def split_categories(rows, relation_config: RelationConfig = None):
"""Split rows into the four evaluation categories.
Uses generic is_scene/has_object keys from hf_rows, or legacy
bathroom/toilet keys from CSV loading.
"""
if relation_config is not None:
cat_names = relation_config.category_names
else:
cat_names = ["bathroom_no_toilet", "bathroom_with_toilet",
"non_bathroom_with_toilet", "unrelated"]
cats = {name: [] for name in cat_names}
for row in rows:
# Support both generic (is_scene/has_object) and legacy (bathroom/toilet) keys
b = row.get("is_scene", row.get("bathroom", 0))
t = row.get("has_object", row.get("toilet", 0))
if b == 1 and t == 0:
cats[cat_names[0]].append(row)
elif b == 1 and t == 1:
cats[cat_names[1]].append(row)
elif b == 0 and t == 1:
cats[cat_names[2]].append(row)
else:
cats[cat_names[3]].append(row)
for k, v in cats.items():
print(f" {k}: {len(v)}")
return cats
def _csv_train_val_split(image_ids: list[str]):
"""Legacy 80/20 split for CSV-only path (no HF split info available)."""
from sklearn.model_selection import train_test_split
train_ids, val_ids = train_test_split(
image_ids, test_size=SPLIT_TEST_SIZE, random_state=SPLIT_SEED,
)
return set(train_ids), set(val_ids)
# ---------------------------------------------------------------------------
# Target generation (requires GPU)
# ---------------------------------------------------------------------------
def generate_captions(
image_sources: dict[str, object],
model_name: str,
prompt: str = CAPTION_PROMPT,
device: str = "cuda",
batch_size: int = 1,
object_re: re.Pattern = None,
) -> dict[str, dict]:
"""Run the original LLaVA model to generate per-image captions.
For each image, produces:
original_caption: what the unedited model says (may hallucinate object)
cleaned_caption: original with object mentions removed
had_object: whether the original mentioned the object
is_usable: whether the cleaned version has enough content
Args:
image_sources: dict mapping image_id → file path (str) or PIL.Image
model_name: HuggingFace model ID
prompt: the captioning prompt
device: cuda device
object_re: Compiled regex for matching object keywords.
Returns:
dict mapping image_id → {original, cleaned, had_toilet, is_usable}
"""
if object_re is None:
object_re = _TOILET_RE
import torch
from PIL import Image
from transformers import AutoProcessor, AutoModelForPreTraining
from tqdm import tqdm
print(f"\nGenerating captions with {model_name} on {len(image_sources)} images...")
processor = AutoProcessor.from_pretrained(model_name)
model = AutoModelForPreTraining.from_pretrained(
model_name, torch_dtype=torch.float16, device_map={"": device},
)
model.eval()
results = {}
items = list(image_sources.items())
for image_id, source in tqdm(items, desc="Captioning"):
try:
if isinstance(source, str):
image = Image.open(source).convert("RGB")
else:
image = source.convert("RGB")
except Exception as e:
print(f" Skipping {image_id}: {e}")
continue
inputs = processor(
images=image,
text=f"<image>\nUSER: {prompt}\nASSISTANT:",
return_tensors="pt",
).to(device)
with torch.no_grad():
output_ids = model.generate(
**inputs, max_new_tokens=256, do_sample=False,
)
# Decode only generated tokens
generated = processor.decode(
output_ids[0][inputs["input_ids"].shape[1]:],
skip_special_tokens=True,
).strip()
cleaned = clean_object_mentions(generated, object_re=object_re)
had_object = bool(object_re.search(generated))
results[image_id] = {
"original": generated,
"cleaned": cleaned,
"had_toilet": had_object, # key kept for backward compat
"is_usable": has_substance(cleaned),
}
del model
torch.cuda.empty_cache()
# Stats
n_had_object = sum(1 for r in results.values() if r["had_toilet"])
n_usable = sum(1 for r in results.values() if r["is_usable"])
print(f" Generated {len(results)} captions")
print(f" {n_had_object}/{len(results)} mentioned object (hallucinated)")
print(f" {n_usable}/{len(results)} usable after cleaning")
return results
def generate_locality_captions(
image_sources: dict[str, object],
model_name: str,
prompt: str = CAPTION_PROMPT,
device: str = "cuda",
) -> dict[str, str]:
"""Generate original-model captions for locality images.
These serve as the ground-truth reference for locality evaluation:
the edited model's output on these images should match the original's.
Args:
image_sources: dict mapping image_id → file path (str) or PIL.Image
"""
import torch
from PIL import Image
from transformers import AutoProcessor, AutoModelForPreTraining
from tqdm import tqdm
print(f"\nGenerating locality captions for {len(image_sources)} images...")
processor = AutoProcessor.from_pretrained(model_name)
model = AutoModelForPreTraining.from_pretrained(
model_name, torch_dtype=torch.float16, device_map={"": device},
)
model.eval()
results = {}
for image_id, source in tqdm(image_sources.items(), desc="Locality captions"):
try:
if isinstance(source, str):
image = Image.open(source).convert("RGB")
else:
image = source.convert("RGB")
except Exception:
continue
inputs = processor(
images=image,
text=f"<image>\nUSER: {prompt}\nASSISTANT:",
return_tensors="pt",
).to(device)
with torch.no_grad():
output_ids = model.generate(
**inputs, max_new_tokens=256, do_sample=False,
)
generated = processor.decode(
output_ids[0][inputs["input_ids"].shape[1]:],
skip_special_tokens=True,
).strip()
results[image_id] = generated
del model
torch.cuda.empty_cache()
return results
# ---------------------------------------------------------------------------
# Build the edit set
# ---------------------------------------------------------------------------
def load_caption_targets(caption_targets_path: str,
relation_config: RelationConfig = None) -> tuple[dict, dict]:
"""Load pre-built caption targets from build_caption_targets.py.
Returns:
caption_data: {image_id: {"original": ..., "cleaned": ..., "had_toilet": ..., "is_usable": ...}}
locality_captions: {image_id: original_caption_str}
"""
efficacy_cat = relation_config.efficacy_category if relation_config else "bathroom_no_toilet"
with open(caption_targets_path) as f:
targets = json.load(f)
caption_data = {}
locality_captions = {}
for iid, entry in targets["images"].items():
cat = entry.get("category", "")
original = entry.get("original_caption")
if cat == efficacy_cat and original is not None:
caption_data[iid] = {
"original": original,
"cleaned": entry.get("cleaned_caption"),
"had_toilet": entry.get("had_toilet_mention_llm") or entry.get("had_toilet_mention_regex") or entry.get("had_toilet_mention", False),
"is_hallucinating": entry.get("is_hallucinating", False),
"is_usable": entry.get("is_usable", True),
}
elif original is not None:
# Locality images — store original caption as ground truth
locality_captions[iid] = original
print(f"Loaded caption targets: {len(caption_data)} edit, "
f"{len(locality_captions)} locality")
return caption_data, locality_captions
def build_edit_set(
csv_path: str = None,
image_dir: str = None,
dataset_id: str = HF_DATASET_ID,
max_edit_instances: Optional[int] = None,
max_locality_per_category: Optional[int] = None,
max_eval_per_category: int = DEFAULT_EVAL_PER_CATEGORY,
caption_data: Optional[dict] = None,
locality_captions: Optional[dict] = None,
n_seed_tries: int = 100,
relation_config: RelationConfig = None,
):
"""Build the full edit set dictionary.
Data sources:
- HuggingFace (default): uses the dataset's official train/validation splits.
edit_instances.train = HF train split bathroom_no_toilet (for LoRA etc.)
eval_instances = HF val split, up to max_eval_per_category per category
(for DualEdit: both editing and evaluation use these)
- CSV (legacy): loads all rows then does a local 80/20 split.
Args:
csv_path: (Legacy) Path to bathroom_toilet_labels.csv
image_dir: (Legacy) Image directory.
dataset_id: HuggingFace dataset ID.
max_edit_instances: Cap on edit_instances.train (HF train BNT images).
max_locality_per_category: Cap on locality_instances (HF train non-BNT).
max_eval_per_category: Cap per category for eval_instances (HF val). Default 50.
caption_data: {image_id: {"original","cleaned","had_toilet","is_usable"}}
locality_captions: {image_id: original_caption_str}
"""
# Resolve category names
efficacy_cat = relation_config.efficacy_category if relation_config else "bathroom_no_toilet"
locality_cat_names = list(relation_config.locality_categories) if relation_config else [
"bathroom_with_toilet", "non_bathroom_with_toilet", "unrelated"]
if csv_path is not None and image_dir is not None:
# ---- Legacy CSV path: local 80/20 split (no HF split available) ----
rows = load_csv(csv_path, image_dir)
cats = split_categories(rows, relation_config)
bnt_ids = [r["image_id"] for r in cats[efficacy_cat]]
train_ids, val_ids = _csv_train_val_split(bnt_ids)
bnt_train = [r for r in cats[efficacy_cat] if r["image_id"] in train_ids]
bnt_val = [r for r in cats[efficacy_cat] if r["image_id"] in val_ids]
if max_edit_instances:
bnt_train = bnt_train[:max_edit_instances]
bnt_val = bnt_val[:max_edit_instances]
locality_cats = cats
eval_cats = None # no separate eval set in CSV mode
data_config = {"csv_path": csv_path, "image_dir": image_dir,
"split_seed": SPLIT_SEED, "split_test_size": SPLIT_TEST_SIZE}
else:
# ---- HuggingFace path: use official train/validation splits ----
hf_kwargs = {}
if relation_config is not None:
hf_kwargs = {"scene_col": relation_config.scene_key,
"object_col": relation_config.object_key}
print(f"Loading HuggingFace train split ({dataset_id})...")
train_rows = _hf_rows(dataset_id, split="train", **hf_kwargs)
print(f"Loading HuggingFace validation split ({dataset_id})...")
val_rows = _hf_rows(dataset_id, split="val", **hf_kwargs)
print("\nTrain split categories:")
train_cats = split_categories(train_rows, relation_config)
print("Validation split categories:")
val_cats = split_categories(val_rows, relation_config)
bnt_train = train_cats[efficacy_cat]
if max_edit_instances:
bnt_train = bnt_train[:max_edit_instances]
# Val efficacy: take the first max_eval_per_category images deterministically.
bnt_val_pool = val_cats[efficacy_cat]
n_sample = min(max_eval_per_category, len(bnt_val_pool))
bnt_val = bnt_val_pool[:n_sample]
print(f" Val {efficacy_cat}: first {n_sample} images (deterministic)")
# Locality comes from the train split (no leakage from val)
locality_cats = train_cats
# Eval set: efficacy uses the same sample; other categories take first N from val
eval_cats = {
efficacy_cat: bnt_val,
**{
cat: val_cats[cat][:max_eval_per_category]
for cat in locality_cat_names
}
}
print(f"\nEval set (val split, ≤{max_eval_per_category} per category):")
for cat, rows in eval_cats.items():
print(f" {cat}: {len(rows)}")
data_config = {"dataset_id": dataset_id, "source": "huggingface",
"max_eval_per_category": max_eval_per_category}
print(f"\nEdit instances: {len(bnt_train)} train, {len(bnt_val)} val")
# ---- Build edit instances ----
def make_edit_instance(row, split):
iid = row["image_id"]
inst = {
"image_id": iid,
"image_path": row.get("image_path", iid),
"is_scene": row.get("is_scene", row.get("bathroom", 0)),
"has_object": row.get("has_object", row.get("toilet", 0)),
"split": split,
}
if caption_data and iid in caption_data:
cd = caption_data[iid]
inst["original_caption"] = cd["original"]
inst["target"] = cd["cleaned"]
inst["had_toilet"] = cd["had_toilet"]
inst["is_usable"] = cd["is_usable"]
else:
inst["original_caption"] = None
inst["target"] = None
inst["had_toilet"] = None
inst["is_usable"] = None
return inst
edit_train = [make_edit_instance(r, "train") for r in bnt_train]
edit_val = [make_edit_instance(r, "val") for r in bnt_val]
# ---- Build locality instances (from train split / CSV pool) ----
locality = {}
for cat_name in locality_cat_names:
cat_rows = locality_cats[cat_name]
if max_locality_per_category:
cat_rows = cat_rows[:max_locality_per_category]
locality[cat_name] = []
for row in cat_rows:
iid = row["image_id"]
loc_inst = {
"image_id": iid,
"image_path": row.get("image_path", iid),
"is_scene": row.get("is_scene", row.get("bathroom", 0)),
"has_object": row.get("has_object", row.get("toilet", 0)),
"original_caption": locality_captions.get(iid) if locality_captions else None,
}
locality[cat_name].append(loc_inst)
# ---- Build eval instances (HF val split, all categories) ----
def make_eval_instance(row):
iid = row["image_id"]
inst = {
"image_id": iid,
"image_path": row.get("image_path", iid),
"is_scene": row.get("is_scene", row.get("bathroom", 0)),
"has_object": row.get("has_object", row.get("toilet", 0)),
}
if caption_data and iid in caption_data:
cd = caption_data[iid]
inst["original_caption"] = cd["original"]
inst["target"] = cd["cleaned"]
inst["had_toilet"] = cd["had_toilet"]
inst["is_usable"] = cd["is_usable"]
elif locality_captions and iid in locality_captions:
inst["original_caption"] = locality_captions[iid]
return inst
eval_instances = None
if eval_cats is not None:
eval_instances = {
cat: [make_eval_instance(r) for r in rows]
for cat, rows in eval_cats.items()
}
# ---- Stats ----
all_edit = edit_train + edit_val
n_with_targets = sum(1 for e in all_edit if e["target"] is not None)
n_hallucinated = sum(1 for e in all_edit if e.get("had_toilet"))
n_usable = sum(1 for e in all_edit if e.get("is_usable"))
# Resolve relation-specific values
rc = relation_config
object_keywords = rc.object_keywords if rc else TOILET_KEYWORDS
caption_prompt = CAPTION_PROMPT
train_prompts_list = rc.train_prompts if rc else TRAIN_PROMPTS
generality_prompts_list = rc.generality_prompts if rc else GENERALITY_PROMPTS
relation_key = rc.relation_key if rc else "bathroom_toilet"
edit_set = {
"edit_descriptor": {
"relation": relation_key,
"concept": f"{efficacy_cat}",
"target_tokens": object_keywords,
"edit_type": "caption_suppression",
"edit_prompt": caption_prompt,
"description": (
f"For each {efficacy_cat} image, the model's caption "
f"hallucinating the object is edited to the same caption with "
f"object mentions removed."
),
},
"prompts": {
"edit_prompt": caption_prompt,
"train_prompts": train_prompts_list,
"generality_prompts": generality_prompts_list,
},
"edit_instances": {
"train": edit_train, # HF train BNT — for LoRA / fine-tuning methods
"val": edit_val, # HF val BNT — same as eval_instances BNT
},
"locality_instances": locality,
"stats": {
"relation": relation_key,
"n_edit_train": len(edit_train),
"n_edit_val": len(edit_val),
"n_with_targets": n_with_targets,
"n_hallucinated": n_hallucinated,
"n_usable": n_usable,
**{f"n_locality_{cat}": len(insts) for cat, insts in locality.items()},
},
"data_config": data_config,
}
if eval_instances is not None:
edit_set["eval_instances"] = eval_instances
edit_set["stats"].update({
f"n_eval_{cat}": len(insts)
for cat, insts in eval_instances.items()
})
return edit_set
# ---------------------------------------------------------------------------
# Fill targets into an existing edit_set.json
# ---------------------------------------------------------------------------
def fill_targets(edit_set_path: str, model_name: str, device: str = "cuda"):
"""Generate caption targets and fill them into an existing edit_set.json."""
with open(edit_set_path) as f:
edit_set = json.load(f)
# Collect all efficacy-category instances that need targets (edit + eval sets)
efficacy_cat = edit_set.get("edit_descriptor", {}).get("concept", "bathroom_no_toilet")
all_bnt = (
edit_set["edit_instances"]["train"]
+ edit_set["edit_instances"]["val"]
+ edit_set.get("eval_instances", {}).get(efficacy_cat, [])
)
# Deduplicate by image_id
seen = set()
all_bnt_unique = []
for inst in all_bnt:
if inst["image_id"] not in seen:
seen.add(inst["image_id"])
all_bnt_unique.append(inst)
need_targets = {
inst["image_id"]: inst.get("image_path", inst["image_id"])
for inst in all_bnt_unique
if inst.get("target") is None
}
if not need_targets:
print("All edit instances already have targets.")
return edit_set
# Generate captions
caption_data = generate_captions(
need_targets, model_name=model_name, device=device,
)
def _apply_caption(inst):
iid = inst["image_id"]
if iid in caption_data:
cd = caption_data[iid]
inst["original_caption"] = cd["original"]
inst["target"] = cd["cleaned"]
inst["had_toilet"] = cd["had_toilet"]
inst["is_usable"] = cd["is_usable"]
# Fill targets into edit_instances
for split_name in ["train", "val"]:
for inst in edit_set["edit_instances"][split_name]:
_apply_caption(inst)
# Fill targets into eval_instances efficacy category
for inst in edit_set.get("eval_instances", {}).get(efficacy_cat, []):
_apply_caption(inst)
# Also generate locality captions if missing
locality_need = {}
for cat_name, instances in edit_set["locality_instances"].items():
for inst in instances:
if inst.get("original_caption") is None:
locality_need[inst["image_id"]] = inst["image_path"]
if locality_need:
loc_captions = generate_locality_captions(
locality_need, model_name=model_name, device=device,
)
for cat_name, instances in edit_set["locality_instances"].items():
for inst in instances:
if inst["image_id"] in loc_captions:
inst["original_caption"] = loc_captions[inst["image_id"]]
# Update stats
all_instances = edit_set["edit_instances"]["train"] + edit_set["edit_instances"]["val"]
edit_set["stats"]["n_with_targets"] = sum(
1 for e in all_instances if e.get("target") is not None
)
edit_set["stats"]["n_hallucinated"] = sum(
1 for e in all_instances if e.get("had_toilet")
)
edit_set["stats"]["n_usable"] = sum(
1 for e in all_instances if e.get("is_usable")
)
# Save back
with open(edit_set_path, "w") as f:
json.dump(edit_set, f, indent=2)
print(f"\nUpdated {edit_set_path}")
print(f" {edit_set['stats']}")
return edit_set
# ---------------------------------------------------------------------------
# CLI
# ---------------------------------------------------------------------------
def main():
parser = argparse.ArgumentParser(description="Build KME edit set for captioning")
# Relation selection
parser.add_argument("--relation", type=str, default="bathroom_toilet",
help="Relation key from relations.json (default: bathroom_toilet)")
# Step 1: build structure
parser.add_argument("--csv", type=str, default=None,
help="(Legacy) Path to CSV. If omitted, loads from HuggingFace.")
parser.add_argument("--image_dir", type=str, default=None,
help="(Legacy) Image directory. If omitted, loads from HuggingFace.")
parser.add_argument("--dataset_id", type=str, default=None,
help="HuggingFace dataset ID (default: auto from relation config)")
parser.add_argument("--output", type=str,
default="experiment/knowledge_editing/edit_set.json")
parser.add_argument("--max_edit_instances", type=int, default=None,
help="Cap number of edit_instances.train (HF train BNT images)")
parser.add_argument("--max_locality_per_category", type=int, default=50)
parser.add_argument("--max_eval_per_category", type=int,
default=DEFAULT_EVAL_PER_CATEGORY,
help="Max images per category in eval_instances (HF val split). "
f"Default: {DEFAULT_EVAL_PER_CATEGORY}")
parser.add_argument("--n_seed_tries", type=int, default=100,
help="Try this many random seeds for val BNT sampling and keep "
"the sample with the most hallucinating entries. Default: 1")
# Pre-built caption targets (preferred — from build_caption_targets.py)
parser.add_argument("--caption_targets", type=str, default=None,
help="Path to caption_targets.json from build_caption_targets.py. "
"If provided, skips inline caption generation entirely.")
# Legacy: generate targets inline (needs GPU, prefer --caption_targets)
parser.add_argument("--generate_targets", action="store_true",
help="[Legacy] Generate caption targets using regex cleaning. "
"Prefer --caption_targets for LLM-cleaned captions.")
parser.add_argument("--model", type=str, default="llava-hf/llava-1.5-7b-hf",
help="Model for generating captions (original, pre-edit)")
parser.add_argument("--device", type=str, default="cuda")
# Alternative: fill targets into existing file
parser.add_argument("--fill_targets", type=str, default=None,
help="Path to existing edit_set.json to fill targets into")
args = parser.parse_args()
# Mode: fill targets into existing file
if args.fill_targets:
fill_targets(args.fill_targets, model_name=args.model, device=args.device)
return
# Load relation config
rc = get_relation_config(args.relation)
dataset_id = args.dataset_id or rc.dataset_id
object_re = _build_object_re(rc.object_keywords)
print(f"Relation: {rc}")
print(f"Dataset: {dataset_id}")
# Mode: build from scratch
caption_data = None
locality_captions = None
if args.caption_targets:
# Load from pre-built caption targets (LLM-cleaned)
caption_data, locality_captions = load_caption_targets(args.caption_targets, relation_config=rc)
elif args.generate_targets:
# Legacy: inline generation with regex cleaning
rows = load_csv(args.csv, args.image_dir, args.dataset_id)
cats = split_categories(rows)
# Generate captions for edit images — use PIL image or path
edit_sources = {
r["image_id"]: r.get("image_path") or r.get("image")
for r in cats["bathroom_no_toilet"]
}
caption_data = generate_captions(
edit_sources, model_name=args.model, device=args.device,
)
# Generate captions for locality images
loc_sources = {}
for cat_name in ["bathroom_with_toilet", "non_bathroom_with_toilet", "unrelated"]:
for r in cats[cat_name][:args.max_locality_per_category]:
loc_sources[r["image_id"]] = r.get("image_path") or r.get("image")
locality_captions = generate_locality_captions(
loc_sources, model_name=args.model, device=args.device,
)
edit_set = build_edit_set(
csv_path=args.csv,
image_dir=args.image_dir,
dataset_id=dataset_id,
max_edit_instances=args.max_edit_instances,
max_locality_per_category=args.max_locality_per_category,
max_eval_per_category=args.max_eval_per_category,
caption_data=caption_data,
locality_captions=locality_captions,
n_seed_tries=args.n_seed_tries,
relation_config=rc,
)
out_dir = os.path.dirname(os.path.abspath(args.output))
os.makedirs(out_dir, exist_ok=True)
with open(args.output, "w") as f:
json.dump(edit_set, f, indent=2)
print(f"\nEdit set saved to {args.output}")
print(f" {edit_set['stats']}")
# eval_bnt_ids.json no longer needed — editing and evaluation both use
# the first N images deterministically from the HF val split.
if __name__ == "__main__":
main()