214 lines
7.9 KiB
Python
214 lines
7.9 KiB
Python
from __future__ import annotations
|
|
|
|
import pickle
|
|
from dataclasses import dataclass
|
|
from pathlib import Path
|
|
from typing import Any
|
|
|
|
import numpy as np
|
|
|
|
|
|
ROOT = Path(__file__).resolve().parents[3]
|
|
ATTACHMENT2 = ROOT / "E题数据" / "附件2-数据集特征文件"
|
|
MODALITIES = ("text", "audio", "vision")
|
|
|
|
|
|
@dataclass
|
|
class Split:
|
|
x: tuple[np.ndarray, np.ndarray, np.ndarray]
|
|
mask: np.ndarray # N x T x 3
|
|
y_cls: np.ndarray
|
|
y_reg: np.ndarray
|
|
ids: list[str]
|
|
|
|
@property
|
|
def n(self) -> int:
|
|
return len(self.y_cls)
|
|
|
|
@property
|
|
def steps(self) -> int:
|
|
return int(self.x[0].shape[1])
|
|
|
|
|
|
@dataclass
|
|
class RobustStats:
|
|
center: tuple[np.ndarray, np.ndarray, np.ndarray]
|
|
scale: tuple[np.ndarray, np.ndarray, np.ndarray]
|
|
|
|
def save(self, path: Path) -> None:
|
|
path.parent.mkdir(parents=True, exist_ok=True)
|
|
np.savez_compressed(
|
|
path,
|
|
text_center=self.center[0], text_scale=self.scale[0],
|
|
audio_center=self.center[1], audio_scale=self.scale[1],
|
|
vision_center=self.center[2], vision_scale=self.scale[2],
|
|
)
|
|
|
|
@classmethod
|
|
def load(cls, path: Path) -> "RobustStats":
|
|
with np.load(path) as data:
|
|
return cls(
|
|
tuple(data[f"{m}_center"].astype(np.float32) for m in MODALITIES),
|
|
tuple(data[f"{m}_scale"].astype(np.float32) for m in MODALITIES),
|
|
)
|
|
|
|
|
|
def _unpickle(path: Path) -> dict[str, Any]:
|
|
with path.open("rb") as stream:
|
|
return pickle.load(stream, encoding="latin1")
|
|
|
|
|
|
def _ids_and_targets(part: dict[str, Any]) -> tuple[list[str], np.ndarray, np.ndarray]:
|
|
ids = [str(x) for x in part["id"]]
|
|
y_cls = np.asarray(part["classification_labels"], dtype=np.int64).reshape(-1)
|
|
y_reg = np.asarray(part["regression_labels"], dtype=np.float32).reshape(-1)
|
|
return ids, y_cls, y_reg
|
|
|
|
|
|
def _text_mask(part: dict[str, Any]) -> np.ndarray:
|
|
tokens = np.asarray(part["text_bert"])
|
|
if tokens.ndim != 3 or tokens.shape[1] < 2:
|
|
raise ValueError(f"unexpected text_bert shape: {tokens.shape}")
|
|
# MOSEI text_bert rows are input_ids, input_mask, segment_ids.
|
|
return tokens[:, 1, :].astype(bool)
|
|
|
|
|
|
def load_aligned(path: Path | None = None) -> dict[str, Split]:
|
|
path = path or ATTACHMENT2 / "aligned_50.pkl"
|
|
raw = _unpickle(path)
|
|
result: dict[str, Split] = {}
|
|
for name in ("train", "valid"):
|
|
part = raw[name]
|
|
xs = tuple(np.asarray(part[m], dtype=np.float32) for m in MODALITIES)
|
|
masks = [
|
|
_text_mask(part),
|
|
np.any(np.isfinite(xs[1]) & (xs[1] != 0), axis=-1),
|
|
np.any(np.isfinite(xs[2]) & (xs[2] != 0), axis=-1),
|
|
]
|
|
mask = np.stack(masks, axis=-1)
|
|
ids, y_cls, y_reg = _ids_and_targets(part)
|
|
if any(x.shape[1] != 50 for x in xs):
|
|
raise ValueError(f"{name} aligned feature tensors must have 50 slots")
|
|
result[name] = Split(xs, mask, y_cls, y_reg, ids)
|
|
train_videos = {x.split("$_$", 1)[0] for x in result["train"].ids}
|
|
valid_videos = {x.split("$_$", 1)[0] for x in result["valid"].ids}
|
|
overlap = train_videos & valid_videos
|
|
if overlap:
|
|
raise ValueError(f"official train/valid split leaks {len(overlap)} source video ids")
|
|
return result
|
|
|
|
|
|
def _resample_rows_to_50(values: np.ndarray, lengths: list[int] | np.ndarray) -> tuple[np.ndarray, np.ndarray]:
|
|
n, source_steps, dim = values.shape
|
|
output = np.zeros((n, 50, dim), dtype=np.float32)
|
|
mask = np.zeros((n, 50), dtype=bool)
|
|
lengths_arr = np.asarray(lengths, dtype=np.int64).reshape(-1)
|
|
for i in range(n):
|
|
length = int(np.clip(lengths_arr[i], 0, source_steps))
|
|
if length == 0:
|
|
continue
|
|
source = np.nan_to_num(values[i, :length], nan=0.0, posinf=0.0, neginf=0.0)
|
|
observed = np.any(source != 0, axis=-1)
|
|
for j in range(50):
|
|
left = int(np.floor(j * length / 50))
|
|
right = max(left + 1, int(np.ceil((j + 1) * length / 50)))
|
|
right = min(right, length)
|
|
use = observed[left:right]
|
|
if use.any():
|
|
output[i, j] = source[left:right][use].mean(axis=0)
|
|
mask[i, j] = True
|
|
return output, mask
|
|
|
|
|
|
def load_fixed_window(path: Path | None = None) -> dict[str, Split]:
|
|
"""Build a matched 50-slot equal-window control from the unaligned file."""
|
|
path = path or ATTACHMENT2 / "unaligned_50.pkl"
|
|
raw = _unpickle(path)
|
|
result: dict[str, Split] = {}
|
|
for name in ("train", "valid"):
|
|
part = raw[name]
|
|
text = np.asarray(part["text"], dtype=np.float32)
|
|
audio, audio_mask = _resample_rows_to_50(part["audio"], part["audio_lengths"])
|
|
vision, vision_mask = _resample_rows_to_50(part["vision"], part["vision_lengths"])
|
|
text_mask = _text_mask(part)
|
|
xs = (text, audio, vision)
|
|
mask = np.stack((text_mask, audio_mask, vision_mask), axis=-1)
|
|
ids, y_cls, y_reg = _ids_and_targets(part)
|
|
result[name] = Split(xs, mask, y_cls, y_reg, ids)
|
|
return result
|
|
|
|
|
|
def fit_robust_stats(split: Split) -> RobustStats:
|
|
centers: list[np.ndarray] = []
|
|
scales: list[np.ndarray] = []
|
|
for modality in range(3):
|
|
observed = split.mask[:, :, modality].reshape(-1)
|
|
values = split.x[modality].reshape(-1, split.x[modality].shape[-1])[observed]
|
|
if not len(values):
|
|
raise ValueError(f"no observed values for {MODALITIES[modality]}")
|
|
values = np.nan_to_num(values, nan=0.0, posinf=0.0, neginf=0.0)
|
|
center = np.median(values, axis=0)
|
|
mad = np.median(np.abs(values - center), axis=0)
|
|
scale = 1.4826 * mad
|
|
std = np.std(values, axis=0)
|
|
scale = np.where(scale > 1e-6, scale, std)
|
|
scale = np.where(scale > 1e-6, scale, 1.0)
|
|
centers.append(center.astype(np.float32))
|
|
scales.append(scale.astype(np.float32))
|
|
return RobustStats(tuple(centers), tuple(scales))
|
|
|
|
|
|
def apply_robust_stats(split: Split, stats: RobustStats) -> Split:
|
|
xs: list[np.ndarray] = []
|
|
for modality in range(3):
|
|
values = (split.x[modality] - stats.center[modality]) / stats.scale[modality]
|
|
values = np.nan_to_num(values, nan=0.0, posinf=0.0, neginf=0.0)
|
|
values *= split.mask[:, :, modality, None]
|
|
xs.append(values.astype(np.float32, copy=False))
|
|
return Split(tuple(xs), split.mask.copy(), split.y_cls, split.y_reg, split.ids)
|
|
|
|
|
|
def corrupt_masks(
|
|
base: np.ndarray,
|
|
ratio: float,
|
|
modalities: tuple[int, ...],
|
|
seed: int,
|
|
) -> np.ndarray:
|
|
result = base.copy()
|
|
rng = np.random.default_rng(seed)
|
|
n, steps, _ = result.shape
|
|
width = max(1, min(steps, int(round(ratio * steps))))
|
|
starts = rng.integers(0, steps - width + 1, size=n)
|
|
for row, start in enumerate(starts.tolist()):
|
|
result[row, start:start + width, list(modalities)] = False
|
|
return result
|
|
|
|
|
|
def augment_masks(base: np.ndarray, rng: np.random.Generator) -> np.ndarray:
|
|
result = base.copy()
|
|
n, steps, _ = result.shape
|
|
for row in range(n):
|
|
if rng.random() >= 0.85:
|
|
continue
|
|
count = int(rng.integers(1, 4))
|
|
modalities = rng.choice(3, size=count, replace=False)
|
|
ratio = float(rng.choice((0.10, 0.20, 0.30)))
|
|
width = max(1, int(round(ratio * steps)))
|
|
start = int(rng.integers(0, steps - width + 1))
|
|
result[row, start:start + width, modalities] = False
|
|
return result
|
|
|
|
|
|
def shift_audio_vision(split: Split, seed: int, max_shift: int = 10) -> Split:
|
|
rng = np.random.default_rng(seed)
|
|
xs = [x.copy() for x in split.x]
|
|
masks = split.mask.copy()
|
|
for row in range(split.n):
|
|
for modality in (1, 2):
|
|
shift = int(rng.integers(1, max_shift + 1))
|
|
if rng.random() < 0.5:
|
|
shift = -shift
|
|
xs[modality][row] = np.roll(xs[modality][row], shift, axis=0)
|
|
masks[row, :, modality] = np.roll(masks[row, :, modality], shift)
|
|
return Split(tuple(xs), masks, split.y_cls, split.y_reg, split.ids)
|