from __future__ import annotations import csv from dataclasses import dataclass from pathlib import Path from typing import Mapping, Sequence import numpy as np import torch from torch import Tensor from .types import MODALITIES, SequenceBatch @dataclass(frozen=True) class FeatureSample: sample_id: str group_id: str duration_s: float sentiment: float polarity: int word_intervals: np.ndarray features: Mapping[str, np.ndarray] times: Mapping[str, np.ndarray] valid: Mapping[str, np.ndarray] @dataclass(frozen=True) class FeatureStats: mean: Mapping[str, np.ndarray] scale: Mapping[str, np.ndarray] def load_feature_samples(feature_dir: Path, manifest_path: Path) -> list[FeatureSample]: """Read the extracted NPZ files and their source-time/label manifest.""" with manifest_path.open("r", encoding="utf-8-sig", newline="") as file: rows = list(csv.DictReader(file)) if not rows: raise ValueError(f"sample manifest is empty: {manifest_path}") samples: list[FeatureSample] = [] seen: set[str] = set() for row in rows: video_id = row["video_id"] clip_id = row["clip_id"] sample_id = f"{video_id}/{clip_id}" if sample_id in seen: raise ValueError(f"duplicate sample in manifest: {sample_id}") seen.add(sample_id) path = feature_dir / f"{video_id}__{clip_id}.npz" if not path.is_file(): raise FileNotFoundError(f"feature file missing for {sample_id}: {path}") with np.load(path, allow_pickle=False) as archive: features = { "text": np.asarray(archive["text_features"], dtype=np.float32), "audio": np.asarray(archive["audio_features"], dtype=np.float32), "vision": np.asarray(archive["vision_features"], dtype=np.float32), } times = { "text": np.asarray(archive["word_intervals_s"], dtype=np.float32).mean(axis=1), "audio": np.asarray(archive["audio_times_s"], dtype=np.float32), "vision": np.asarray(archive["vision_times_s"], dtype=np.float32), } valid = { "text": np.ones(features["text"].shape[0], dtype=np.bool_), "audio": np.asarray(archive["audio_valid"], dtype=np.bool_), "vision": np.asarray(archive["vision_valid"], dtype=np.bool_), } word_intervals = np.asarray(archive["word_intervals_s"], dtype=np.float32) for name in MODALITIES: if features[name].ndim != 2 or times[name].shape != (features[name].shape[0],): raise ValueError(f"invalid {name} feature/timestamp shape in {sample_id}") if valid[name].shape != times[name].shape or not valid[name].any(): raise ValueError(f"{sample_id} has no valid {name} sequence positions") if not np.isfinite(features[name][valid[name]]).all(): raise ValueError(f"non-finite valid {name} values in {sample_id}") if not np.isfinite(times[name][valid[name]]).all(): raise ValueError(f"non-finite valid {name} timestamps in {sample_id}") annotation = row["annotation"].strip().lower() polarity_by_name = {"negative": 0, "neutral": 1, "positive": 2} if annotation not in polarity_by_name: raise ValueError(f"unknown polarity label {annotation!r} in {sample_id}") samples.append( FeatureSample( sample_id=sample_id, group_id=row.get("group_id") or video_id, duration_s=float(row["duration_s"]), sentiment=float(row["label"]), polarity=polarity_by_name[annotation], word_intervals=word_intervals, features=features, times=times, valid=valid, ) ) return samples def fit_feature_stats(samples: Sequence[FeatureSample]) -> FeatureStats: """Fit per-modality z-score parameters using only the training fold.""" if not samples: raise ValueError("cannot fit feature statistics on an empty sample list") means: dict[str, np.ndarray] = {} scales: dict[str, np.ndarray] = {} for name in MODALITIES: values = np.concatenate( [sample.features[name][sample.valid[name]] for sample in samples], axis=0 ).astype(np.float64, copy=False) mean = values.mean(axis=0) scale = values.std(axis=0) scale[scale < 1e-6] = 1.0 means[name] = mean.astype(np.float32) scales[name] = scale.astype(np.float32) return FeatureStats(mean=means, scale=scales) def standardized_features(sample: FeatureSample, stats: FeatureStats) -> dict[str, np.ndarray]: return { name: ((sample.features[name] - stats.mean[name]) / stats.scale[name]).astype( np.float32, copy=False ) for name in MODALITIES } def collate_feature_samples( samples: Sequence[FeatureSample], stats: FeatureStats, device: torch.device, ) -> tuple[dict[str, SequenceBatch], Tensor, list[Tensor]]: """Pad one variable-length batch in memory; padding is masked and never saved.""" if not samples: raise ValueError("cannot collate an empty sample list") batch_size = len(samples) sequences: dict[str, SequenceBatch] = {} for name in MODALITIES: lengths = [sample.features[name].shape[0] for sample in samples] max_length = max(lengths) dimension = samples[0].features[name].shape[1] feature_batch = torch.zeros(batch_size, max_length, dimension, dtype=torch.float32) time_batch = torch.zeros(batch_size, max_length, dtype=torch.float32) valid_batch = torch.zeros(batch_size, max_length, dtype=torch.bool) for index, sample in enumerate(samples): features = standardized_features(sample, stats)[name] length = len(features) feature_batch[index, :length] = torch.from_numpy(features) time_batch[index, :length] = torch.from_numpy(sample.times[name]) valid_batch[index, :length] = torch.from_numpy(sample.valid[name]) sequences[name] = SequenceBatch( features=feature_batch.to(device), times=time_batch.to(device), valid=valid_batch.to(device), ) durations = torch.tensor([sample.duration_s for sample in samples], dtype=torch.float32, device=device) word_intervals = [torch.from_numpy(sample.word_intervals).to(device) for sample in samples] return sequences, durations, word_intervals