建立分批同步基线(基础文件)
This commit is contained in:
@@ -0,0 +1,159 @@
|
||||
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
|
||||
Reference in New Issue
Block a user