Files

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)