整理 Q1-Q3 实验代码与结果
This commit is contained in:
@@ -0,0 +1,84 @@
|
||||
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
|
||||
Reference in New Issue
Block a user