""" Datasets for the hallucination-removal experiment. Contains: - FinetuneDataset: Image+prompt pairs for fine-tuning. Data is loaded from the HuggingFace Hub by default. The dataset ID and column names are determined by the active relation (e.g. bathroom_toilet uses columns "bathroom"/"toilet", kitchen_microwave uses "kitchen"/"microwave"). Legacy CSV+image_dir loading is still supported via the csv_path / image_dir constructor arguments. """ import os import csv import random from typing import Optional, Literal, Iterator import torch from torch.utils.data import Dataset from PIL import Image from sklearn.model_selection import train_test_split from experiment.config.train_config import PromptConfig from experiment.data.hf_loader import ( HF_DATASET_ID, DEFAULT_SCENE_COL, DEFAULT_OBJECT_COL, load_hf_dataset, ) SPLIT_SEED = 42 SPLIT_TEST_SIZE = 0.2 _IMAGE_EXTS = {".jpg", ".jpeg", ".png", ".webp", ".bmp", ".gif"} def _iter_image_files(root: str, recursive: bool) -> Iterator[str]: if recursive: for dirpath, _, files in os.walk(root): for f in files: if os.path.splitext(f)[1].lower() in _IMAGE_EXTS: yield os.path.join(dirpath, f) else: with os.scandir(root) as it: for entry in it: if entry.is_file() and os.path.splitext(entry.name)[1].lower() in _IMAGE_EXTS: yield entry.path def _reservoir_sample_image_paths( root: str, k: int, seed: int, recursive: bool, ) -> list[str]: """Sample up to k random image paths in one pass (no full listing in memory).""" rng = random.Random(seed) reservoir: list[str] = [] n = 0 for path in _iter_image_files(root, recursive): n += 1 if len(reservoir) < k: reservoir.append(path) else: j = rng.randint(1, n) if j <= k: reservoir[j - 1] = path return reservoir def finetune_dataset_extra_kwargs(config) -> dict: """Optional kwargs for FinetuneDataset from TrainConfig (CC3M / general mix).""" return { "general_image_dir": config.general_image_dir, "general_dataset_id": getattr(config, "general_dataset_id", None), "num_general_samples": config.num_general_samples, "general_image_seed": config.general_image_seed, "general_image_recursive": config.general_image_recursive, } def get_split_image_ids(csv_path: str = None, split: Literal["train", "val"] = "train", dataset_id: str = HF_DATASET_ID) -> set[str]: """Return the set of image_ids belonging to a train or val split. If csv_path is provided (legacy), uses sklearn train_test_split. Otherwise loads splits directly from the HuggingFace dataset. """ if csv_path is not None: with open(csv_path, "r") as f: reader = csv.DictReader(f) all_ids = [row["image_id"] for row in reader] train_ids, val_ids = train_test_split( all_ids, test_size=SPLIT_TEST_SIZE, random_state=SPLIT_SEED, ) return set(train_ids) if split == "train" else set(val_ids) # HuggingFace dataset — splits are built-in ds = load_hf_dataset(dataset_id, split=split) return set(ds["image_id"]) # --------------------------------------------------------------------------- # Fine-tuning dataset (from step3_data.py) # --------------------------------------------------------------------------- class FinetuneDataset(Dataset): """Dataset for fine-tuning. Each sample is an image paired with a text prompt, processed into model inputs. """ def __init__( self, processor, prompt_config: PromptConfig, dataset_id: str = HF_DATASET_ID, scene_col: str = DEFAULT_SCENE_COL, object_col: str = DEFAULT_OBJECT_COL, csv_path: str = None, image_dir: str = None, max_samples: Optional[int] = None, filter_label: Optional[int] = None, split: Optional[Literal["train", "val"]] = None, upsample_categories: Optional[list[tuple[int, int, int]]] = None, general_image_dir: Optional[str] = None, general_dataset_id: Optional[str] = None, num_general_samples: int = 0, general_image_seed: int = 42, general_image_recursive: bool = False, lm_supervision: bool = False, lm_max_length: int = 640, ): """ Args: processor: HuggingFace processor (tokenizer + image processor). prompt_config: Which prompts to use and how to sample them. dataset_id: HuggingFace dataset ID. scene_col: Column name for scene label (e.g. "bathroom", "kitchen"). object_col: Column name for object label (e.g. "toilet", "microwave"). csv_path: (Legacy) Path to labels CSV. image_dir: (Legacy) Directory containing ``{image_id}.jpg`` files. max_samples: Cap the number of samples. None = all data. filter_label: If set, only keep rows where object == filter_label. split: Deterministic split: "train" (80%), "val" (20%), None = all. upsample_categories: List of (is_scene, has_object, multiplier) tuples. Matching rows are repeated `multiplier` times. Applied after split. general_image_dir: If set and num_general_samples > 0, append that many random images from this local folder. Ignored if general_dataset_id is set. general_dataset_id: If set and num_general_samples > 0, sample that many random images from this HuggingFace dataset (e.g. "username/cc3m-general-2k"). Each row must have an "image" column with PIL images. Labels: is_scene=0, label=0 (unrelated / general). Takes priority over general_image_dir. num_general_samples: How many general images to mix in (0 = disabled). general_image_seed: RNG seed for reproducible sampling. general_image_recursive: If True, walk subfolders for images; else top-level only. Only used with general_image_dir, not general_dataset_id. lm_supervision: If True, tokenise ``USER: … ASSISTANT: `` and return ``labels`` for causal LM cross-entropy (non-caption rows are dropped after load). lm_max_length: Max sequence length when ``lm_supervision`` is True. """ self.processor = processor self.prompt_config = prompt_config self.scene_col = scene_col self.object_col = object_col self._general_hf_ds = None self._general_hf_indices = None self.lm_supervision = lm_supervision self._lm_max_length = lm_max_length self._assistant_marker_ids = processor.tokenizer.encode("ASSISTANT:", add_special_tokens=False) self.data = [] if csv_path is not None and image_dir is not None: # Legacy CSV loading self._load_from_csv(csv_path, image_dir, max_samples, filter_label, split) else: # HuggingFace dataset self._load_from_hf(dataset_id, max_samples, filter_label, split) if num_general_samples and not general_image_dir and not general_dataset_id: print(" WARNING: num_general_samples > 0 but no general image source set; skipping general mix.") n_general = 0 if num_general_samples and general_dataset_id: n_general = self._append_general_from_hf( general_dataset_id, num_general_samples, general_image_seed, ) elif num_general_samples and general_image_dir: n_general = self._append_general_images( general_image_dir, num_general_samples, general_image_seed, general_image_recursive, ) # Per-category counts before upsample cat_counts: dict[str, int] = {} for d in self.data: key = f"scene={d['is_scene']},object={d['label']}" cat_counts[key] = cat_counts.get(key, 0) + 1 # Upsample specified categories if upsample_categories: extra = [] for scene_val, object_val, multiplier in upsample_categories: if multiplier <= 1: continue matching = [d for d in self.data if d["is_scene"] == scene_val and d["label"] == object_val] # Add (multiplier - 1) copies (original already in self.data) for _ in range(multiplier - 1): extra.extend(matching) self.data.extend(extra) if lm_supervision: before = len(self.data) self.data = [d for d in self.data if (d.get("caption") or "").strip()] print(f" lm_supervision: {len(self.data)} samples with caption (dropped {before - len(self.data)} without)") n_pos = sum(d["label"] for d in self.data) print(f"FinetuneDataset: {len(self.data)} samples " f"(has_object={n_pos}, no_object={len(self.data) - n_pos})") if n_general: print(f" general_images: {n_general} (is_scene=0, label=0)") print(f" per-category (before upsample): {cat_counts}") if upsample_categories: cat_counts_after: dict[str, int] = {} for d in self.data: key = f"scene={d['is_scene']},object={d['label']}" cat_counts_after[key] = cat_counts_after.get(key, 0) + 1 print(f" per-category (after upsample): {cat_counts_after}") if len(self.data) == 0: print(f" WARNING: 0 samples loaded!") self._round_robin_idx = 0 def _load_from_csv(self, csv_path, image_dir, max_samples, filter_label, split): """Legacy: load from CSV + image directory.""" print(f" csv_path: {os.path.abspath(csv_path)}") print(f" image_dir: {os.path.abspath(image_dir)}") if split: print(f" split: {split}") split_ids = get_split_image_ids(csv_path, split) if split else None total_rows = 0 missing_images = 0 split_filtered = 0 with open(csv_path, "r") as f: reader = csv.DictReader(f) for row in reader: total_rows += 1 # Legacy CSV uses "toilet" and "bathroom" column names object_val = int(row.get(self.object_col, 0)) if filter_label is not None and object_val != filter_label: continue if split_ids is not None and row["image_id"] not in split_ids: split_filtered += 1 continue image_path = os.path.join(image_dir, f"{row['image_id']}.jpg") if not os.path.exists(image_path): missing_images += 1 continue self.data.append({ "image_path": image_path, "label": object_val, "is_scene": int(row.get(self.scene_col, 0)), "caption": (row.get("caption") or "").strip(), }) if max_samples and len(self.data) >= max_samples: break if split: print(f" split={split}, {split_filtered} rows filtered out") if len(self.data) == 0: print(f" WARNING: 0 samples loaded! " f"CSV had {total_rows} rows, {missing_images} images not found on disk.") def _load_from_hf(self, dataset_id, max_samples, filter_label, split): """Load from HuggingFace dataset.""" print(f" dataset: {dataset_id}") if split: print(f" split: {split}") ds = load_hf_dataset(dataset_id, split=split) else: ds = load_hf_dataset(dataset_id) if hasattr(ds, "keys"): from datasets import concatenate_datasets ds = concatenate_datasets([ds[s] for s in ds]) for item in ds: object_val = int(item[self.object_col]) if filter_label is not None and object_val != filter_label: continue cap = "" if item.get("caption"): cap = str(item["caption"]).strip() self.data.append({ "image": item["image"], "label": object_val, "is_scene": int(item[self.scene_col]), "caption": cap, }) if max_samples and len(self.data) >= max_samples: break def _append_general_images( self, root: str, k: int, seed: int, recursive: bool, ) -> int: root = os.path.expanduser(root) if not os.path.isdir(root): print(f" WARNING: general_image_dir not found or not a directory: {root}") return 0 paths = _reservoir_sample_image_paths(root, k, seed, recursive) if not paths: print(f" WARNING: no image files found under {root}") return 0 if len(paths) < k: print(f" WARNING: only {len(paths)} general images found (requested {k})") for p in paths: self.data.append({ "image_path": p, "label": 0, "is_scene": 0, "caption": "", }) print(f" general_image_dir: {os.path.abspath(root)} (recursive={recursive})") return len(paths) def _append_general_from_hf( self, dataset_id: str, k: int, seed: int, ) -> int: """Sample k random images from a HuggingFace dataset and append as general (is_scene=0, label=0). The dataset must have an 'image' column containing PIL images. Instead of eagerly loading all images, stores the dataset reference and sampled indices; images are loaded on-demand in __getitem__. """ ds = load_hf_dataset(dataset_id, split="train") rng = random.Random(seed) n = len(ds) if n == 0: print(f" WARNING: HF general dataset {dataset_id} has 0 rows") return 0 indices = rng.sample(range(n), min(k, n)) self._general_hf_ds = ds self._general_hf_indices = indices for i in indices: self.data.append({ "general_hf_idx": i, "label": 0, "is_scene": 0, "caption": "", }) print(f" general_dataset_id: {dataset_id} ({min(k, n)}/{n} sampled, lazy load)") return len(indices) def _select_prompt(self, index: int) -> str: prompts = self.prompt_config.prompts if self.prompt_config.sampling == "round_robin": return prompts[index % len(prompts)] else: # uniform return random.choice(prompts) def __len__(self): return len(self.data) def __getitem__(self, idx): item = self.data[idx] prompt = self._select_prompt(idx) if "image_path" in item: image = Image.open(item["image_path"]).convert("RGB") elif "general_hf_idx" in item: image = self._general_hf_ds[item["general_hf_idx"]]["image"].convert("RGB") else: image = item["image"].convert("RGB") if self.lm_supervision: caption = (item.get("caption") or "").strip() text = f"\nUSER: {prompt}\nASSISTANT: {caption}" inputs = self.processor( images=image, text=text, return_tensors="pt", padding="max_length", max_length=self._lm_max_length, truncation=True, ) labels = inputs["input_ids"].clone() row_ids = labels[0].tolist() marker = self._assistant_marker_ids L = len(marker) start = -1 for j in range(len(row_ids) - L + 1): if row_ids[j : j + L] == marker: start = j + L break if start > 0: labels[:, :start] = -100 attn = inputs["attention_mask"] labels[attn == 0] = -100 return { "pixel_values": inputs["pixel_values"][0], "input_ids": inputs["input_ids"][0], "attention_mask": inputs["attention_mask"][0], "labels": labels[0], "has_object": item["label"], "is_scene": item["is_scene"], } inputs = self.processor( images=image, text=f"\n{prompt}", return_tensors="pt", padding="max_length", max_length=640, truncation=True, ) return { "pixel_values": inputs["pixel_values"][0], "input_ids": inputs["input_ids"][0], "attention_mask": inputs["attention_mask"][0], "has_object": item["label"], "is_scene": item["is_scene"], }