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)