hallucination / experiment /config /relation_config.py
ToiTenBao's picture
Upload hallucination folder
a2ffd07 verified
Raw
History Blame Contribute Delete
4.28 kB
"""
Relation Configuration — Defines scene→object hallucination relations.
Each relation describes a (scene, object) pair where VLMs tend to hallucinate
the object in images of the scene that don't actually contain it.
Example: bathroom→toilet — LLaVA hallucinates toilets in bathroom images
that don't contain toilets.
"""
from __future__ import annotations
import json
import os
from dataclasses import dataclass, field
from functools import lru_cache
from typing import Optional
_RELATIONS_JSON = os.path.join(os.path.dirname(__file__), "relations.json")
@dataclass
class RelationConfig:
"""Configuration for a single scene→object hallucination relation."""
relation_key: str # e.g. "bathroom_toilet"
scene_key: str # e.g. "bathroom"
object_key: str # e.g. "toilet"
dataset_id: str # e.g. "pbcong/bathroom-toilet"
# Keywords for token suppression during training
object_keywords: list[str] = field(default_factory=list)
# Keywords for KeywordMentionDetector during evaluation
mention_keywords: list[str] = field(default_factory=list)
# Prompts
train_prompts: list[str] = field(default_factory=lambda: ["Describe this image."])
generality_prompts: list[str] = field(default_factory=list)
# Human-readable object name for judge prompts
judge_object_name: str = ""
# -- Derived properties (category names) ----------------------------------
@property
def scene_no_object(self) -> str:
"""Category: scene present, object absent (efficacy target)."""
return f"{self.scene_key}_no_{self.object_key}"
@property
def scene_with_object(self) -> str:
"""Category: scene present, object present (locality positive)."""
return f"{self.scene_key}_with_{self.object_key}"
@property
def non_scene_with_object(self) -> str:
"""Category: scene absent, object present (locality)."""
return f"non_{self.scene_key}_with_{self.object_key}"
@property
def category_names(self) -> list[str]:
"""All 4 evaluation categories in canonical order."""
return [
self.scene_no_object,
self.scene_with_object,
self.non_scene_with_object,
"unrelated",
]
@property
def efficacy_category(self) -> str:
"""The category where the edit should suppress the object."""
return self.scene_no_object
@property
def locality_categories(self) -> set[str]:
"""Categories where the edit should NOT change outputs."""
return {self.scene_with_object, self.non_scene_with_object, "unrelated"}
def __repr__(self) -> str:
return (
f"RelationConfig({self.relation_key}: "
f"{self.scene_key}{self.object_key}, "
f"dataset={self.dataset_id})"
)
@lru_cache(maxsize=1)
def _load_relations_registry() -> dict[str, RelationConfig]:
"""Load all relation configs from relations.json."""
with open(_RELATIONS_JSON, "r") as f:
raw = json.load(f)
registry = {}
for key, data in raw.items():
registry[key] = RelationConfig(
relation_key=key,
scene_key=data["scene_key"],
object_key=data["object_key"],
dataset_id=data["dataset_id"],
object_keywords=data.get("object_keywords", []),
mention_keywords=data.get("mention_keywords", []),
train_prompts=data.get("train_prompts", ["Describe this image."]),
generality_prompts=data.get("generality_prompts", []),
judge_object_name=data.get("judge_object_name", data["object_key"]),
)
return registry
def get_relation_config(relation_key: str) -> RelationConfig:
"""Look up a RelationConfig by key (e.g. 'bathroom_toilet')."""
registry = _load_relations_registry()
if relation_key not in registry:
available = ", ".join(sorted(registry.keys()))
raise ValueError(
f"Unknown relation key {relation_key!r}. Available: {available}"
)
return registry[relation_key]
def list_relation_keys() -> list[str]:
"""Return all available relation keys."""
return sorted(_load_relations_registry().keys())