hallucination / experiment /data /hf_loader.py
ToiTenBao's picture
Upload hallucination folder
a2ffd07 verified
Raw
History Blame Contribute Delete
3.49 kB
"""
Load scene/object hallucination datasets from HuggingFace Hub.
Supports any relation defined in experiment/config/relations.json.
The default relation is bathroom/toilet (dataset: pbcong/bathroom-toilet)
with columns: image (PIL), image_id (str), caption (str), bathroom (ClassLabel),
toilet (ClassLabel).
Other relations use their own column names (e.g. kitchen/microwave).
"""
from __future__ import annotations
from typing import Literal, Optional
from datasets import load_dataset
# Legacy default — kept for backward compatibility when no relation is specified.
HF_DATASET_ID = "pbcong/bathroom-toilet"
# Legacy default column names (bathroom_toilet relation).
DEFAULT_SCENE_COL = "bathroom"
DEFAULT_OBJECT_COL = "toilet"
def load_hf_dataset(
dataset_id: str = HF_DATASET_ID,
split: Optional[str] = None,
):
"""Load the HuggingFace dataset (or a specific split).
Args:
dataset_id: HuggingFace dataset identifier.
split: "train", "val", or None for both splits as a DatasetDict.
Returns:
A ``datasets.Dataset`` (if split given) or ``datasets.DatasetDict``.
"""
if split == "val":
split = "validation"
return load_dataset(dataset_id, split=split)
def get_hf_split_image_ids(
dataset_id: str = HF_DATASET_ID,
split: Literal["train", "val"] = "train",
) -> set[str]:
"""Return the set of image_ids belonging to a split."""
ds = load_hf_dataset(dataset_id, split=split)
return set(ds["image_id"])
def hf_rows(
dataset_id: str = HF_DATASET_ID,
split: Optional[str] = None,
scene_col: str = DEFAULT_SCENE_COL,
object_col: str = DEFAULT_OBJECT_COL,
) -> list[dict]:
"""Load HF dataset and return a list of row dicts.
Each row dict has generic keys compatible with the rest of the codebase:
image_id, is_scene (int), has_object (int), image (PIL.Image), category (str)
Args:
dataset_id: HuggingFace dataset identifier.
split: "train", "val", or None for all splits.
scene_col: Column name for the scene label (e.g. "bathroom", "kitchen").
object_col: Column name for the object label (e.g. "toilet", "microwave").
"""
ds = load_hf_dataset(dataset_id, split=split)
# If DatasetDict, concatenate all splits
if hasattr(ds, "keys"):
from datasets import concatenate_datasets
parts = []
split_map = {}
for s in ds:
for i in range(len(ds[s])):
split_map[ds[s][i]["image_id"]] = "train" if s == "train" else "val"
parts.append(ds[s])
ds = concatenate_datasets(parts)
else:
split_map = None
rows = []
for item in ds:
b = int(item[scene_col])
t = int(item[object_col])
if b == 1 and t == 0:
cat = f"{scene_col}_no_{object_col}"
elif b == 1 and t == 1:
cat = f"{scene_col}_with_{object_col}"
elif b == 0 and t == 1:
cat = f"non_{scene_col}_with_{object_col}"
else:
cat = "unrelated"
row = {
"image_id": item["image_id"],
"is_scene": b,
"has_object": t,
"image": item["image"],
"category": cat,
}
if "caption" in item and item["caption"]:
row["caption"] = item["caption"]
if split_map is not None:
row["split"] = split_map.get(item["image_id"], "train")
rows.append(row)
return rows