from __future__ import annotations import pickle from pathlib import Path from typing import Any import numpy as np import torch from torch.utils.data import Dataset DEFAULT_FEATURE_FILE = Path(__file__).resolve().parent / "features" / "aligned_50.pkl" def load_aligned50(path: str | Path | None = None, text_variant: str = "hard") -> dict[str, Any]: """Load Q1's 50-bin features, selecting the B1 hard or B2 posterior text view.""" if text_variant not in {"hard", "posterior"}: raise ValueError("text_variant must be 'hard' or 'posterior'") source = Path(path) if path is not None else DEFAULT_FEATURE_FILE with source.open("rb") as stream: payload = pickle.load(stream) if payload.get("metadata", {}).get("schema") != "q1-aligned50-v1": raise ValueError(f"Unsupported aligned feature schema in {source}") arrays = payload["all"].copy() if text_variant == "posterior": arrays["text"] = arrays["text_posterior"] arrays["text_mask"] = arrays["text_posterior_mask"] arrays["text_coverage"] = arrays["text_posterior_activity"] return { **payload, "all": arrays, "metadata": {**payload["metadata"], "selected_text_variant": text_variant}, } def group_holdout_indices( payload: dict[str, Any], validation_fraction: float = 0.2, seed: int = 42 ) -> tuple[np.ndarray, np.ndarray]: """Return train/validation indices while keeping each source video in one split.""" if not 0.0 < validation_fraction < 1.0: raise ValueError("validation_fraction must be between 0 and 1") arrays = payload["all"] if "all" in payload else payload groups = np.asarray(arrays["video_id"], dtype=str) unique_groups = np.unique(groups) if len(unique_groups) < 2: raise ValueError("At least two distinct video_id groups are required") shuffled = np.random.default_rng(seed).permutation(unique_groups) target = max(1, int(np.ceil(len(groups) * validation_fraction))) validation_groups: list[str] = [] for group in shuffled: if validation_groups and np.count_nonzero(np.isin(groups, validation_groups)) >= target: break validation_groups.append(str(group)) if len(validation_groups) == len(unique_groups): validation_groups.pop() validation_mask = np.isin(groups, validation_groups) return np.flatnonzero(~validation_mask), np.flatnonzero(validation_mask) class Q1TorchDataset(Dataset): """PyTorch Dataset yielding feature tensors, masks, labels, times, and video IDs.""" def __init__(self, path: str | Path | None = None, text_variant: str = "hard") -> None: self.payload = load_aligned50(path, text_variant=text_variant) self.arrays = self.payload["all"] self.text_variant = text_variant def __len__(self) -> int: return len(self.arrays["sample_id"]) def __getitem__(self, index: int) -> dict[str, Any]: arrays = self.arrays sample: dict[str, Any] = {} for name in ("text", "audio", "vision", "time_bounds_s", "time_s", "progress"): sample[name] = torch.as_tensor(np.asarray(arrays[name][index], dtype=np.float32)) for name in ("text_mask", "audio_mask", "vision_mask"): sample[name] = torch.as_tensor(np.asarray(arrays[name][index], dtype=bool)) for name in ("text_coverage", "audio_coverage", "vision_coverage"): sample[name] = torch.as_tensor(np.asarray(arrays[name][index], dtype=np.float32)) sample["classification_labels"] = torch.tensor(int(arrays["classification_labels"][index]), dtype=torch.long) sample["regression_labels"] = torch.tensor(float(arrays["regression_labels"][index]), dtype=torch.float32) sample["valid_length"] = torch.tensor(int(arrays["valid_lengths"][index]), dtype=torch.long) for name in ("id", "sample_id", "video_id", "clip_id", "raw_text", "annotations", "source_video_sha256"): sample[name] = str(arrays[name][index]) return sample