tosc-model-weights / eval /eval_tosc.py
julyanghar's picture
Add self-contained TOSC evaluation pipeline
7e01f18
Raw
History Blame Contribute Delete
19.1 kB
#!/usr/bin/env python3
"""Evaluate Triadic Object-State Consistency (TOSC).
The evaluator consumes captions generated for origin / removed / replaced
images, extracts canonical COCO objects with local synonym and lemmatization
rules, and computes TOSC metrics.
"""
from __future__ import annotations
import argparse
import ast
import json
import re
from dataclasses import dataclass
from pathlib import Path
from typing import Any
REPO_ROOT = Path(__file__).resolve().parents[1]
DEFAULT_DATASET_FILE = REPO_ROOT / "benchmark" / "TOSC_dataset.jsonl"
DEFAULT_INSERTIONS_FILE = REPO_ROOT / "benchmark" / "insertions.jsonl"
STATE_SUFFIXES = {
"origin": "_origin",
"removed": "_removed",
"replaced": "_replaced",
}
# Triplets are linked by the leading integer index shared across the three
# states (e.g. "0", "1000"); image ids look like "0_origin",
# "0_masked_clock_removed", "0_masked_clock_replaced".
INDEX_ID_RE = re.compile(r"^(\d+)")
MASKED_OBJECT_RE = re.compile(r"_masked_(.+)$")
try:
from nltk.stem import WordNetLemmatizer
_LEMMATIZER = WordNetLemmatizer()
except Exception:
_LEMMATIZER = None
@dataclass(frozen=True)
class DatasetSample:
sample_id: str
state: str
image_path: str
question: str | None
base_id: str | None
coco_id: str
@dataclass(frozen=True)
class Triplet:
base_id: str
coco_id: str
original_object: str
replacement_object: str
origin: DatasetSample
removed: DatasetSample
replaced: DatasetSample
class ObjectNormalizer:
def __init__(self) -> None:
self.synonym_to_canonical: dict[str, str] = {}
self.phrase_to_canonicals: dict[str, set[str]] = {}
self.canonical_objects: set[str] = set()
self._load_eval_masked_synonyms()
def _load_eval_masked_synonyms(self) -> None:
synonyms_path = Path(__file__).with_name("eval_masked_obj_generative.py")
payload = ast.parse(synonyms_path.read_text(encoding="utf-8"))
object_synonyms: dict[str, list[str]] | None = None
for node in payload.body:
if not isinstance(node, ast.Assign):
continue
if not any(isinstance(target, ast.Name) and target.id == "OBJECT_SYNONYMS" for target in node.targets):
continue
object_synonyms = ast.literal_eval(node.value)
break
if not object_synonyms:
raise ValueError(f"Could not load OBJECT_SYNONYMS from {synonyms_path}")
for canonical, synonyms in object_synonyms.items():
self._add_synonym_group(canonical, [canonical, *synonyms], overwrite=True)
def _add_synonym_group(self, canonical: str, synonyms: list[str], *, overwrite: bool) -> None:
canonical_norm = normalize_phrase(canonical)
if not canonical_norm:
return
self.canonical_objects.add(canonical_norm)
for synonym in synonyms:
self._add_synonym(synonym, canonical_norm, overwrite=overwrite)
def _add_synonym(self, synonym: str, canonical: str, *, overwrite: bool) -> None:
canonical_norm = normalize_phrase(canonical)
if not canonical_norm:
return
self.canonical_objects.add(canonical_norm)
variants = {
normalize_phrase(synonym),
lemmatize_phrase(normalize_phrase(synonym)),
simple_singular_phrase(normalize_phrase(synonym)),
}
for variant in variants:
if not variant:
continue
self.phrase_to_canonicals.setdefault(variant, set()).add(canonical_norm)
if overwrite or variant not in self.synonym_to_canonical:
self.synonym_to_canonical[variant] = canonical_norm
def canonicalize(self, phrase: str) -> str | None:
normalized = normalize_phrase(phrase)
if not normalized:
return None
for candidate in (
normalized,
lemmatize_phrase(normalized),
simple_singular_phrase(normalized),
):
if candidate in self.synonym_to_canonical:
return self.synonym_to_canonical[candidate]
if normalized in self.canonical_objects:
return normalized
return None
def normalize_objects(self, raw_objects: list[str]) -> tuple[list[str], list[str]]:
canonical: set[str] = set()
unmapped: list[str] = []
seen_unmapped: set[str] = set()
for raw_object in raw_objects:
mapped = self.canonicalize(raw_object)
if mapped:
canonical.add(mapped)
continue
cleaned = normalize_phrase(raw_object)
if cleaned and cleaned not in seen_unmapped:
unmapped.append(cleaned)
seen_unmapped.add(cleaned)
return sorted(canonical), unmapped
def extract_caption_objects(self, caption: str) -> dict[str, Any]:
caption_variants = caption_text_variants(caption)
raw_objects: set[str] = set()
normalized_objects: set[str] = set()
for phrase, canonicals in self.phrase_to_canonicals.items():
if not phrase:
continue
if any(contains_token_phrase(text, phrase) for text in caption_variants):
raw_objects.add(phrase)
normalized_objects.update(canonicals)
return {
"raw_objects": sorted(raw_objects),
"normalized_objects": sorted(normalized_objects),
"unmapped_raw_objects": [],
"extraction_source": "local_synonym_match",
}
def normalize_phrase(value: str) -> str:
value = str(value).lower().strip()
value = value.replace("_", " ").replace("-", " ")
value = re.sub(r"[^a-z0-9\s]", " ", value)
value = re.sub(r"\b(a|an|the)\b", " ", value)
return re.sub(r"\s+", " ", value).strip()
def lemmatize_word(word: str) -> str:
if not word:
return word
if _LEMMATIZER is None:
return simple_singular_word(word)
try:
return _LEMMATIZER.lemmatize(word)
except Exception:
return simple_singular_word(word)
def lemmatize_phrase(phrase: str) -> str:
return " ".join(lemmatize_word(part) for part in phrase.split())
def simple_singular_word(word: str) -> str:
if len(word) > 3 and word.endswith("ies"):
return word[:-3] + "y"
if len(word) > 3 and word.endswith("es") and not word.endswith(("ses", "ies")):
return word[:-2]
if len(word) > 3 and word.endswith("s") and not word.endswith("ss"):
return word[:-1]
return word
def simple_singular_phrase(phrase: str) -> str:
return " ".join(simple_singular_word(part) for part in phrase.split())
def caption_text_variants(caption: str) -> set[str]:
normalized = normalize_phrase(caption)
variants = {
normalized,
lemmatize_phrase(normalized),
simple_singular_phrase(normalized),
simple_singular_phrase(lemmatize_phrase(normalized)),
}
return {variant for variant in variants if variant}
def contains_token_phrase(text: str, phrase: str) -> bool:
return f" {phrase} " in f" {text} "
def read_json_or_jsonl(path: Path) -> Any:
if path.suffix == ".jsonl":
with path.open("r", encoding="utf-8") as handle:
return [json.loads(line) for line in handle if line.strip()]
with path.open("r", encoding="utf-8") as handle:
return json.load(handle)
def write_json(path: Path, payload: Any) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
with path.open("w", encoding="utf-8") as handle:
json.dump(payload, handle, ensure_ascii=False, indent=2)
handle.write("\n")
def get_record_id(record: dict[str, Any]) -> str | None:
if "question_id" in record:
return str(record["question_id"])
if "image_id" in record:
return str(record["image_id"])
return None
def get_caption(record: dict[str, Any]) -> str | None:
for key in ("text", "caption", "answer"):
value = record.get(key)
if isinstance(value, str) and value.strip():
return value.strip()
return None
def parse_dataset_sample(record: dict[str, Any]) -> DatasetSample | None:
sample_id = str(record.get("image_id", ""))
image_path = str(record.get("image_path", ""))
if not sample_id or not image_path:
return None
state = ""
stem = sample_id
for candidate_state, suffix in STATE_SUFFIXES.items():
if sample_id.endswith(suffix):
state = candidate_state
stem = sample_id[: -len(suffix)]
break
if not state:
return None
if state == "origin":
coco_id = stem
base_id = None
else:
base_id = stem
match = INDEX_ID_RE.match(base_id)
if not match:
return None
coco_id = match.group(1)
return DatasetSample(
sample_id=sample_id,
state=state,
image_path=image_path,
question=record.get("question"),
base_id=base_id,
coco_id=coco_id,
)
def parse_original_object(base_id: str) -> str:
match = MASKED_OBJECT_RE.search(base_id)
if not match:
raise ValueError(f"Could not parse original object from base id: {base_id}")
return match.group(1).replace("_", " ")
def parse_replacement_from_path(path: str) -> str | None:
stem = Path(path).stem
if "_insert_" not in stem:
return None
return stem.split("_insert_", 1)[1].replace("_", " ")
def load_insertions(path: Path) -> dict[str, dict[str, Any]]:
if not path.exists():
return {}
records: dict[str, dict[str, Any]] = {}
for item in read_json_or_jsonl(path):
image_id = item.get("image_id")
if image_id:
records[str(image_id)] = item
return records
def load_dataset_triplets(dataset_file: Path, insertions_file: Path) -> tuple[list[Triplet], list[dict[str, Any]]]:
rows = read_json_or_jsonl(dataset_file)
origin_by_coco: dict[str, DatasetSample] = {}
removed_by_base: dict[str, DatasetSample] = {}
replaced_by_base: dict[str, DatasetSample] = {}
skipped: list[dict[str, Any]] = []
for row in rows:
sample = parse_dataset_sample(row)
if sample is None:
skipped.append({"reason": "invalid_dataset_sample", "sample": row})
continue
if sample.state == "origin":
origin_by_coco[sample.coco_id] = sample
elif sample.state == "removed" and sample.base_id:
removed_by_base[sample.base_id] = sample
elif sample.state == "replaced" and sample.base_id:
replaced_by_base[sample.base_id] = sample
insertions = load_insertions(insertions_file)
triplets: list[Triplet] = []
for base_id, replaced in sorted(replaced_by_base.items()):
coco_match = INDEX_ID_RE.match(base_id)
if not coco_match:
skipped.append({"reason": "invalid_replaced_base_id", "base_id": base_id})
continue
coco_id = coco_match.group(1)
origin = origin_by_coco.get(coco_id)
removed = removed_by_base.get(base_id)
insertion = insertions.get(base_id, {})
replacement_object = insertion.get("replacement_object") or parse_replacement_from_path(replaced.image_path)
if origin is None or removed is None or not replacement_object:
skipped.append(
{
"reason": "incomplete_triplet",
"base_id": base_id,
"has_origin": origin is not None,
"has_removed": removed is not None,
"has_replacement_object": bool(replacement_object),
}
)
continue
triplets.append(
Triplet(
base_id=base_id,
coco_id=coco_id,
original_object=insertion.get("original_object") or parse_original_object(base_id),
replacement_object=str(replacement_object),
origin=origin,
removed=removed,
replaced=replaced,
)
)
return triplets, skipped
def load_captions(inference_file: Path) -> dict[str, dict[str, Any]]:
captions: dict[str, dict[str, Any]] = {}
for row in read_json_or_jsonl(inference_file):
sample_id = get_record_id(row)
caption = get_caption(row)
if not sample_id or not caption:
continue
captions[sample_id] = {"caption": caption, "record": row}
return captions
def evaluate_tosc(
*,
inference_file: Path,
dataset_file: Path,
insertions_file: Path,
save_file: Path | None,
) -> dict[str, Any]:
normalizer = ObjectNormalizer()
triplets, dataset_skipped = load_dataset_triplets(dataset_file, insertions_file)
captions = load_captions(inference_file)
missing_caption_skipped: list[dict[str, Any]] = []
for triplet in triplets:
for state, sample in (
("origin", triplet.origin),
("removed", triplet.removed),
("replaced", triplet.replaced),
):
if sample.sample_id not in captions:
missing_caption_skipped.append(
{
"reason": "missing_caption",
"base_id": triplet.base_id,
"state": state,
"sample_id": sample.sample_id,
}
)
continue
details: list[dict[str, Any]] = []
skipped: list[dict[str, Any]] = [*dataset_skipped, *missing_caption_skipped]
seen_missing_caption_triplets = {item["base_id"] for item in missing_caption_skipped}
for triplet in triplets:
if triplet.base_id in seen_missing_caption_triplets:
continue
state_outputs: dict[str, dict[str, Any]] = {}
for state, sample in (
("origin", triplet.origin),
("removed", triplet.removed),
("replaced", triplet.replaced),
):
caption = captions[sample.sample_id]["caption"]
extraction = normalizer.extract_caption_objects(caption)
state_outputs[state] = {
"sample_id": sample.sample_id,
"image_path": sample.image_path,
"caption": caption,
**extraction,
}
original_canonical = normalizer.canonicalize(triplet.original_object) or normalize_phrase(triplet.original_object)
replacement_canonical = normalizer.canonicalize(triplet.replacement_object) or normalize_phrase(
triplet.replacement_object
)
origin_objects = set(state_outputs["origin"]["normalized_objects"])
removed_objects = set(state_outputs["removed"]["normalized_objects"])
replaced_objects = set(state_outputs["replaced"]["normalized_objects"])
m_orig_o = int(original_canonical in origin_objects)
m_rem_o = int(original_canonical in removed_objects)
m_rep_o = int(original_canonical in replaced_objects)
m_rep_r = int(replacement_canonical in replaced_objects)
details.append(
{
"base_id": triplet.base_id,
"coco_id": triplet.coco_id,
"original_object": triplet.original_object,
"original_object_canonical": original_canonical,
"replacement_object": triplet.replacement_object,
"replacement_object_canonical": replacement_canonical,
"states": state_outputs,
"indicators": {
"m_orig_o": m_orig_o,
"m_rem_o": m_rem_o,
"m_rep_o": m_rep_o,
"m_rep_r": m_rep_r,
},
"contributions": {
"OPA": m_orig_o,
"RCA": 1 - m_rem_o,
"RUA": m_rep_r * (1 - m_rep_o),
"TOSC": m_orig_o * (1 - m_rem_o) * m_rep_r * (1 - m_rep_o),
"OldPersist_rem": m_rem_o,
"OldPersist_rep": m_rep_o,
"RepFail": 1 - m_rep_r,
"CC": (1 - m_rep_o) * (1 - m_rep_r),
"MixConf": m_rep_o * m_rep_r,
},
}
)
metrics = compute_metrics(details)
metrics.update(
{
"dataset_triplets": len(triplets),
"evaluated_triplets": len(details),
"skipped_triplets": len(skipped),
"object_extraction_source": "local_synonym_match",
}
)
output = {
"metrics": metrics,
"detailed_results": details,
"skipped": skipped,
}
if save_file:
write_json(save_file, output)
summary_path = save_file.with_name(save_file.stem + "_summary.json")
write_json(summary_path, metrics)
return output
def compute_metrics(details: list[dict[str, Any]]) -> dict[str, Any]:
total = len(details)
metric_names = [
"OPA",
"RCA",
"RUA",
"TOSC",
"OldPersist_rem",
"OldPersist_rep",
"RepFail",
"CC",
"MixConf",
]
sums = {name: sum(item["contributions"][name] for item in details) for name in metric_names}
return {
"total_samples": total,
**{f"{name}_count": sums[name] for name in metric_names},
**{name: (sums[name] / total if total else 0.0) for name in metric_names},
**{f"{name}_percent": ((sums[name] / total * 100) if total else 0.0) for name in metric_names},
}
def print_metrics(metrics: dict[str, Any]) -> None:
print("\n" + "=" * 70)
print("Triadic Object-State Consistency Benchmark")
print("=" * 70)
print(f"Dataset Triplets: {metrics.get('dataset_triplets', 0)}")
print(f"Evaluated Triplets: {metrics.get('evaluated_triplets', 0)}")
print(f"Skipped Triplets: {metrics.get('skipped_triplets', 0)}")
print("-" * 70)
for name in ("OPA", "RCA", "RUA", "TOSC", "OldPersist_rem", "OldPersist_rep", "RepFail", "CC", "MixConf"):
print(f"{name:16s}: {metrics.get(name, 0.0) * 100:6.2f}%")
print("=" * 70 + "\n")
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Evaluate Triadic Object-State Consistency.")
parser.add_argument("--inference_file", type=Path, required=True, help="Caption inference JSONL.")
parser.add_argument("--dataset_file", type=Path, default=DEFAULT_DATASET_FILE)
parser.add_argument("--insertions_file", type=Path, default=DEFAULT_INSERTIONS_FILE)
parser.add_argument("--save_file", type=Path, default=None)
return parser.parse_args()
def main() -> int:
args = parse_args()
output = evaluate_tosc(
inference_file=args.inference_file,
dataset_file=args.dataset_file,
insertions_file=args.insertions_file,
save_file=args.save_file,
)
print_metrics(output["metrics"])
return 0
if __name__ == "__main__":
raise SystemExit(main())