160 lines
6.5 KiB
Python
160 lines
6.5 KiB
Python
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
|