85 lines
3.9 KiB
Python
85 lines
3.9 KiB
Python
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
|