Complete standalone final deliverable and unaligned Q2 results
This commit is contained in:
@@ -0,0 +1 @@
|
||||
"""Q2 training, comparison, and inference pipelines."""
|
||||
@@ -0,0 +1 @@
|
||||
"""Maintained Q2 deep-learning schemes."""
|
||||
@@ -0,0 +1 @@
|
||||
"""Q2 multimodal emotion-recognition experiments."""
|
||||
@@ -0,0 +1,214 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import pickle
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
|
||||
from ....data_paths import ATTACHMENT2, PROJECT_ROOT
|
||||
|
||||
|
||||
ROOT = PROJECT_ROOT
|
||||
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)
|
||||
@@ -0,0 +1,614 @@
|
||||
"""Score the frozen EarlyConcat and MoFE checkpoints using the math-Q2 protocol.
|
||||
|
||||
This script performs no training and selects no models. It evaluates the saved
|
||||
three-seed checkpoints on the official labeled test split once, and reuses the
|
||||
fixed 42-scenario validation-mask audit as the controlled-missingness protocol.
|
||||
This source is retained for protocol helpers used by the standalone training runner.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import csv
|
||||
import argparse
|
||||
import hashlib
|
||||
import json
|
||||
import math
|
||||
import statistics
|
||||
import time
|
||||
from collections import defaultdict
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from sklearn.metrics import accuracy_score, f1_score, mean_absolute_error, mean_squared_error
|
||||
from torch import nn
|
||||
|
||||
from .data import (
|
||||
ATTACHMENT2,
|
||||
MODALITIES,
|
||||
RobustStats,
|
||||
Split,
|
||||
_ids_and_targets,
|
||||
_text_mask,
|
||||
_unpickle,
|
||||
apply_robust_stats,
|
||||
fit_robust_stats,
|
||||
load_aligned,
|
||||
)
|
||||
from .models import AlignedFusionModel
|
||||
from .mofe import MixtureOfFusionExperts
|
||||
from .train_mofe import MODEL_CONFIG, _predict, _device_for, EARLYCONCAT, MOFE7_MLP
|
||||
|
||||
|
||||
Q2_ROOT = Path(__file__).resolve().parents[1]
|
||||
REPO_ROOT = Q2_ROOT.parents[1]
|
||||
REFERENCE_DIR = Q2_ROOT / "outputs" / "followups" / "R01_selected_model_reevaluation"
|
||||
OUTPUT_DIR = Q2_ROOT / "outputs" / "followups" / "R02_math_protocol_evaluation"
|
||||
SEEDS = (42, 3407, 2026)
|
||||
BOOTSTRAP_REPS = 1000
|
||||
TEST_BOOTSTRAP_SEED = 20260925
|
||||
AURC_BOOTSTRAP_SEED = 20260926
|
||||
SCENARIO_SEED = 20261833
|
||||
METHODS = (EARLYCONCAT, MOFE7_MLP)
|
||||
CURVE_MODES = ("single", "sync", "partial", "async")
|
||||
CURVE_RATES = (0.0, 0.1, 0.3, 0.5, 0.7)
|
||||
|
||||
|
||||
def write_csv(path: Path, rows: list[dict[str, Any]]) -> None:
|
||||
if not rows:
|
||||
return
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
fields = list(dict.fromkeys(key for row in rows for key in row))
|
||||
with path.open("w", newline="", encoding="utf-8-sig") as stream:
|
||||
writer = csv.DictWriter(stream, fieldnames=fields)
|
||||
writer.writeheader()
|
||||
writer.writerows(rows)
|
||||
|
||||
|
||||
def sha256(path: Path) -> str:
|
||||
digest = hashlib.sha256()
|
||||
with path.open("rb") as stream:
|
||||
for block in iter(lambda: stream.read(1024 * 1024), b""):
|
||||
digest.update(block)
|
||||
return digest.hexdigest()
|
||||
|
||||
|
||||
def split_from_part(part: dict[str, Any]) -> Split:
|
||||
xs = tuple(np.asarray(part[name], dtype=np.float32) for name 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),
|
||||
]
|
||||
ids, y_cls, y_reg = _ids_and_targets(part)
|
||||
return Split(xs, np.stack(masks, axis=-1), y_cls, y_reg, ids)
|
||||
|
||||
|
||||
def load_splits(feature_path: Path) -> dict[str, Split]:
|
||||
raw = _unpickle(feature_path)
|
||||
usual = load_aligned(feature_path)
|
||||
splits = {"train": usual["train"], "valid": usual["valid"], "test": split_from_part(raw["test"])}
|
||||
groups = {
|
||||
name: {sample_id.split("$_$", 1)[0] for sample_id in split.ids}
|
||||
for name, split in splits.items()
|
||||
}
|
||||
for first, second in (("train", "valid"), ("train", "test"), ("valid", "test")):
|
||||
overlap = groups[first] & groups[second]
|
||||
if overlap:
|
||||
raise ValueError(f"official {first}/{second} source-video groups overlap: {len(overlap)}")
|
||||
for name, split in splits.items():
|
||||
expected = np.where(split.y_reg < 0, 0, np.where(split.y_reg == 0, 1, 2))
|
||||
if not np.array_equal(expected, split.y_cls):
|
||||
raise ValueError(f"{name}: classification labels disagree with strict sign of regression labels")
|
||||
return splits
|
||||
|
||||
|
||||
def metrics(split: Split, logits: np.ndarray, intensity: np.ndarray, indices: np.ndarray | None = None) -> dict[str, float]:
|
||||
if indices is None:
|
||||
indices = np.arange(split.n)
|
||||
y_cls = split.y_cls[indices]
|
||||
y_reg = split.y_reg[indices]
|
||||
pred_cls = np.asarray(logits)[indices].argmax(axis=-1)
|
||||
pred_reg = np.clip(np.asarray(intensity).reshape(-1)[indices], -3.0, 3.0)
|
||||
return {
|
||||
"accuracy": float(accuracy_score(y_cls, pred_cls)),
|
||||
"macro_f1": float(f1_score(y_cls, pred_cls, labels=[0, 1, 2], average="macro", zero_division=0)),
|
||||
"mae": float(mean_absolute_error(y_reg, pred_reg)),
|
||||
"rmse": float(math.sqrt(mean_squared_error(y_reg, pred_reg))),
|
||||
"pearson": float(np.corrcoef(y_reg, pred_reg)[0, 1]) if np.std(y_reg) > 0 and np.std(pred_reg) > 0 else float("nan"),
|
||||
}
|
||||
|
||||
|
||||
def load_model(method: str, seed: int, dims: tuple[int, int, int], device: torch.device) -> nn.Module:
|
||||
if method == EARLYCONCAT:
|
||||
checkpoint = REFERENCE_DIR / "models" / "baselines" / "concat" / f"seed_{seed}" / "model_best.pt"
|
||||
model: nn.Module = AlignedFusionModel("concat", dims=dims).to(device)
|
||||
state = torch.load(checkpoint, map_location=device, weights_only=False)
|
||||
if state.get("kind") != "concat" or int(state.get("seed", -1)) != seed:
|
||||
raise ValueError(f"unexpected EarlyConcat checkpoint: {checkpoint}")
|
||||
elif method == MOFE7_MLP:
|
||||
checkpoint = REFERENCE_DIR / "models" / MOFE7_MLP / f"seed_{seed}" / "model_best.pt"
|
||||
model = MixtureOfFusionExperts(dims=dims, **MODEL_CONFIG).to(device)
|
||||
state = torch.load(checkpoint, map_location=device, weights_only=False)
|
||||
if state.get("config") != MODEL_CONFIG or int(state.get("seed", -1)) != seed:
|
||||
raise ValueError(f"unexpected MoFE checkpoint: {checkpoint}")
|
||||
else:
|
||||
raise ValueError(f"unknown model {method}")
|
||||
if tuple(state.get("dims", ())) != dims:
|
||||
raise ValueError(f"feature dimensions do not match checkpoint: {checkpoint}")
|
||||
model.load_state_dict(state["state_dict"])
|
||||
model.eval()
|
||||
return model
|
||||
|
||||
|
||||
def best_interval(visible: np.ndarray, wanted: int, cap: int, location: str, rng: np.random.Generator) -> tuple[int, int] | None:
|
||||
steps = len(visible)
|
||||
candidates: list[tuple[int, int, int, int]] = []
|
||||
for left in range(steps):
|
||||
hits = 0
|
||||
for right in range(left, steps):
|
||||
hits += int(visible[right])
|
||||
count = min(hits, cap)
|
||||
if count:
|
||||
candidates.append((abs(count - wanted), right - left + 1, left, right))
|
||||
if not candidates:
|
||||
return None
|
||||
best = min((error, span) for error, span, _, _ in candidates)
|
||||
tied = [(left, right) for error, span, left, right in candidates if (error, span) == best]
|
||||
if location == "start":
|
||||
return min(tied, key=lambda pair: (pair[0], pair[1]))
|
||||
if location == "end":
|
||||
return max(tied, key=lambda pair: (pair[1], pair[0]))
|
||||
if location == "middle":
|
||||
center = (steps - 1) / 2
|
||||
return min(tied, key=lambda pair: (abs((pair[0] + pair[1]) / 2 - center), pair[0]))
|
||||
if location != "random":
|
||||
raise ValueError(f"unknown interval location: {location}")
|
||||
return tied[int(rng.integers(0, len(tied)))]
|
||||
|
||||
|
||||
def spread_short_spans(visible: np.ndarray, wanted: int, cap: int) -> np.ndarray:
|
||||
positions = np.flatnonzero(visible)
|
||||
count = min(int(wanted), int(cap), len(positions))
|
||||
chosen = np.zeros(len(visible), dtype=bool)
|
||||
if count <= 0:
|
||||
return chosen
|
||||
n_spans = min(3, count)
|
||||
chunks = np.array_split(positions, n_spans)
|
||||
allocations = [count // n_spans + int(i < count % n_spans) for i in range(n_spans)]
|
||||
for chunk, amount in zip(chunks, allocations):
|
||||
if amount <= 0 or len(chunk) == 0:
|
||||
continue
|
||||
amount = min(amount, len(chunk))
|
||||
start = max(0, (len(chunk) - amount) // 2)
|
||||
chosen[chunk[start:start + amount]] = True
|
||||
return chosen
|
||||
|
||||
|
||||
def continuous_mask(
|
||||
original: np.ndarray,
|
||||
rate: float,
|
||||
mode: str,
|
||||
rng: np.random.Generator,
|
||||
*,
|
||||
modalities: tuple[int, ...] | None = None,
|
||||
location: str = "random",
|
||||
span_structure: str = "long",
|
||||
) -> np.ndarray:
|
||||
"""Reproduce math/Q2 continuous masking on this model's observed positions."""
|
||||
observed = np.asarray(original, dtype=bool)
|
||||
result = observed.copy()
|
||||
if rate <= 0 or mode == "none":
|
||||
return result
|
||||
steps, modality_count = observed.shape
|
||||
present = [m for m in range(modality_count) if observed[:, m].any()]
|
||||
if not present:
|
||||
return result
|
||||
if modalities is not None:
|
||||
selected = [int(m) for m in modalities if int(m) in present]
|
||||
if not selected:
|
||||
return result
|
||||
elif mode == "single":
|
||||
selected = [int(rng.choice(present))]
|
||||
elif mode in {"sync", "partial", "async"}:
|
||||
if len(present) == 1:
|
||||
selected = present
|
||||
else:
|
||||
count = int(rng.integers(2, min(3, len(present)) + 1))
|
||||
selected = sorted(int(v) for v in rng.choice(present, size=count, replace=False))
|
||||
else:
|
||||
raise ValueError(f"unknown mask mode: {mode}")
|
||||
|
||||
def max_hide(modality: int) -> int:
|
||||
count = int(observed[:, modality].sum())
|
||||
keep = max(1, int(math.ceil(0.2 * count)))
|
||||
return max(0, count - keep)
|
||||
|
||||
target = {m: min(max_hide(m), int(round(rate * int(observed[:, m].sum())))) for m in selected}
|
||||
if mode == "sync":
|
||||
span = max(1, int(round(rate * steps)))
|
||||
if location == "start":
|
||||
left = 0
|
||||
elif location == "end":
|
||||
left = steps - span
|
||||
elif location == "middle":
|
||||
left = (steps - span) // 2
|
||||
elif location == "random":
|
||||
left = int(rng.integers(0, max(1, steps - span + 1)))
|
||||
else:
|
||||
raise ValueError(f"unknown interval location: {location}")
|
||||
right = min(steps - 1, left + span - 1)
|
||||
for m in selected:
|
||||
candidates = np.flatnonzero(observed[left:right + 1, m]) + left
|
||||
amount = min(len(candidates), max_hide(m), target[m])
|
||||
if amount:
|
||||
offset = 0 if location != "end" else len(candidates) - amount
|
||||
result[candidates[max(0, offset):max(0, offset) + amount], m] = False
|
||||
else:
|
||||
common_span = max(1, int(round(rate * steps)))
|
||||
for rank, m in enumerate(selected):
|
||||
wanted = target[m]
|
||||
if wanted <= 0:
|
||||
continue
|
||||
cap = max_hide(m)
|
||||
if span_structure == "multi_short":
|
||||
hide = spread_short_spans(observed[:, m], wanted, cap)
|
||||
elif span_structure != "long":
|
||||
raise ValueError(f"unknown span structure: {span_structure}")
|
||||
elif mode == "single" and location != "random":
|
||||
# Place a contiguous block at the requested relative location
|
||||
# among observed positions, while keeping the selected-source
|
||||
# missing amount fixed. This avoids treating padding as time.
|
||||
interval = best_interval(observed[:, m], wanted, cap, location, rng)
|
||||
hide = np.zeros(steps, dtype=bool)
|
||||
if interval is not None:
|
||||
left, right = interval
|
||||
candidates = np.flatnonzero(observed[left:right + 1, m]) + left
|
||||
amount = min(len(candidates), wanted, cap)
|
||||
if amount:
|
||||
offset = 0 if location != "end" else len(candidates) - amount
|
||||
hide[candidates[max(0, offset):max(0, offset) + amount]] = True
|
||||
elif mode in {"partial", "async"}:
|
||||
if mode == "partial":
|
||||
base_left = int(rng.integers(0, max(1, steps - common_span + 1))) if location == "random" else (
|
||||
0 if location == "start" else steps - common_span if location == "end" else (steps - common_span) // 2
|
||||
)
|
||||
offset = int(round(rank * common_span * 0.5))
|
||||
else:
|
||||
base_left = 0 if location == "random" else (
|
||||
0 if location == "start" else steps - common_span if location == "end" else (steps - common_span) // 2
|
||||
)
|
||||
available = max(1, steps - common_span + 1)
|
||||
offsets = np.rint(np.linspace(0, max(0, available - 1), len(selected))).astype(int)
|
||||
if location == "random":
|
||||
rng.shuffle(offsets)
|
||||
offset = int(offsets[rank])
|
||||
left = min(max(0, base_left + offset), max(0, steps - common_span))
|
||||
right = min(steps - 1, left + common_span - 1)
|
||||
hide = np.zeros(steps, dtype=bool)
|
||||
candidates = np.flatnonzero(observed[left:right + 1, m]) + left
|
||||
amount = min(len(candidates), wanted, cap)
|
||||
if amount:
|
||||
hide[candidates[:amount]] = True
|
||||
else:
|
||||
interval = best_interval(observed[:, m], wanted, cap, location, rng)
|
||||
hide = np.zeros(steps, dtype=bool)
|
||||
if interval is not None:
|
||||
left, right = interval
|
||||
candidates = np.flatnonzero(observed[left:right + 1, m]) + left
|
||||
amount = min(len(candidates), wanted, cap)
|
||||
if amount:
|
||||
offset = 0 if location != "end" else len(candidates) - amount
|
||||
hide[candidates[max(0, offset):max(0, offset) + amount]] = True
|
||||
result[hide, m] = False
|
||||
return result
|
||||
|
||||
|
||||
def scenario_seed(seed: int, sample_id: str, key: str) -> int:
|
||||
return int.from_bytes(hashlib.sha256(f"{seed}:{sample_id}:{key}".encode()).digest()[:8], "little")
|
||||
|
||||
|
||||
def make_scenarios(valid: Split, seed: int = SCENARIO_SEED) -> dict[str, np.ndarray]:
|
||||
scenarios = {"0.0/none": valid.mask.copy()}
|
||||
for rate in CURVE_RATES[1:]:
|
||||
for mode in CURVE_MODES:
|
||||
key = f"{rate:.1f}/{mode}"
|
||||
scenarios[key] = np.stack([
|
||||
continuous_mask(mask, rate, mode, np.random.default_rng(scenario_seed(seed, sample_id, key)))
|
||||
for sample_id, mask in zip(valid.ids, valid.mask)
|
||||
])
|
||||
modality_sets = (((0,), "T"), ((1,), "A"), ((2,), "V"), ((0, 1), "TA"), ((0, 2), "TV"), ((1, 2), "AV"), ((0, 1, 2), "TAV"))
|
||||
for selected, label in modality_sets:
|
||||
key = f"0.3/modality_{label}"
|
||||
scenarios[key] = np.stack([
|
||||
continuous_mask(mask, 0.3, "sync", np.random.default_rng(scenario_seed(seed, sample_id, key)), modalities=selected)
|
||||
for sample_id, mask in zip(valid.ids, valid.mask)
|
||||
])
|
||||
for modality_index, label in enumerate(("T", "A", "V")):
|
||||
for location in ("start", "middle", "end"):
|
||||
key = f"0.3/location_{location}_{label}"
|
||||
scenarios[key] = np.stack([
|
||||
continuous_mask(mask, 0.3, "single", np.random.default_rng(scenario_seed(seed, sample_id, key)), modalities=(modality_index,), location=location)
|
||||
for sample_id, mask in zip(valid.ids, valid.mask)
|
||||
])
|
||||
for structure in ("long", "multi_short"):
|
||||
key = f"0.3/span_{structure}_{label}"
|
||||
scenarios[key] = np.stack([
|
||||
continuous_mask(mask, 0.3, "single", np.random.default_rng(scenario_seed(seed, sample_id, key)), modalities=(modality_index,), span_structure=structure)
|
||||
for sample_id, mask in zip(valid.ids, valid.mask)
|
||||
])
|
||||
for mode in ("sync", "partial", "async"):
|
||||
key = f"0.3/synchrony_{mode}"
|
||||
scenarios[key] = np.stack([
|
||||
continuous_mask(mask, 0.3, mode, np.random.default_rng(scenario_seed(seed, sample_id, key)), modalities=(0, 1, 2))
|
||||
for sample_id, mask in zip(valid.ids, valid.mask)
|
||||
])
|
||||
return scenarios
|
||||
|
||||
|
||||
def actual_additional_rates(base: np.ndarray, scenarios: dict[str, np.ndarray]) -> dict[str, np.ndarray]:
|
||||
result = {}
|
||||
observed = base.sum(axis=1)
|
||||
for scenario, current in scenarios.items():
|
||||
newly_hidden = base & ~current
|
||||
hidden_count = newly_hidden.sum(axis=1)
|
||||
by_modality = np.divide(
|
||||
hidden_count,
|
||||
observed,
|
||||
out=np.full(hidden_count.shape, np.nan, dtype=np.float64),
|
||||
where=observed > 0,
|
||||
)
|
||||
result[scenario] = np.nanmean(by_modality, axis=1)
|
||||
return result
|
||||
|
||||
|
||||
def aurc_from_curve(rates: list[float], maes: list[float]) -> float:
|
||||
order = np.argsort(np.asarray(rates), kind="stable")
|
||||
x = np.asarray(rates, dtype=np.float64)[order]
|
||||
y = np.asarray(maes, dtype=np.float64)[order]
|
||||
unique_x, inverse = np.unique(x, return_inverse=True)
|
||||
unique_y = np.asarray([y[inverse == i].mean() for i in range(len(unique_x))])
|
||||
if len(unique_x) <= 1 or unique_x[-1] <= 0:
|
||||
return float(maes[0])
|
||||
return float(np.trapezoid(unique_y, unique_x) / unique_x[-1])
|
||||
|
||||
|
||||
def curve_scenarios(mode: str) -> list[str]:
|
||||
return ["0.0/none"] + [f"{rate:.1f}/{mode}" for rate in CURVE_RATES[1:]]
|
||||
|
||||
|
||||
def group_indices(ids: list[str]) -> tuple[list[str], dict[str, np.ndarray]]:
|
||||
groups = sorted({sample_id.split("$_$", 1)[0] for sample_id in ids})
|
||||
mapping = {group: np.flatnonzero(np.asarray([x.split("$_$", 1)[0] == group for x in ids])) for group in groups}
|
||||
return groups, mapping
|
||||
|
||||
|
||||
def bootstrap_clean_test(
|
||||
split: Split,
|
||||
preds: dict[tuple[str, int], dict[str, np.ndarray]],
|
||||
) -> list[dict[str, Any]]:
|
||||
groups, mapping = group_indices(split.ids)
|
||||
rng = np.random.default_rng(TEST_BOOTSTRAP_SEED)
|
||||
draws: dict[str, list[float]] = defaultdict(list)
|
||||
for _ in range(BOOTSTRAP_REPS):
|
||||
chosen = rng.choice(groups, size=len(groups), replace=True)
|
||||
indices = np.concatenate([mapping[group] for group in chosen])
|
||||
per_method = {}
|
||||
for method in METHODS:
|
||||
per_seed = [metrics(split, preds[(method, seed)]["logits"], preds[(method, seed)]["intensity"], indices) for seed in SEEDS]
|
||||
per_method[method] = {key: float(np.mean([row[key] for row in per_seed])) for key in per_seed[0]}
|
||||
for metric in per_method[EARLYCONCAT]:
|
||||
draws[metric].append(per_method[MOFE7_MLP][metric] - per_method[EARLYCONCAT][metric])
|
||||
rows = []
|
||||
for metric, values in draws.items():
|
||||
rows.append({
|
||||
"comparison": "MoFE-7 + MLP Router minus EarlyConcat + BiGRU",
|
||||
"metric": metric,
|
||||
"delta_mean_over_seeds": float(np.mean([r[metric] for r in [
|
||||
metrics(split, preds[(MOFE7_MLP, seed)]["logits"], preds[(MOFE7_MLP, seed)]["intensity"])
|
||||
for seed in SEEDS
|
||||
]]) - np.mean([r[metric] for r in [
|
||||
metrics(split, preds[(EARLYCONCAT, seed)]["logits"], preds[(EARLYCONCAT, seed)]["intensity"])
|
||||
for seed in SEEDS
|
||||
]])),
|
||||
"bootstrap_ci_2p5": float(np.quantile(values, 0.025)),
|
||||
"bootstrap_ci_97p5": float(np.quantile(values, 0.975)),
|
||||
"bootstrap_probability_delta_gt_0": float(np.mean(np.asarray(values) > 0)),
|
||||
"replicates": BOOTSTRAP_REPS,
|
||||
"resampling_unit": "source video id",
|
||||
"paired": True,
|
||||
"seed": TEST_BOOTSTRAP_SEED,
|
||||
})
|
||||
return rows
|
||||
|
||||
|
||||
def bootstrap_aurc(
|
||||
valid: Split,
|
||||
predictions: dict[tuple[str, int, str], dict[str, np.ndarray]],
|
||||
scenarios: dict[str, np.ndarray],
|
||||
rates_by_sample: dict[str, np.ndarray],
|
||||
) -> list[dict[str, Any]]:
|
||||
groups, mapping = group_indices(valid.ids)
|
||||
rng = np.random.default_rng(AURC_BOOTSTRAP_SEED)
|
||||
delta_by_mode: dict[str, list[float]] = {mode: [] for mode in CURVE_MODES}
|
||||
for _ in range(BOOTSTRAP_REPS):
|
||||
chosen = rng.choice(groups, size=len(groups), replace=True)
|
||||
indices = np.concatenate([mapping[group] for group in chosen])
|
||||
for mode in CURVE_MODES:
|
||||
keys = curve_scenarios(mode)
|
||||
model_aucs: dict[str, list[float]] = {method: [] for method in METHODS}
|
||||
for method in METHODS:
|
||||
for seed in SEEDS:
|
||||
xs = [float(np.nanmean(rates_by_sample[key][indices])) for key in keys]
|
||||
ys = [float(np.abs(valid.y_reg[indices] - predictions[(method, seed, key)]["intensity"][indices]).mean()) for key in keys]
|
||||
model_aucs[method].append(aurc_from_curve(xs, ys))
|
||||
delta_by_mode[mode].append(float(np.mean(model_aucs[MOFE7_MLP]) - np.mean(model_aucs[EARLYCONCAT])))
|
||||
point = {}
|
||||
for mode in CURVE_MODES:
|
||||
model_aucs = {}
|
||||
for method in METHODS:
|
||||
model_aucs[method] = []
|
||||
for seed in SEEDS:
|
||||
keys = curve_scenarios(mode)
|
||||
xs = [float(np.nanmean(rates_by_sample[key])) for key in keys]
|
||||
ys = [float(np.abs(valid.y_reg - predictions[(method, seed, key)]["intensity"]).mean()) for key in keys]
|
||||
model_aucs[method].append(aurc_from_curve(xs, ys))
|
||||
point[mode] = float(np.mean(model_aucs[MOFE7_MLP]) - np.mean(model_aucs[EARLYCONCAT]))
|
||||
rows = []
|
||||
for mode, values in delta_by_mode.items():
|
||||
rows.append({
|
||||
"mode": mode,
|
||||
"delta_aurc_mae_mofe_minus_earlyconcat": point[mode],
|
||||
"bootstrap_ci_2p5": float(np.quantile(values, 0.025)),
|
||||
"bootstrap_ci_97p5": float(np.quantile(values, 0.975)),
|
||||
"bootstrap_probability_delta_lt_0": float(np.mean(np.asarray(values) < 0)),
|
||||
"replicates": BOOTSTRAP_REPS,
|
||||
"resampling_unit": "source video id",
|
||||
"paired": True,
|
||||
"seed": AURC_BOOTSTRAP_SEED,
|
||||
})
|
||||
return rows
|
||||
|
||||
|
||||
def run(device_name: str = "auto", batch_size: int = 64, masks_only: bool = False) -> None:
|
||||
OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
|
||||
device = _device_for(device_name)
|
||||
torch.set_num_threads(4)
|
||||
torch.backends.cudnn.deterministic = True
|
||||
torch.backends.cudnn.benchmark = False
|
||||
|
||||
feature_path = ATTACHMENT2 / "aligned_50.pkl"
|
||||
with (REFERENCE_DIR / "run_manifest.json").open("r", encoding="utf-8") as stream:
|
||||
reference_manifest = json.load(stream)
|
||||
if sha256(feature_path) != reference_manifest["feature_sha256"]:
|
||||
raise ValueError("current official feature file hash differs from the checkpoint evaluation manifest")
|
||||
|
||||
raw_splits = load_splits(feature_path)
|
||||
train = raw_splits["train"]
|
||||
valid = raw_splits["valid"]
|
||||
test = raw_splits["test"]
|
||||
scaler_path = REFERENCE_DIR / "aligned_robust_stats.npz"
|
||||
stats = RobustStats.load(scaler_path)
|
||||
computed = fit_robust_stats(train)
|
||||
scaler_diff = max(
|
||||
max(float(np.max(np.abs(a - b))) for a, b in zip(computed.center, stats.center)),
|
||||
max(float(np.max(np.abs(a - b))) for a, b in zip(computed.scale, stats.scale)),
|
||||
)
|
||||
if scaler_diff > 1e-6:
|
||||
raise ValueError(f"checkpoint scaler is not the train-only scaler (max difference {scaler_diff})")
|
||||
valid = apply_robust_stats(valid, stats)
|
||||
test = apply_robust_stats(test, stats)
|
||||
dims = tuple(x.shape[-1] for x in train.x)
|
||||
|
||||
if not masks_only:
|
||||
# Final, clean official-test evaluation; no retraining or selection occurs here.
|
||||
test_predictions: dict[tuple[str, int], dict[str, np.ndarray]] = {}
|
||||
test_rows: list[dict[str, Any]] = []
|
||||
for method in METHODS:
|
||||
for seed in SEEDS:
|
||||
model = load_model(method, seed, dims, device)
|
||||
prediction = _predict(model, test, test.mask, device, batch_size)
|
||||
test_predictions[(method, seed)] = prediction
|
||||
test_rows.append({"method": method, "seed": seed, "n_test": test.n, **metrics(test, prediction["logits"], prediction["intensity"])})
|
||||
del model
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
summary_rows = []
|
||||
for method in METHODS:
|
||||
subset = [row for row in test_rows if row["method"] == method]
|
||||
for metric in ("accuracy", "macro_f1", "mae", "rmse", "pearson"):
|
||||
values = [float(row[metric]) for row in subset]
|
||||
summary_rows.append({"method": method, "metric": metric, "mean": float(np.mean(values)), "sd_across_seeds": float(np.std(values, ddof=1))})
|
||||
write_csv(OUTPUT_DIR / "official_test_metrics_by_seed.csv", test_rows)
|
||||
write_csv(OUTPUT_DIR / "official_test_summary.csv", summary_rows)
|
||||
write_csv(OUTPUT_DIR / "official_test_paired_bootstrap.csv", bootstrap_clean_test(test, test_predictions))
|
||||
|
||||
# Reproduce the math-Q2 42-scenario design with a per-sample stable seed,
|
||||
# while applying it to the observation masks used to train these models.
|
||||
scenario_masks = make_scenarios(valid)
|
||||
rates_by_sample = actual_additional_rates(valid.mask, scenario_masks)
|
||||
condition_predictions: dict[tuple[str, int, str], dict[str, np.ndarray]] = {}
|
||||
condition_rows: list[dict[str, Any]] = []
|
||||
for method in METHODS:
|
||||
for seed in SEEDS:
|
||||
model = load_model(method, seed, dims, device)
|
||||
for scenario, masks in scenario_masks.items():
|
||||
prediction = _predict(model, valid, masks, device, batch_size)
|
||||
condition_predictions[(method, seed, scenario)] = prediction
|
||||
values = metrics(valid, prediction["logits"], prediction["intensity"])
|
||||
condition_rows.append({
|
||||
"method": method,
|
||||
"seed": seed,
|
||||
"scenario": scenario,
|
||||
"realized_additional_global_rate": float(np.nanmean(rates_by_sample[scenario])),
|
||||
"n_valid": valid.n,
|
||||
**values,
|
||||
})
|
||||
del model
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
write_csv(OUTPUT_DIR / "controlled_metrics_by_scenario.csv", condition_rows)
|
||||
|
||||
auc_rows: list[dict[str, Any]] = []
|
||||
for method in METHODS:
|
||||
for seed in SEEDS:
|
||||
for mode in CURVE_MODES:
|
||||
keys = curve_scenarios(mode)
|
||||
xs = [float(np.nanmean(rates_by_sample[key])) for key in keys]
|
||||
ys = [float(np.abs(valid.y_reg - condition_predictions[(method, seed, key)]["intensity"]).mean()) for key in keys]
|
||||
auc_rows.append({"method": method, "seed": seed, "mask_mode": mode, "aurc_mae": aurc_from_curve(xs, ys), "rates_realized": json.dumps(xs)})
|
||||
write_csv(OUTPUT_DIR / "aurc_mae_by_mode_seed.csv", auc_rows)
|
||||
auc_summary = []
|
||||
for method in METHODS:
|
||||
for mode in CURVE_MODES:
|
||||
values = [row["aurc_mae"] for row in auc_rows if row["method"] == method and row["mask_mode"] == mode]
|
||||
auc_summary.append({"method": method, "mask_mode": mode, "mean": float(np.mean(values)), "sd_across_seeds": float(np.std(values, ddof=1))})
|
||||
write_csv(OUTPUT_DIR / "aurc_mae_summary.csv", auc_summary)
|
||||
write_csv(OUTPUT_DIR / "aurc_mae_paired_bootstrap.csv", bootstrap_aurc(valid, condition_predictions, scenario_masks, rates_by_sample))
|
||||
|
||||
manifest = {
|
||||
"experiment": "Frozen EarlyConcat vs MoFE-7 evaluation under math/Q2 test protocol",
|
||||
"created_unix": time.time(),
|
||||
"device": str(device),
|
||||
"cuda_device": torch.cuda.get_device_name(0) if device.type == "cuda" else None,
|
||||
"feature_file": str(feature_path),
|
||||
"feature_sha256": sha256(feature_path),
|
||||
"representation": "official aligned_50 ordered positions; not physical-time bins",
|
||||
"train_valid_test_counts": {name: split.n for name, split in raw_splits.items()},
|
||||
"source_video_groups": {name: len({sample_id.split("$_$", 1)[0] for sample_id in split.ids}) for name, split in raw_splits.items()},
|
||||
"official_group_splits_disjoint": True,
|
||||
"test_evaluation": (
|
||||
"one final clean evaluation on official labeled test split; no training/model selection/calibration"
|
||||
if not masks_only else "test outputs preserved from the earlier single evaluation; no test prediction was rerun"
|
||||
),
|
||||
"test_prediction_performed_this_invocation": not masks_only,
|
||||
"seeds": list(SEEDS),
|
||||
"checkpoint_source": str(REFERENCE_DIR / "models"),
|
||||
"train_only_scaler": str(scaler_path),
|
||||
"scaler_max_abs_difference_from_train_refit": scaler_diff,
|
||||
"test_labels_used_for_training_or_selection": False,
|
||||
"controlled_missingness": {
|
||||
"scenario_seed": SCENARIO_SEED,
|
||||
"scenario_design": "math/Q2 42-scenario design regenerated on the Q2 models' BERT attention-mask base",
|
||||
"scenarios": len(scenario_masks),
|
||||
"AURC": "normalized trapezoidal area of MAE over realized equal-modality-weighted added missing rate, at 0/.1/.3/.5/.7 for single/sync/partial/async",
|
||||
},
|
||||
"bootstrap": {
|
||||
"replicates": BOOTSTRAP_REPS,
|
||||
"test_seed": TEST_BOOTSTRAP_SEED,
|
||||
"aurc_seed": AURC_BOOTSTRAP_SEED,
|
||||
"unit": "source video id",
|
||||
"paired": True,
|
||||
},
|
||||
}
|
||||
(OUTPUT_DIR / "run_manifest.json").write_text(json.dumps(manifest, indent=2), encoding="utf-8")
|
||||
print(f"wrote math-protocol comparison to {OUTPUT_DIR}")
|
||||
print(f"n_test={test.n}; n_valid={valid.n}; device={device}; scenarios={len(scenario_masks)}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--masks-only", action="store_true", help="Recompute validation mask scenarios without rerunning official-test inference")
|
||||
args = parser.parse_args()
|
||||
run(masks_only=args.masks_only)
|
||||
@@ -0,0 +1,4 @@
|
||||
"""Compatibility import for the maintained model registry."""
|
||||
from ....model.early_concat import AlignedFusionModel
|
||||
|
||||
__all__ = ["AlignedFusionModel"]
|
||||
@@ -0,0 +1,4 @@
|
||||
"""Compatibility import for the maintained model registry."""
|
||||
from ....model.mofe import EXPERT_NAMES, SUBSETS, MixtureOfFusionExperts
|
||||
|
||||
__all__ = ["EXPERT_NAMES", "SUBSETS", "MixtureOfFusionExperts"]
|
||||
@@ -0,0 +1,490 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import csv
|
||||
import hashlib
|
||||
import json
|
||||
import math
|
||||
import random
|
||||
import shutil
|
||||
import time
|
||||
from collections import Counter
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import matplotlib.pyplot as plt
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from sklearn.metrics import accuracy_score, f1_score, mean_absolute_error
|
||||
from torch import nn
|
||||
|
||||
from .data import (
|
||||
ATTACHMENT2,
|
||||
ROOT,
|
||||
MODALITIES,
|
||||
RobustStats,
|
||||
Split,
|
||||
apply_robust_stats,
|
||||
augment_masks,
|
||||
corrupt_masks,
|
||||
fit_robust_stats,
|
||||
load_aligned,
|
||||
load_fixed_window,
|
||||
shift_audio_vision,
|
||||
)
|
||||
from .models import AlignedFusionModel
|
||||
|
||||
|
||||
PATTERNS = {
|
||||
"text": (0,),
|
||||
"audio": (1,),
|
||||
"vision": (2,),
|
||||
"audio_vision": (1, 2),
|
||||
"all_modalities": (0, 1, 2),
|
||||
}
|
||||
KINDS = ("concat",)
|
||||
|
||||
|
||||
def seed_everything(seed: int) -> None:
|
||||
random.seed(seed)
|
||||
np.random.seed(seed)
|
||||
torch.manual_seed(seed)
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.manual_seed_all(seed)
|
||||
torch.backends.cudnn.deterministic = True
|
||||
torch.backends.cudnn.benchmark = False
|
||||
|
||||
|
||||
def _tensor_split(split: Split, device: torch.device) -> tuple[tuple[torch.Tensor, ...], torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
xs = tuple(torch.as_tensor(x, dtype=torch.float32, device=device) for x in split.x)
|
||||
mask = torch.as_tensor(split.mask, dtype=torch.bool, device=device)
|
||||
y_cls = torch.as_tensor(split.y_cls, dtype=torch.long, device=device)
|
||||
y_reg = torch.as_tensor(split.y_reg, dtype=torch.float32, device=device)
|
||||
return xs, mask, y_cls, y_reg
|
||||
|
||||
|
||||
def _loss(output: dict[str, torch.Tensor], y_cls: torch.Tensor, y_reg: torch.Tensor) -> torch.Tensor:
|
||||
class_loss = F.cross_entropy(output["logits"], y_cls)
|
||||
intensity_loss = F.smooth_l1_loss(output["intensity"] / 3.0, y_reg / 3.0)
|
||||
return class_loss + 0.5 * intensity_loss
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def _score_arrays(
|
||||
model: AlignedFusionModel,
|
||||
split: Split,
|
||||
mask: np.ndarray,
|
||||
device: torch.device,
|
||||
batch_size: int = 128,
|
||||
) -> tuple[dict[str, float], dict[str, np.ndarray]]:
|
||||
model.eval()
|
||||
predictions: dict[str, list[np.ndarray]] = {"logits": [], "intensity": []}
|
||||
xs = split.x
|
||||
for start in range(0, split.n, batch_size):
|
||||
end = min(start + batch_size, split.n)
|
||||
xb = tuple(torch.as_tensor(x[start:end], dtype=torch.float32, device=device) for x in xs)
|
||||
mb = torch.as_tensor(mask[start:end], dtype=torch.bool, device=device)
|
||||
output = model(xb, mb)
|
||||
predictions["logits"].append(output["logits"].float().cpu().numpy())
|
||||
predictions["intensity"].append(output["intensity"].float().cpu().numpy())
|
||||
logits = np.concatenate(predictions["logits"], axis=0)
|
||||
intensity = np.clip(np.concatenate(predictions["intensity"], axis=0), -3.0, 3.0)
|
||||
pred_cls = logits.argmax(axis=-1)
|
||||
pearson = _pearson(split.y_reg, intensity)
|
||||
metrics = {
|
||||
"accuracy": float(accuracy_score(split.y_cls, pred_cls)),
|
||||
"macro_f1": float(f1_score(split.y_cls, pred_cls, labels=[0, 1, 2], average="macro", zero_division=0)),
|
||||
"mae": float(mean_absolute_error(split.y_reg, intensity)),
|
||||
"pearson": pearson,
|
||||
}
|
||||
return metrics, {"logits": logits, "intensity": intensity, "class": pred_cls}
|
||||
|
||||
|
||||
def _pearson(y: np.ndarray, pred: np.ndarray) -> float:
|
||||
a = np.asarray(y, dtype=np.float64)
|
||||
b = np.asarray(pred, dtype=np.float64)
|
||||
if a.std() < 1e-12 or b.std() < 1e-12:
|
||||
return 0.0
|
||||
return float(np.corrcoef(a, b)[0, 1])
|
||||
|
||||
|
||||
def _validation_loss(model: AlignedFusionModel, valid: Split, device: torch.device, batch_size: int) -> float:
|
||||
model.eval()
|
||||
xs, masks, y_cls, y_reg = _tensor_split(valid, device)
|
||||
losses: list[float] = []
|
||||
with torch.inference_mode():
|
||||
for start in range(0, valid.n, batch_size):
|
||||
idx = slice(start, min(start + batch_size, valid.n))
|
||||
output = model(tuple(x[idx] for x in xs), masks[idx])
|
||||
losses.append(float(_loss(output, y_cls[idx], y_reg[idx]).item()))
|
||||
return float(np.average(losses, weights=[min(batch_size, valid.n - i) for i in range(0, valid.n, batch_size)]))
|
||||
|
||||
|
||||
def _write_csv(path: Path, rows: list[dict[str, Any]]) -> None:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
if not rows:
|
||||
return
|
||||
fields = list(dict.fromkeys(key for row in rows for key in row))
|
||||
with path.open("w", newline="", encoding="utf-8-sig") as stream:
|
||||
writer = csv.DictWriter(stream, fieldnames=fields)
|
||||
writer.writeheader()
|
||||
writer.writerows(rows)
|
||||
|
||||
|
||||
def _train_one(
|
||||
kind: str,
|
||||
train: Split,
|
||||
valid: Split,
|
||||
output_dir: Path,
|
||||
device: torch.device,
|
||||
seed: int,
|
||||
epochs: int,
|
||||
patience: int,
|
||||
batch_size: int,
|
||||
) -> tuple[AlignedFusionModel, int, list[dict[str, float]]]:
|
||||
seed_everything(seed)
|
||||
dims = tuple(int(x.shape[-1]) for x in train.x)
|
||||
model = AlignedFusionModel(kind, dims=dims).to(device)
|
||||
optimizer = torch.optim.AdamW(model.parameters(), lr=1.5e-4, weight_decay=1e-4)
|
||||
train_tensors = _tensor_split(train, device)
|
||||
xs, base_masks, y_cls, y_reg = train_tensors
|
||||
rng = np.random.default_rng(seed + 809)
|
||||
best_loss = math.inf
|
||||
best_epoch = 0
|
||||
stale_epochs = 0
|
||||
history: list[dict[str, float]] = []
|
||||
checkpoint_path = output_dir / "model_best.pt"
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
for epoch in range(1, epochs + 1):
|
||||
model.train()
|
||||
order = rng.permutation(train.n)
|
||||
batch_losses: list[float] = []
|
||||
for start in range(0, train.n, batch_size):
|
||||
ids_np = order[start:start + batch_size]
|
||||
ids = torch.as_tensor(ids_np, dtype=torch.long, device=device)
|
||||
masks_np = augment_masks(train.mask[ids_np], rng)
|
||||
masks = torch.as_tensor(masks_np, dtype=torch.bool, device=device)
|
||||
output = model(tuple(x.index_select(0, ids) for x in xs), masks)
|
||||
loss = _loss(output, y_cls.index_select(0, ids), y_reg.index_select(0, ids))
|
||||
optimizer.zero_grad(set_to_none=True)
|
||||
loss.backward()
|
||||
nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
|
||||
optimizer.step()
|
||||
batch_losses.append(float(loss.detach().item()))
|
||||
valid_loss = _validation_loss(model, valid, device, batch_size)
|
||||
row = {"epoch": float(epoch), "train_loss": float(np.mean(batch_losses)), "valid_clean_loss": valid_loss}
|
||||
history.append(row)
|
||||
print(f"[{kind}] epoch={epoch:02d} train={row['train_loss']:.4f} valid={valid_loss:.4f}", flush=True)
|
||||
if valid_loss < best_loss - 1e-4:
|
||||
best_loss = valid_loss
|
||||
best_epoch = epoch
|
||||
stale_epochs = 0
|
||||
torch.save({"kind": kind, "dims": dims, "state_dict": model.state_dict(), "seed": seed, "best_epoch": epoch}, checkpoint_path)
|
||||
else:
|
||||
stale_epochs += 1
|
||||
if stale_epochs >= patience:
|
||||
break
|
||||
|
||||
saved = torch.load(checkpoint_path, map_location=device, weights_only=False)
|
||||
model.load_state_dict(saved["state_dict"])
|
||||
model.eval()
|
||||
_write_csv(output_dir / "training_history.csv", history)
|
||||
return model, best_epoch, history
|
||||
|
||||
|
||||
def _conditions(valid: Split, seed: int) -> list[tuple[str, float, np.ndarray]]:
|
||||
result = [("clean", 0.0, valid.mask.copy())]
|
||||
for rate in (0.10, 0.20, 0.30):
|
||||
for pattern_id, (pattern, mods) in enumerate(PATTERNS.items()):
|
||||
result.append((pattern, rate, corrupt_masks(valid.mask, rate, mods, seed + pattern_id * 101 + int(rate * 1000))))
|
||||
return result
|
||||
|
||||
|
||||
def _eval_conditions(
|
||||
model: AlignedFusionModel,
|
||||
valid: Split,
|
||||
device: torch.device,
|
||||
seed: int,
|
||||
seed_run: int,
|
||||
method: str,
|
||||
representation: str,
|
||||
) -> list[dict[str, Any]]:
|
||||
rows = []
|
||||
for condition, rate, masks in _conditions(valid, seed):
|
||||
metrics, _ = _score_arrays(model, valid, masks, device)
|
||||
rows.append({"method": method, "representation": representation, "seed": seed_run, "condition": condition,
|
||||
"missing_rate": rate, "n_valid": valid.n, **metrics})
|
||||
print(f"[{method}/{representation}] {condition:14s} rate={rate:.1f} "
|
||||
f"F1={metrics['macro_f1']:.3f} MAE={metrics['mae']:.3f} "
|
||||
f"P={metrics['pearson']:.3f}", flush=True)
|
||||
return rows
|
||||
|
||||
|
||||
def _summary(rows: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
||||
groups = list(dict.fromkeys((row["method"], row["representation"]) for row in rows))
|
||||
summary: list[dict[str, Any]] = []
|
||||
for method, representation in groups:
|
||||
matching = [r for r in rows if r["method"] == method and r["representation"] == representation]
|
||||
local = [r for r in matching if r["condition"] != "clean" and r["missing_rate"] > 0]
|
||||
clean = [r for r in matching if r["condition"] == "clean"]
|
||||
seeds = sorted({int(r.get("seed", 0)) for r in matching})
|
||||
|
||||
def per_seed_mean(selected: list[dict[str, Any]], metric: str) -> list[float]:
|
||||
return [float(np.mean([r[metric] for r in selected if int(r.get("seed", 0)) == seed]))
|
||||
for seed in seeds if any(int(r.get("seed", 0)) == seed for r in selected)]
|
||||
|
||||
clean_f1 = per_seed_mean(clean, "macro_f1")
|
||||
clean_accuracy = per_seed_mean(clean, "accuracy")
|
||||
clean_mae = per_seed_mean(clean, "mae")
|
||||
clean_pearson = per_seed_mean(clean, "pearson")
|
||||
corrupt_f1 = per_seed_mean(local, "macro_f1")
|
||||
corrupt_accuracy = per_seed_mean(local, "accuracy")
|
||||
corrupt_mae = per_seed_mean(local, "mae")
|
||||
corrupt_pearson = per_seed_mean(local, "pearson")
|
||||
row: dict[str, Any] = {
|
||||
"method": method,
|
||||
"representation": representation,
|
||||
"n_seeds": len(seeds),
|
||||
"clean_accuracy": float(np.mean(clean_accuracy)),
|
||||
"clean_accuracy_sd": float(np.std(clean_accuracy, ddof=1)) if len(clean_accuracy) > 1 else 0.0,
|
||||
"clean_macro_f1": float(np.mean(clean_f1)),
|
||||
"clean_macro_f1_sd": float(np.std(clean_f1, ddof=1)) if len(clean_f1) > 1 else 0.0,
|
||||
"clean_mae": float(np.mean(clean_mae)),
|
||||
"clean_mae_sd": float(np.std(clean_mae, ddof=1)) if len(clean_mae) > 1 else 0.0,
|
||||
"clean_pearson": float(np.mean(clean_pearson)),
|
||||
"clean_pearson_sd": float(np.std(clean_pearson, ddof=1)) if len(clean_pearson) > 1 else 0.0,
|
||||
"corrupt_accuracy_mean": float(np.mean(corrupt_accuracy)),
|
||||
"corrupt_accuracy_sd": float(np.std(corrupt_accuracy, ddof=1)) if len(corrupt_accuracy) > 1 else 0.0,
|
||||
"corrupt_macro_f1_mean": float(np.mean(corrupt_f1)),
|
||||
"corrupt_macro_f1_sd": float(np.std(corrupt_f1, ddof=1)) if len(corrupt_f1) > 1 else 0.0,
|
||||
"corrupt_macro_f1_worst": float(np.min([r["macro_f1"] for r in local])),
|
||||
"corrupt_mae_mean": float(np.mean(corrupt_mae)),
|
||||
"corrupt_mae_sd": float(np.std(corrupt_mae, ddof=1)) if len(corrupt_mae) > 1 else 0.0,
|
||||
"corrupt_pearson_mean": float(np.mean(corrupt_pearson)),
|
||||
"corrupt_pearson_sd": float(np.std(corrupt_pearson, ddof=1)) if len(corrupt_pearson) > 1 else 0.0,
|
||||
}
|
||||
for rate in (0.10, 0.20, 0.30):
|
||||
at_rate = [r for r in local if r["missing_rate"] == rate]
|
||||
f1_by_seed = per_seed_mean(at_rate, "macro_f1")
|
||||
accuracy_by_seed = per_seed_mean(at_rate, "accuracy")
|
||||
mae_by_seed = per_seed_mean(at_rate, "mae")
|
||||
row[f"f1_rate_{int(rate * 100)}"] = float(np.mean(f1_by_seed))
|
||||
row[f"accuracy_rate_{int(rate * 100)}"] = float(np.mean(accuracy_by_seed))
|
||||
row[f"mae_rate_{int(rate * 100)}"] = float(np.mean(mae_by_seed))
|
||||
summary.append(row)
|
||||
for row in summary:
|
||||
row["pareto_nondominated"] = not any(
|
||||
other is not row and other["representation"] == row["representation"]
|
||||
and other["corrupt_macro_f1_mean"] >= row["corrupt_macro_f1_mean"]
|
||||
and other["corrupt_mae_mean"] <= row["corrupt_mae_mean"]
|
||||
and other["corrupt_pearson_mean"] >= row["corrupt_pearson_mean"]
|
||||
and (
|
||||
other["corrupt_macro_f1_mean"] > row["corrupt_macro_f1_mean"]
|
||||
or other["corrupt_mae_mean"] < row["corrupt_mae_mean"]
|
||||
or other["corrupt_pearson_mean"] > row["corrupt_pearson_mean"]
|
||||
)
|
||||
for other in summary
|
||||
)
|
||||
return summary
|
||||
|
||||
|
||||
def _plot(summary: list[dict[str, Any]], rows: list[dict[str, Any]], path: Path) -> None:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
colors = {"concat": "#4e79a7"}
|
||||
fig, axes = plt.subplots(1, 2, figsize=(11, 4.4), constrained_layout=True)
|
||||
for row in summary:
|
||||
kind = row["method"]
|
||||
y_f1 = [row["clean_macro_f1"]] + [row[f"f1_rate_{r}"] for r in (10, 20, 30)]
|
||||
y_mae = [row["clean_mae"]] + [row[f"mae_rate_{r}"] for r in (10, 20, 30)]
|
||||
axes[0].plot([0, 10, 20, 30], y_f1, marker="o", label=kind, color=colors.get(kind))
|
||||
axes[1].plot([0, 10, 20, 30], y_mae, marker="o", label=kind, color=colors.get(kind))
|
||||
axes[0].set(title="Polarity under contiguous local missingness", xlabel="masked slots (%)", ylabel="Macro-F1 (higher is better)")
|
||||
axes[1].set(title="Intensity under contiguous local missingness", xlabel="masked slots (%)", ylabel="MAE (lower is better)")
|
||||
for ax in axes:
|
||||
ax.grid(alpha=0.25)
|
||||
ax.legend(frameon=False)
|
||||
fig.savefig(path, dpi=180)
|
||||
plt.close(fig)
|
||||
|
||||
|
||||
def _sha256(path: Path) -> str:
|
||||
digest = hashlib.sha256()
|
||||
with path.open("rb") as stream:
|
||||
for block in iter(lambda: stream.read(1024 * 1024), b""):
|
||||
digest.update(block)
|
||||
return digest.hexdigest()
|
||||
|
||||
|
||||
def _run(args: argparse.Namespace) -> None:
|
||||
seed_everything(args.seeds[0])
|
||||
if args.device == "auto":
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
else:
|
||||
device = torch.device(args.device)
|
||||
torch.set_num_threads(args.threads)
|
||||
output = Path(args.output_dir)
|
||||
output.mkdir(parents=True, exist_ok=True)
|
||||
aligned_raw = load_aligned()
|
||||
stats = fit_robust_stats(aligned_raw["train"])
|
||||
stats.save(output / "aligned_robust_stats.npz")
|
||||
aligned = {k: apply_robust_stats(v, stats) for k, v in aligned_raw.items()}
|
||||
audit = {
|
||||
"source": str(ATTACHMENT2 / "aligned_50.pkl"),
|
||||
"train_samples": aligned["train"].n,
|
||||
"valid_samples": aligned["valid"].n,
|
||||
"train_classes": np.bincount(aligned["train"].y_cls, minlength=3).tolist(),
|
||||
"valid_classes": np.bincount(aligned["valid"].y_cls, minlength=3).tolist(),
|
||||
"mean_observed_slots": {
|
||||
MODALITIES[m]: float(aligned["train"].mask[:, :, m].sum(axis=1).mean()) for m in range(3)
|
||||
},
|
||||
"train_valid_video_overlap": 0,
|
||||
}
|
||||
with (output / "data_audit.json").open("w", encoding="utf-8") as stream:
|
||||
json.dump(audit, stream, ensure_ascii=False, indent=2)
|
||||
print(f"device={device}; train={audit['train_samples']}; valid={audit['valid_samples']}; audit={audit}", flush=True)
|
||||
|
||||
metric_rows: list[dict[str, Any]] = []
|
||||
best_epochs: dict[str, int] = {}
|
||||
for kind in KINDS:
|
||||
for seed in args.seeds:
|
||||
seed_dir = output / "models" / "aligned" / kind / f"seed_{seed}"
|
||||
model, best_epoch, _ = _train_one(
|
||||
kind, aligned["train"], aligned["valid"], seed_dir,
|
||||
device, seed, args.epochs, args.patience, args.batch_size,
|
||||
)
|
||||
best_epochs[f"{kind}_seed_{seed}"] = best_epoch
|
||||
metric_rows.extend(_eval_conditions(model, aligned["valid"], device, seed + 13, seed, kind, "provided_word_aligned_50"))
|
||||
if seed == args.seeds[0]:
|
||||
shutil.copy2(seed_dir / "model_best.pt", output / "models" / "aligned" / kind / "model_best.pt")
|
||||
del model
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
summary = _summary(metric_rows)
|
||||
selected = sorted(summary, key=lambda r: (-r["corrupt_macro_f1_mean"], r["corrupt_mae_mean"], r["method"]))[0]["method"]
|
||||
(output / "selected_method.txt").write_text(
|
||||
f"Macro-F1-first validation selection: {selected}. See summary.csv for the full multi-metric tradeoff.\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
# Matched audio/vision temporal-shift control for the selected architecture and every seed.
|
||||
for seed in args.seeds:
|
||||
aligned_payload = torch.load(output / "models" / "aligned" / selected / f"seed_{seed}" / "model_best.pt",
|
||||
map_location=device, weights_only=False)
|
||||
aligned_model = AlignedFusionModel(selected, tuple(aligned_payload["dims"])).to(device)
|
||||
aligned_model.load_state_dict(aligned_payload["state_dict"])
|
||||
shifted = shift_audio_vision(aligned["valid"], seed=seed + 2026, max_shift=10)
|
||||
shift_metrics, _ = _score_arrays(aligned_model, shifted, shifted.mask, device)
|
||||
metric_rows.append({"method": selected, "representation": "provided_word_aligned_50", "seed": seed,
|
||||
"condition": "audio_vision_shifted_1_to_10_slots", "missing_rate": 0.0,
|
||||
"n_valid": shifted.n, **shift_metrics})
|
||||
del aligned_model
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# Same selected fusion architecture, but equal-window audio/vision pooling of the unaligned source.
|
||||
print(f"selected_by_corrupt_macro_f1={selected}; starting fixed-window alignment control", flush=True)
|
||||
fixed_raw = load_fixed_window()
|
||||
fixed_stats = fit_robust_stats(fixed_raw["train"])
|
||||
fixed_stats.save(output / "fixed_window_robust_stats.npz")
|
||||
fixed = {k: apply_robust_stats(v, fixed_stats) for k, v in fixed_raw.items()}
|
||||
for seed in args.seeds:
|
||||
fixed_model, fixed_epoch, _ = _train_one(
|
||||
selected, fixed["train"], fixed["valid"], output / "models" / "fixed_window" / selected / f"seed_{seed}",
|
||||
device, seed, args.epochs, args.patience, args.batch_size,
|
||||
)
|
||||
best_epochs[f"fixed_window_{selected}_seed_{seed}"] = fixed_epoch
|
||||
metric_rows.extend(_eval_conditions(fixed_model, fixed["valid"], device, seed + 13, seed, selected,
|
||||
"equal_window_resampled_unaligned"))
|
||||
fixed_shifted = shift_audio_vision(fixed["valid"], seed=seed + 2026, max_shift=10)
|
||||
fixed_shift_metrics, _ = _score_arrays(fixed_model, fixed_shifted, fixed_shifted.mask, device)
|
||||
metric_rows.append({"method": selected, "representation": "equal_window_resampled_unaligned", "seed": seed,
|
||||
"condition": "audio_vision_shifted_1_to_10_slots", "missing_rate": 0.0,
|
||||
"n_valid": fixed_shifted.n, **fixed_shift_metrics})
|
||||
del fixed_model
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
all_summary = _summary(metric_rows)
|
||||
_write_csv(output / "validation_metrics_by_condition.csv", metric_rows)
|
||||
_write_csv(output / "summary.csv", all_summary)
|
||||
aligned_summary = [r for r in all_summary if r["representation"] == "provided_word_aligned_50"]
|
||||
_plot(aligned_summary, metric_rows, output / "missing_rate_comparison.png")
|
||||
alignment_rows = []
|
||||
for rep in ("provided_word_aligned_50", "equal_window_resampled_unaligned"):
|
||||
for condition in ("clean", "audio_vision_shifted_1_to_10_slots"):
|
||||
match = [r for r in metric_rows if r["method"] == selected and r["representation"] == rep
|
||||
and r["condition"] == condition]
|
||||
if match:
|
||||
row = {"method": selected, "representation": rep, "condition": condition,
|
||||
"n_valid": aligned["valid"].n, "n_seeds": len(match)}
|
||||
for metric in ("accuracy", "macro_f1", "mae", "pearson"):
|
||||
values = [r[metric] for r in match]
|
||||
row[metric] = float(np.mean(values))
|
||||
row[f"{metric}_sd"] = float(np.std(values, ddof=1)) if len(values) > 1 else 0.0
|
||||
alignment_rows.append(row)
|
||||
corrupt = [r for r in metric_rows if r["method"] == selected and r["representation"] == rep
|
||||
and r["condition"] != "clean" and r["missing_rate"] > 0]
|
||||
if corrupt:
|
||||
per_seed = []
|
||||
for seed in args.seeds:
|
||||
local = [r for r in corrupt if int(r["seed"]) == seed]
|
||||
if local:
|
||||
per_seed.append({metric: float(np.mean([r[metric] for r in local])) for metric in
|
||||
("accuracy", "macro_f1", "mae", "pearson")})
|
||||
alignment_rows.append({
|
||||
"method": selected, "representation": rep, "condition": "all_local_corruption_mean",
|
||||
"missing_rate": float(np.mean([r["missing_rate"] for r in corrupt])),
|
||||
"n_valid": aligned["valid"].n, "n_seeds": len(per_seed),
|
||||
**{metric: float(np.mean([r[metric] for r in per_seed])) for metric in ("accuracy", "macro_f1", "mae", "pearson")},
|
||||
**{f"{metric}_sd": float(np.std([r[metric] for r in per_seed], ddof=1)) if len(per_seed) > 1 else 0.0
|
||||
for metric in ("accuracy", "macro_f1", "mae", "pearson")},
|
||||
})
|
||||
_write_csv(output / "alignment_transfer_ablation.csv", alignment_rows)
|
||||
|
||||
source_path = ATTACHMENT2 / "aligned_50.pkl"
|
||||
manifest = {
|
||||
"source_feature": str(source_path),
|
||||
"source_sha256": _sha256(source_path),
|
||||
"device": str(device),
|
||||
"cuda_name": torch.cuda.get_device_name(0) if device.type == "cuda" else None,
|
||||
"seeds": args.seeds,
|
||||
"epochs_max": args.epochs,
|
||||
"patience": args.patience,
|
||||
"batch_size": args.batch_size,
|
||||
"best_epochs": best_epochs,
|
||||
"selected_macro_f1_first": selected,
|
||||
"selection_policy": "report Macro-F1, MAE, and Pearson separately; selected model maximizes mean validation Macro-F1 across 15 contiguous corruption conditions, then uses MAE and lexical model name only as tie-breaks",
|
||||
"models": list(KINDS),
|
||||
"corruption_rates": [0.10, 0.20, 0.30],
|
||||
"corruption_patterns": list(PATTERNS),
|
||||
"feature_scaling": "training split median/MAD; fallback to standard deviation for zero-MAD dimensions",
|
||||
"test_labels_used": False,
|
||||
"alignment_transfer_limit": "The official aligned_50 data use a 50-slot wordpiece sequence with no per-slot seconds or stored Q1 B1 time_bounds. The fixed-window comparison is a downstream alignment control, not a re-run of Q1 B1 on the full dataset.",
|
||||
"python": __import__("sys").version,
|
||||
"torch": torch.__version__,
|
||||
"numpy": np.__version__,
|
||||
"created_unix": time.time(),
|
||||
}
|
||||
with (output / "run_manifest.json").open("w", encoding="utf-8") as stream:
|
||||
json.dump(manifest, stream, ensure_ascii=False, indent=2)
|
||||
print(f"saved selection artifacts to {output}; selected={selected}; seeds={args.seeds}", flush=True)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description="Train the EarlyConcat baseline and its alignment-transfer control")
|
||||
parser.add_argument("--seeds", type=int, nargs="+", default=[42, 3407, 2026])
|
||||
parser.add_argument("--epochs", type=int, default=32)
|
||||
parser.add_argument("--patience", type=int, default=6)
|
||||
parser.add_argument("--batch-size", type=int, default=64)
|
||||
parser.add_argument("--threads", type=int, default=4)
|
||||
parser.add_argument("--device", default="auto")
|
||||
parser.add_argument("--output-dir", default=str(Path(__file__).resolve().parents[1] / "outputs" / "followups" / "earlyconcat_standalone"))
|
||||
args = parser.parse_args()
|
||||
_run(args)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,547 @@
|
||||
"""Retrain the two maintained Q2 models under the shared V2 protocol.
|
||||
|
||||
The model architectures and joint CE + SmoothL1 objective stay unchanged.
|
||||
Training masks, official splits, validation scenarios, and final-test handling
|
||||
follow the corresponding Q2 protocol where those choices apply.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import csv
|
||||
import hashlib
|
||||
import json
|
||||
import math
|
||||
import random
|
||||
import time
|
||||
from collections import Counter, defaultdict
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from sklearn.metrics import accuracy_score, f1_score, mean_absolute_error, mean_squared_error
|
||||
from torch import nn
|
||||
|
||||
from .data import ATTACHMENT2, RobustStats, Split, apply_robust_stats, fit_robust_stats
|
||||
from .evaluate_math_protocol import (
|
||||
AURC_BOOTSTRAP_SEED,
|
||||
BOOTSTRAP_REPS,
|
||||
CURVE_MODES,
|
||||
METHODS,
|
||||
SCENARIO_SEED,
|
||||
TEST_BOOTSTRAP_SEED,
|
||||
actual_additional_rates,
|
||||
aurc_from_curve,
|
||||
continuous_mask,
|
||||
curve_scenarios,
|
||||
load_splits,
|
||||
make_scenarios,
|
||||
metrics,
|
||||
scenario_seed,
|
||||
sha256,
|
||||
write_csv,
|
||||
)
|
||||
from .models import AlignedFusionModel
|
||||
from .mofe import MixtureOfFusionExperts
|
||||
from .train_mofe import EARLYCONCAT, MODEL_CONFIG, MOFE7_MLP, _predict
|
||||
from .train_compare import _loss, seed_everything
|
||||
|
||||
|
||||
Q2_ROOT = Path(__file__).resolve().parents[1]
|
||||
OUTPUT_DIR = Q2_ROOT / "outputs" / "followups" / "R03_math_protocol_retraining"
|
||||
SEED = 20260924
|
||||
TRAIN_MASK_SEED = 20261227
|
||||
BATCH_SIZE = 64
|
||||
EPOCH_LIMIT = 12
|
||||
PATIENCE = 3
|
||||
LEARNING_RATE = 3e-4
|
||||
WEIGHT_DECAY = 1e-3
|
||||
SELECTION_SCENARIOS = ("0.0/none", "0.3/single", "0.3/sync", "0.5/async")
|
||||
TRAIN_RATES = (0.0, 0.1, 0.3, 0.5, 0.7)
|
||||
TRAIN_MODES = ("single", "sync", "partial", "async")
|
||||
|
||||
|
||||
def device_for(name: str) -> torch.device:
|
||||
if name == "auto":
|
||||
return torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
return torch.device(name)
|
||||
|
||||
|
||||
def set_deterministic(seed: int) -> None:
|
||||
seed_everything(seed)
|
||||
torch.set_num_threads(4)
|
||||
torch.backends.cudnn.deterministic = True
|
||||
torch.backends.cudnn.benchmark = False
|
||||
|
||||
|
||||
def build_model(method: str, dims: tuple[int, int, int], device: torch.device) -> nn.Module:
|
||||
if method == EARLYCONCAT:
|
||||
return AlignedFusionModel("concat", dims=dims).to(device)
|
||||
if method == MOFE7_MLP:
|
||||
return MixtureOfFusionExperts(dims=dims, **MODEL_CONFIG).to(device)
|
||||
raise ValueError(f"unknown method: {method}")
|
||||
|
||||
|
||||
def model_state(model: nn.Module, method: str) -> dict[str, Any]:
|
||||
state: dict[str, Any] = {
|
||||
"method": method,
|
||||
"dims": tuple(int(x) for x in model_dims(model)),
|
||||
"state_dict": model.state_dict(),
|
||||
"seed": SEED,
|
||||
"protocol": "Q2 V2 adapted deterministic-model training",
|
||||
}
|
||||
if method == EARLYCONCAT:
|
||||
state["kind"] = "concat"
|
||||
else:
|
||||
state["config"] = MODEL_CONFIG
|
||||
return state
|
||||
|
||||
|
||||
def model_dims(model: nn.Module) -> tuple[int, int, int]:
|
||||
if isinstance(model, AlignedFusionModel):
|
||||
return tuple(layer[0].in_features for layer in model.projections) # type: ignore[return-value]
|
||||
if isinstance(model, MixtureOfFusionExperts):
|
||||
return tuple(layer[0].in_features for layer in model.private_projections) # type: ignore[return-value]
|
||||
raise TypeError(type(model))
|
||||
|
||||
|
||||
def train_masks_for_epoch(split: Split, epoch: int) -> tuple[np.ndarray, Counter[str]]:
|
||||
"""Sample reproducible math-protocol rates/patterns per training example."""
|
||||
rows: list[np.ndarray] = []
|
||||
counts: Counter[str] = Counter()
|
||||
for sample_id, observed in zip(split.ids, split.mask):
|
||||
rng = np.random.default_rng(scenario_seed(TRAIN_MASK_SEED + SEED, sample_id, f"train/{epoch}"))
|
||||
rate = float(rng.choice(TRAIN_RATES))
|
||||
mode = str(rng.choice(TRAIN_MODES))
|
||||
key = f"{rate:.1f}/{mode}"
|
||||
counts[key] += 1
|
||||
row = continuous_mask(observed, rate, mode, rng)
|
||||
rows.append(row)
|
||||
return np.stack(rows), counts
|
||||
|
||||
|
||||
def _batched_loss(
|
||||
model: nn.Module,
|
||||
split: Split,
|
||||
masks: np.ndarray,
|
||||
device: torch.device,
|
||||
batch_size: int,
|
||||
) -> float:
|
||||
model.eval()
|
||||
losses: list[float] = []
|
||||
weights: list[int] = []
|
||||
with torch.inference_mode():
|
||||
for start in range(0, split.n, batch_size):
|
||||
end = min(start + batch_size, split.n)
|
||||
xs = tuple(torch.as_tensor(x[start:end], dtype=torch.float32, device=device) for x in split.x)
|
||||
mb = torch.as_tensor(masks[start:end], dtype=torch.bool, device=device)
|
||||
y_cls = torch.as_tensor(split.y_cls[start:end], dtype=torch.long, device=device)
|
||||
y_reg = torch.as_tensor(split.y_reg[start:end], dtype=torch.float32, device=device)
|
||||
losses.append(float(_loss(model(xs, mb), y_cls, y_reg).item()))
|
||||
weights.append(end - start)
|
||||
return float(np.average(losses, weights=weights))
|
||||
|
||||
|
||||
def selection_loss(model: nn.Module, valid: Split, scenarios: dict[str, np.ndarray], device: torch.device) -> float:
|
||||
return float(np.mean([
|
||||
_batched_loss(model, valid, scenarios[key], device, BATCH_SIZE)
|
||||
for key in SELECTION_SCENARIOS
|
||||
]))
|
||||
|
||||
|
||||
def train_one(
|
||||
method: str,
|
||||
train: Split,
|
||||
valid: Split,
|
||||
valid_scenarios: dict[str, np.ndarray],
|
||||
orders: list[np.ndarray],
|
||||
output_dir: Path,
|
||||
device: torch.device,
|
||||
) -> tuple[nn.Module, int, list[dict[str, Any]], Counter[str]]:
|
||||
set_deterministic(SEED)
|
||||
model = build_model(method, tuple(x.shape[-1] for x in train.x), device)
|
||||
optimizer = torch.optim.AdamW(model.parameters(), lr=LEARNING_RATE, weight_decay=WEIGHT_DECAY)
|
||||
xs = tuple(torch.as_tensor(x, dtype=torch.float32, device=device) for x in train.x)
|
||||
y_cls = torch.as_tensor(train.y_cls, dtype=torch.long, device=device)
|
||||
y_reg = torch.as_tensor(train.y_reg, dtype=torch.float32, device=device)
|
||||
checkpoint_path = output_dir / "model_best.pt"
|
||||
history: list[dict[str, Any]] = []
|
||||
train_mask_counts: Counter[str] = Counter()
|
||||
best_loss = math.inf
|
||||
best_epoch = 0
|
||||
stale = 0
|
||||
|
||||
for epoch in range(1, EPOCH_LIMIT + 1):
|
||||
model.train()
|
||||
epoch_masks, epoch_counts = train_masks_for_epoch(train, epoch)
|
||||
train_mask_counts.update(epoch_counts)
|
||||
batch_losses: list[float] = []
|
||||
order = orders[epoch - 1]
|
||||
for start in range(0, train.n, BATCH_SIZE):
|
||||
indices_np = order[start:start + BATCH_SIZE]
|
||||
indices = torch.as_tensor(indices_np, dtype=torch.long, device=device)
|
||||
mb = torch.as_tensor(epoch_masks[indices_np], dtype=torch.bool, device=device)
|
||||
output = model(tuple(x.index_select(0, indices) for x in xs), mb)
|
||||
loss = _loss(output, y_cls.index_select(0, indices), y_reg.index_select(0, indices))
|
||||
optimizer.zero_grad(set_to_none=True)
|
||||
loss.backward()
|
||||
nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
|
||||
optimizer.step()
|
||||
batch_losses.append(float(loss.detach().item()))
|
||||
|
||||
valid_selection_loss = selection_loss(model, valid, valid_scenarios, device)
|
||||
row = {
|
||||
"method": method,
|
||||
"seed": SEED,
|
||||
"epoch": epoch,
|
||||
"train_loss": float(np.mean(batch_losses)),
|
||||
"valid_selection_loss": valid_selection_loss,
|
||||
"valid_clean_loss": _batched_loss(model, valid, valid.mask, device, BATCH_SIZE),
|
||||
}
|
||||
history.append(row)
|
||||
print(
|
||||
f"[{method}] epoch={epoch:02d} train={row['train_loss']:.4f} "
|
||||
f"valid_selection={valid_selection_loss:.4f} clean={row['valid_clean_loss']:.4f}",
|
||||
flush=True,
|
||||
)
|
||||
if valid_selection_loss < best_loss - 1e-4:
|
||||
best_loss = valid_selection_loss
|
||||
best_epoch = epoch
|
||||
stale = 0
|
||||
torch.save(model_state(model, method) | {"best_epoch": best_epoch}, checkpoint_path)
|
||||
else:
|
||||
stale += 1
|
||||
if stale >= PATIENCE:
|
||||
break
|
||||
|
||||
saved = torch.load(checkpoint_path, map_location=device, weights_only=False)
|
||||
model.load_state_dict(saved["state_dict"])
|
||||
model.eval()
|
||||
write_csv(output_dir / "training_history.csv", history)
|
||||
return model, best_epoch, history, train_mask_counts
|
||||
|
||||
|
||||
def _group_map(ids: list[str]) -> tuple[list[str], dict[str, np.ndarray]]:
|
||||
source_ids = [sample_id.split("$_$", 1)[0] for sample_id in ids]
|
||||
groups = sorted(set(source_ids))
|
||||
mapping = {
|
||||
group: np.flatnonzero(np.asarray([source == group for source in source_ids]))
|
||||
for group in groups
|
||||
}
|
||||
return groups, mapping
|
||||
|
||||
|
||||
def test_group_bootstrap(test: Split, predictions: dict[str, dict[str, np.ndarray]]) -> list[dict[str, Any]]:
|
||||
groups, mapping = _group_map(test.ids)
|
||||
rng = np.random.default_rng(TEST_BOOTSTRAP_SEED)
|
||||
draws: dict[str, list[float]] = defaultdict(list)
|
||||
for _ in range(BOOTSTRAP_REPS):
|
||||
selected = rng.choice(groups, size=len(groups), replace=True)
|
||||
indices = np.concatenate([mapping[group] for group in selected])
|
||||
values = {
|
||||
method: metrics(test, predictions[method]["logits"], predictions[method]["intensity"], indices)
|
||||
for method in METHODS
|
||||
}
|
||||
for name in values[EARLYCONCAT]:
|
||||
draws[name].append(values[MOFE7_MLP][name] - values[EARLYCONCAT][name])
|
||||
point = {
|
||||
name: metrics(test, predictions[MOFE7_MLP]["logits"], predictions[MOFE7_MLP]["intensity"])[name]
|
||||
- metrics(test, predictions[EARLYCONCAT]["logits"], predictions[EARLYCONCAT]["intensity"])[name]
|
||||
for name in draws
|
||||
}
|
||||
return [{
|
||||
"comparison": f"{MOFE7_MLP} minus {EARLYCONCAT}",
|
||||
"metric": name,
|
||||
"delta": point[name],
|
||||
"bootstrap_ci_2p5": float(np.quantile(values, 0.025)),
|
||||
"bootstrap_ci_97p5": float(np.quantile(values, 0.975)),
|
||||
"bootstrap_probability_delta_gt_0": float(np.mean(np.asarray(values) > 0)),
|
||||
"replicates": BOOTSTRAP_REPS,
|
||||
"resampling_unit": "source video id",
|
||||
"paired": True,
|
||||
"seed": TEST_BOOTSTRAP_SEED,
|
||||
} for name, values in draws.items()]
|
||||
|
||||
|
||||
def validation_aurc_bootstrap(
|
||||
valid: Split,
|
||||
predictions: dict[tuple[str, str], dict[str, np.ndarray]],
|
||||
rates_by_sample: dict[str, np.ndarray],
|
||||
) -> list[dict[str, Any]]:
|
||||
groups, mapping = _group_map(valid.ids)
|
||||
rng = np.random.default_rng(AURC_BOOTSTRAP_SEED)
|
||||
deltas: dict[str, list[float]] = {mode: [] for mode in CURVE_MODES}
|
||||
|
||||
def score(method: str, mode: str, indices: np.ndarray) -> float:
|
||||
keys = curve_scenarios(mode)
|
||||
xs = [float(np.nanmean(rates_by_sample[key][indices])) for key in keys]
|
||||
ys = [
|
||||
float(np.abs(valid.y_reg[indices] - predictions[(method, key)]["intensity"][indices]).mean())
|
||||
for key in keys
|
||||
]
|
||||
return aurc_from_curve(xs, ys)
|
||||
|
||||
for _ in range(BOOTSTRAP_REPS):
|
||||
selected = rng.choice(groups, size=len(groups), replace=True)
|
||||
indices = np.concatenate([mapping[group] for group in selected])
|
||||
for mode in CURVE_MODES:
|
||||
deltas[mode].append(score(MOFE7_MLP, mode, indices) - score(EARLYCONCAT, mode, indices))
|
||||
rows = []
|
||||
for mode in CURVE_MODES:
|
||||
all_indices = np.arange(valid.n)
|
||||
values = deltas[mode]
|
||||
rows.append({
|
||||
"mask_mode": mode,
|
||||
"delta_aurc_mae_mofe_minus_earlyconcat": score(MOFE7_MLP, mode, all_indices) - score(EARLYCONCAT, mode, all_indices),
|
||||
"bootstrap_ci_2p5": float(np.quantile(values, 0.025)),
|
||||
"bootstrap_ci_97p5": float(np.quantile(values, 0.975)),
|
||||
"bootstrap_probability_delta_lt_0": float(np.mean(np.asarray(values) < 0)),
|
||||
"replicates": BOOTSTRAP_REPS,
|
||||
"resampling_unit": "source video id",
|
||||
"paired": True,
|
||||
"seed": AURC_BOOTSTRAP_SEED,
|
||||
})
|
||||
return rows
|
||||
|
||||
|
||||
def run(device_name: str = "auto", output_dir: Path = OUTPUT_DIR,
|
||||
input_version: str = "aligned_50", batch_size: int = BATCH_SIZE) -> None:
|
||||
global BATCH_SIZE
|
||||
if batch_size < 1:
|
||||
raise ValueError("batch_size must be positive")
|
||||
BATCH_SIZE = batch_size
|
||||
if output_dir.exists() and any(output_dir.iterdir()):
|
||||
raise FileExistsError(f"refusing to overwrite non-empty result directory: {output_dir}")
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
device = device_for(device_name)
|
||||
if device.type == "cuda" and not torch.cuda.is_available():
|
||||
raise RuntimeError("CUDA was requested but is unavailable")
|
||||
|
||||
if input_version not in {"aligned_50", "unaligned_50"}:
|
||||
raise ValueError(f"unsupported input version: {input_version}")
|
||||
feature_path = ATTACHMENT2 / f"{input_version}.pkl"
|
||||
if input_version == "unaligned_50":
|
||||
from ....adapter import adapt_official_split
|
||||
from .data import _unpickle, _ids_and_targets
|
||||
|
||||
source = _unpickle(feature_path)
|
||||
raw_splits = {}
|
||||
adapter_audit = {}
|
||||
for name in ("train", "valid", "test"):
|
||||
arrays, mask, audit = adapt_official_split(source[name])
|
||||
ids, y_cls, y_reg = _ids_and_targets(source[name])
|
||||
raw_splits[name] = Split(tuple(arrays[m] for m in ("text", "audio", "vision")),
|
||||
mask, y_cls, y_reg, ids)
|
||||
adapter_audit[name] = audit
|
||||
del source
|
||||
groups = {name: {sid.split("$_$", 1)[0] for sid in split.ids}
|
||||
for name, split in raw_splits.items()}
|
||||
if any(groups[a] & groups[b] for a, b in (("train", "valid"), ("train", "test"), ("valid", "test"))):
|
||||
raise ValueError("official source-video groups overlap")
|
||||
else:
|
||||
raw_splits = load_splits(feature_path)
|
||||
adapter_audit = None
|
||||
train_raw, valid_raw, test_raw = raw_splits["train"], raw_splits["valid"], raw_splits["test"]
|
||||
stats = fit_robust_stats(train_raw)
|
||||
train, valid, test = (apply_robust_stats(s, stats) for s in (train_raw, valid_raw, test_raw))
|
||||
stats_path = output_dir / f"{input_version}_robust_stats.npz"
|
||||
stats.save(stats_path)
|
||||
dims = tuple(int(x.shape[-1]) for x in train.x)
|
||||
valid_scenarios = make_scenarios(valid, SCENARIO_SEED)
|
||||
if len(valid_scenarios) != 42:
|
||||
raise ValueError(f"expected 42 controlled scenarios, got {len(valid_scenarios)}")
|
||||
rates_by_sample = actual_additional_rates(valid.mask, valid_scenarios)
|
||||
|
||||
set_deterministic(SEED)
|
||||
order_rng = np.random.default_rng(SEED + 809)
|
||||
orders = [order_rng.permutation(train.n) for _ in range(EPOCH_LIMIT)]
|
||||
best_epochs: dict[str, int] = {}
|
||||
training_rows: list[dict[str, Any]] = []
|
||||
mask_count_rows: list[dict[str, Any]] = []
|
||||
parameter_rows: list[dict[str, Any]] = []
|
||||
|
||||
for method in METHODS:
|
||||
model_dir = output_dir / "models" / method / f"seed_{SEED}"
|
||||
model_dir.mkdir(parents=True, exist_ok=True)
|
||||
model, best_epoch, history, mask_counts = train_one(
|
||||
method, train, valid, valid_scenarios, orders, model_dir, device
|
||||
)
|
||||
best_epochs[method] = best_epoch
|
||||
training_rows.extend(history)
|
||||
parameter_rows.append({
|
||||
"method": method,
|
||||
"parameters_total": sum(p.numel() for p in model.parameters()),
|
||||
"parameters_trainable": sum(p.numel() for p in model.parameters() if p.requires_grad),
|
||||
"best_epoch": best_epoch,
|
||||
})
|
||||
for key, count in sorted(mask_counts.items()):
|
||||
mask_count_rows.append({"method": method, "seed": SEED, "rate_mode": key, "sample_epoch_assignments": count})
|
||||
del model
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
write_csv(output_dir / "training_history.csv", training_rows)
|
||||
write_csv(output_dir / "training_mask_distribution.csv", mask_count_rows)
|
||||
write_csv(output_dir / "parameter_count.csv", parameter_rows)
|
||||
|
||||
# Reload the selected checkpoints, then conduct one final official-test pass.
|
||||
test_predictions: dict[str, dict[str, np.ndarray]] = {}
|
||||
test_rows: list[dict[str, Any]] = []
|
||||
condition_predictions: dict[tuple[str, str], dict[str, np.ndarray]] = {}
|
||||
condition_rows: list[dict[str, Any]] = []
|
||||
for method in METHODS:
|
||||
checkpoint_path = output_dir / "models" / method / f"seed_{SEED}" / "model_best.pt"
|
||||
saved = torch.load(checkpoint_path, map_location=device, weights_only=False)
|
||||
model = build_model(method, dims, device)
|
||||
model.load_state_dict(saved["state_dict"])
|
||||
model.eval()
|
||||
|
||||
test_prediction = _predict(model, test, test.mask, device, BATCH_SIZE)
|
||||
test_predictions[method] = test_prediction
|
||||
test_rows.append({
|
||||
"method": method,
|
||||
"seed": SEED,
|
||||
"best_epoch": best_epochs[method],
|
||||
"n_test": test.n,
|
||||
**metrics(test, test_prediction["logits"], test_prediction["intensity"]),
|
||||
})
|
||||
|
||||
for scenario, masks in valid_scenarios.items():
|
||||
prediction = _predict(model, valid, masks, device, BATCH_SIZE)
|
||||
condition_predictions[(method, scenario)] = prediction
|
||||
condition_rows.append({
|
||||
"method": method,
|
||||
"seed": SEED,
|
||||
"scenario": scenario,
|
||||
"realized_additional_global_rate": float(np.nanmean(rates_by_sample[scenario])),
|
||||
"n_valid": valid.n,
|
||||
**metrics(valid, prediction["logits"], prediction["intensity"]),
|
||||
})
|
||||
print(f"[valid/{method}] {scenario} done", flush=True)
|
||||
del model
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
write_csv(output_dir / "official_test_metrics_by_seed.csv", test_rows)
|
||||
write_csv(output_dir / "official_test_paired_bootstrap.csv", test_group_bootstrap(test, test_predictions))
|
||||
write_csv(output_dir / "controlled_metrics_by_scenario.csv", condition_rows)
|
||||
|
||||
test_summary = []
|
||||
for method in METHODS:
|
||||
row = next(r for r in test_rows if r["method"] == method)
|
||||
for metric in ("accuracy", "macro_f1", "mae", "rmse", "pearson"):
|
||||
test_summary.append({"method": method, "metric": metric, "mean": row[metric], "sd_across_seeds": 0.0, "n_seeds": 1})
|
||||
write_csv(output_dir / "official_test_summary.csv", test_summary)
|
||||
|
||||
aurc_rows: list[dict[str, Any]] = []
|
||||
for method in METHODS:
|
||||
for mode in CURVE_MODES:
|
||||
keys = curve_scenarios(mode)
|
||||
xs = [float(np.nanmean(rates_by_sample[key])) for key in keys]
|
||||
ys = [
|
||||
float(np.abs(valid.y_reg - condition_predictions[(method, key)]["intensity"]).mean())
|
||||
for key in keys
|
||||
]
|
||||
aurc_rows.append({
|
||||
"method": method,
|
||||
"seed": SEED,
|
||||
"mask_mode": mode,
|
||||
"aurc_mae": aurc_from_curve(xs, ys),
|
||||
"rates_realized": json.dumps(xs),
|
||||
})
|
||||
write_csv(output_dir / "aurc_mae_by_mode_seed.csv", aurc_rows)
|
||||
write_csv(output_dir / "aurc_mae_paired_bootstrap.csv", validation_aurc_bootstrap(valid, condition_predictions, rates_by_sample))
|
||||
|
||||
manifest = {
|
||||
"experiment": "Retrained EarlyConcat and MoFE-7 + MLP Router using the shared Q2 V2 protocol",
|
||||
"created_unix": time.time(),
|
||||
"device": str(device),
|
||||
"cuda_device": torch.cuda.get_device_name(0) if device.type == "cuda" else None,
|
||||
"feature_file": str(feature_path),
|
||||
"feature_sha256": sha256(feature_path),
|
||||
"representation": ("shared Q1 adapter relative-progress projection of official unaligned_50; not physical-time alignment"
|
||||
if input_version == "unaligned_50" else
|
||||
"official aligned_50 ordered positions; not Q1 physical-time bins"),
|
||||
"adapter": "Q1 adapter relative-progress projection" if input_version == "unaligned_50" else None,
|
||||
"adapter_audit": adapter_audit,
|
||||
"train_valid_test_counts": {name: split.n for name, split in raw_splits.items()},
|
||||
"source_video_groups": {name: len({sid.split("$_$", 1)[0] for sid in split.ids}) for name, split in raw_splits.items()},
|
||||
"official_group_splits_disjoint": True,
|
||||
"train_only_scaler": str(stats_path),
|
||||
"scaler_fit": "median and 1.4826*MAD on observed training rows only; zero-MAD fallback to std then 1",
|
||||
"seed": SEED,
|
||||
"model_seeds": [SEED],
|
||||
"training_configuration": {
|
||||
"epoch_limit": EPOCH_LIMIT,
|
||||
"early_stopping_patience": PATIENCE,
|
||||
"batch_size": BATCH_SIZE,
|
||||
"optimizer": "AdamW",
|
||||
"learning_rate": LEARNING_RATE,
|
||||
"weight_decay": WEIGHT_DECAY,
|
||||
"gradient_clip_norm": 1.0,
|
||||
"early_stopping_metric": "mean validation joint CE + 0.5*SmoothL1 over 0.0/none, 0.3/single, 0.3/sync, 0.5/async",
|
||||
"architecture_preserved": {
|
||||
EARLYCONCAT: "EarlyConcat + BiGRU",
|
||||
MOFE7_MLP: "MoFE-7 + MLP Router",
|
||||
},
|
||||
"objective": "cross entropy + 0.5 * SmoothL1(intensity/3, label/3); same objective for both methods",
|
||||
"training_corruption": {
|
||||
"rates": list(TRAIN_RATES),
|
||||
"patterns": list(TRAIN_MODES),
|
||||
"preserve_at_least_fraction_per_selected_modality": 0.2,
|
||||
"generator_seed": TRAIN_MASK_SEED,
|
||||
"same_sample_masks_and_batch_orders_across_models": True,
|
||||
},
|
||||
},
|
||||
"validation_protocol": {
|
||||
"scenario_seed": SCENARIO_SEED,
|
||||
"scenario_count": len(valid_scenarios),
|
||||
"same_fixed_masks_for_both_models": True,
|
||||
"scenario_design": "42 controlled continuous-mask scenarios regenerated on each sample's original observation mask",
|
||||
"selection_scenarios": list(SELECTION_SCENARIOS),
|
||||
"selection_note": "Deterministic-model adaptation; uses joint supervised loss instead of C5's probabilistic selection NLL.",
|
||||
"aurc": "normalized trapezoidal MAE area over realized equal-modality-weighted additional missing rate for single/sync/partial/async at 0/.1/.3/.5/.7",
|
||||
},
|
||||
"test_protocol": {
|
||||
"official_test_final_clean_passes": 1,
|
||||
"test_used_for_training_or_checkpoint_selection": False,
|
||||
"metrics": ["accuracy", "macro_f1", "mae", "rmse", "pearson"],
|
||||
"paired_group_bootstrap_replicates": BOOTSTRAP_REPS,
|
||||
"bootstrap_unit": "source video id",
|
||||
"bootstrap_seed": TEST_BOOTSTRAP_SEED,
|
||||
},
|
||||
}
|
||||
(output_dir / "run_manifest.json").write_text(json.dumps(manifest, indent=2), encoding="utf-8")
|
||||
(output_dir / "hypothesis.md").write_text(
|
||||
"# R03: 按统一 Q2 V2 口径重训两种保留模型\n\n"
|
||||
"## 假设\n\n"
|
||||
"在保持 EarlyConcat + BiGRU 与 MoFE-7 + MLP Router 结构及共同监督目标不变的情况下,"
|
||||
"使用数学方案中的官方划分、连续块缺失训练和 42 个固定验证情景,可以公平比较两种模型的干净测试表现与缺失鲁棒性。\n\n"
|
||||
"## 唯一实验改动\n\n"
|
||||
"相对现有检查点,本轮重新训练时将缺失训练改为 0/10/30/50/70% 与 single/sync/partial/async,"
|
||||
"每个被选模态至少保留 20% 观测;训练和批次顺序在两个模型间配对。数学方案中的 C5 概率损失不适用于现有确定性分类/回归头,"
|
||||
"因此保留项目既有的 CE + 0.5 SmoothL1 联合目标。\n\n"
|
||||
"## 数据使用\n\n"
|
||||
"标准化器只在官方训练集观测行上拟合;官方验证集只用于早停与缺失评估;官方测试集在全部检查点确定后做一次干净评估。\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
print(f"wrote retraining results to {output_dir}", flush=True)
|
||||
print(f"train/valid/test={train.n}/{valid.n}/{test.n}; device={device}; best_epochs={best_epochs}", flush=True)
|
||||
for row in test_rows:
|
||||
print(
|
||||
f"{row['method']}: Acc={row['accuracy']:.4f} Macro-F1={row['macro_f1']:.4f} "
|
||||
f"MAE={row['mae']:.4f} RMSE={row['rmse']:.4f} Pearson={row['pearson']:.4f}",
|
||||
flush=True,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--device", default="auto", choices=("auto", "cuda", "cpu"))
|
||||
parser.add_argument("--output-dir", type=Path, default=OUTPUT_DIR)
|
||||
parser.add_argument("--input-version", choices=("aligned_50", "unaligned_50"), default="unaligned_50")
|
||||
parser.add_argument("--batch-size", type=int, default=BATCH_SIZE)
|
||||
arguments = parser.parse_args()
|
||||
run(device_name=arguments.device, output_dir=arguments.output_dir,
|
||||
input_version=arguments.input_version, batch_size=arguments.batch_size)
|
||||
@@ -0,0 +1,787 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import csv
|
||||
import hashlib
|
||||
import json
|
||||
import math
|
||||
import random
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from sklearn.metrics import accuracy_score, f1_score, mean_absolute_error
|
||||
from torch import nn
|
||||
|
||||
from .data import (
|
||||
ATTACHMENT2,
|
||||
MODALITIES,
|
||||
RobustStats,
|
||||
Split,
|
||||
apply_robust_stats,
|
||||
augment_masks,
|
||||
corrupt_masks,
|
||||
fit_robust_stats,
|
||||
load_aligned,
|
||||
)
|
||||
from .models import AlignedFusionModel
|
||||
from .mofe import EXPERT_NAMES, SUBSETS, MixtureOfFusionExperts
|
||||
from .train_compare import PATTERNS, _loss, _pearson, _train_one, seed_everything
|
||||
|
||||
|
||||
ROOT = Path(__file__).resolve().parents[1]
|
||||
REFERENCE_OUTPUT = ROOT / "outputs" / "mofe_7experts"
|
||||
DEFAULT_OUTPUT = ROOT / "outputs" / "mofe_7experts"
|
||||
EARLYCONCAT = "B0_early_concat"
|
||||
MOFE7_MLP = "B5_mofe_mlp"
|
||||
SEEDS = (42, 3407, 2026)
|
||||
RATES = (0.10, 0.20, 0.30)
|
||||
HIDDEN = 128
|
||||
LATENT_DIM = 64
|
||||
MODEL_CONFIG: dict[str, Any] = {
|
||||
"router": "mlp",
|
||||
"expert_names": EXPERT_NAMES,
|
||||
"availability_mode": "hard",
|
||||
}
|
||||
SUMMARY_METRICS = (
|
||||
"corrupt_macro_f1",
|
||||
"worst_condition_macro_f1",
|
||||
"text_30_macro_f1",
|
||||
"corrupt_mae",
|
||||
"corrupt_pearson",
|
||||
)
|
||||
|
||||
|
||||
def _write_csv(path: Path, rows: list[dict[str, Any]]) -> None:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
if not rows:
|
||||
return
|
||||
fields = list(dict.fromkeys(key for row in rows for key in row))
|
||||
with path.open("w", newline="", encoding="utf-8-sig") as stream:
|
||||
writer = csv.DictWriter(stream, fieldnames=fields)
|
||||
writer.writeheader()
|
||||
writer.writerows(rows)
|
||||
|
||||
|
||||
def _read_csv(path: Path) -> list[dict[str, str]]:
|
||||
if not path.exists():
|
||||
return []
|
||||
with path.open("r", newline="", encoding="utf-8-sig") as stream:
|
||||
return list(csv.DictReader(stream))
|
||||
|
||||
|
||||
def _sha256(path: Path) -> str:
|
||||
digest = hashlib.sha256()
|
||||
with path.open("rb") as stream:
|
||||
for block in iter(lambda: stream.read(1024 * 1024), b""):
|
||||
digest.update(block)
|
||||
return digest.hexdigest()
|
||||
|
||||
|
||||
def _device_for(name: str) -> torch.device:
|
||||
if name == "auto":
|
||||
return torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
return torch.device(name)
|
||||
|
||||
|
||||
def _conditions(valid: Split, seed: int) -> list[tuple[str, float, np.ndarray]]:
|
||||
rows = [("clean", 0.0, valid.mask.copy())]
|
||||
for rate in RATES:
|
||||
for pattern_idx, (pattern, modalities) in enumerate(PATTERNS.items()):
|
||||
masks = corrupt_masks(
|
||||
valid.mask,
|
||||
rate,
|
||||
modalities,
|
||||
seed + 13 + pattern_idx * 101 + int(rate * 1000),
|
||||
)
|
||||
rows.append((f"{pattern}_{int(rate * 100)}", rate, masks))
|
||||
return rows
|
||||
|
||||
|
||||
def _metric_dict(
|
||||
y_cls: np.ndarray,
|
||||
y_reg: np.ndarray,
|
||||
logits: np.ndarray,
|
||||
intensity: np.ndarray,
|
||||
) -> dict[str, float]:
|
||||
predicted_class = np.asarray(logits).argmax(axis=-1)
|
||||
predicted_intensity = np.clip(np.asarray(intensity).reshape(-1), -3.0, 3.0)
|
||||
return {
|
||||
"accuracy": float(accuracy_score(y_cls, predicted_class)),
|
||||
"macro_f1": float(f1_score(y_cls, predicted_class, labels=[0, 1, 2], average="macro", zero_division=0)),
|
||||
"mae": float(mean_absolute_error(y_reg, predicted_intensity)),
|
||||
"pearson": _pearson(y_reg, predicted_intensity),
|
||||
}
|
||||
|
||||
|
||||
def _validation_loss(model: nn.Module, valid: Split, device: torch.device, batch_size: int) -> float:
|
||||
model.eval()
|
||||
values: list[float] = []
|
||||
weights: list[int] = []
|
||||
with torch.inference_mode():
|
||||
for start in range(0, valid.n, batch_size):
|
||||
end = min(start + batch_size, valid.n)
|
||||
xs = tuple(torch.as_tensor(x[start:end], dtype=torch.float32, device=device) for x in valid.x)
|
||||
masks = torch.as_tensor(valid.mask[start:end], dtype=torch.bool, device=device)
|
||||
y_cls = torch.as_tensor(valid.y_cls[start:end], dtype=torch.long, device=device)
|
||||
y_reg = torch.as_tensor(valid.y_reg[start:end], dtype=torch.float32, device=device)
|
||||
values.append(float(_loss(model(xs, masks), y_cls, y_reg).item()))
|
||||
weights.append(end - start)
|
||||
return float(np.average(values, weights=weights))
|
||||
|
||||
|
||||
def _train_mofe(
|
||||
train: Split,
|
||||
valid: Split,
|
||||
output_dir: Path,
|
||||
device: torch.device,
|
||||
seed: int,
|
||||
epochs: int,
|
||||
patience: int,
|
||||
batch_size: int,
|
||||
reuse_checkpoint: bool,
|
||||
) -> tuple[MixtureOfFusionExperts, int, list[dict[str, Any]]]:
|
||||
dims = tuple(int(x.shape[-1]) for x in train.x)
|
||||
checkpoint_path = output_dir / "model_best.pt"
|
||||
history_path = output_dir / "training_history.csv"
|
||||
if reuse_checkpoint and checkpoint_path.exists():
|
||||
saved = torch.load(checkpoint_path, map_location=device, weights_only=False)
|
||||
if saved.get("config") != MODEL_CONFIG or tuple(saved.get("dims", ())) != dims or int(saved.get("seed", -1)) != seed:
|
||||
raise ValueError(f"cached MoFE checkpoint does not match the selected configuration: {checkpoint_path}")
|
||||
model = MixtureOfFusionExperts(dims=dims, **MODEL_CONFIG).to(device)
|
||||
model.load_state_dict(saved["state_dict"])
|
||||
history = [
|
||||
{"method": MOFE7_MLP, "seed": seed, **{key: float(value) for key, value in row.items() if key in {"epoch", "train_loss", "valid_clean_loss"}}}
|
||||
for row in _read_csv(history_path)
|
||||
]
|
||||
return model.eval(), int(saved.get("best_epoch", 0)), history
|
||||
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
seed_everything(seed)
|
||||
model = MixtureOfFusionExperts(dims=dims, **MODEL_CONFIG).to(device)
|
||||
optimizer = torch.optim.AdamW(model.parameters(), lr=1.5e-4, weight_decay=1e-4)
|
||||
xs = tuple(torch.as_tensor(x, dtype=torch.float32, device=device) for x in train.x)
|
||||
base_masks = train.mask
|
||||
y_cls = torch.as_tensor(train.y_cls, dtype=torch.long, device=device)
|
||||
y_reg = torch.as_tensor(train.y_reg, dtype=torch.float32, device=device)
|
||||
rng = np.random.default_rng(seed + 809)
|
||||
best_loss = math.inf
|
||||
best_epoch = 0
|
||||
stale_epochs = 0
|
||||
history: list[dict[str, Any]] = []
|
||||
|
||||
for epoch in range(1, epochs + 1):
|
||||
model.train()
|
||||
order = rng.permutation(train.n)
|
||||
batch_losses: list[float] = []
|
||||
for start in range(0, train.n, batch_size):
|
||||
ids_np = order[start:start + batch_size]
|
||||
ids = torch.as_tensor(ids_np, dtype=torch.long, device=device)
|
||||
masks_np = augment_masks(base_masks[ids_np], rng)
|
||||
masks = torch.as_tensor(masks_np, dtype=torch.bool, device=device)
|
||||
output = model(tuple(x.index_select(0, ids) for x in xs), masks)
|
||||
loss = _loss(output, y_cls.index_select(0, ids), y_reg.index_select(0, ids))
|
||||
optimizer.zero_grad(set_to_none=True)
|
||||
loss.backward()
|
||||
nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
|
||||
optimizer.step()
|
||||
batch_losses.append(float(loss.detach().item()))
|
||||
|
||||
valid_loss = _validation_loss(model, valid, device, batch_size)
|
||||
row = {
|
||||
"method": MOFE7_MLP,
|
||||
"seed": seed,
|
||||
"epoch": epoch,
|
||||
"train_loss": float(np.mean(batch_losses)),
|
||||
"valid_clean_loss": valid_loss,
|
||||
}
|
||||
history.append(row)
|
||||
print(f"[MoFE-7 MLP] seed={seed} epoch={epoch:02d} train={row['train_loss']:.4f} valid={valid_loss:.4f}", flush=True)
|
||||
if valid_loss < best_loss - 1e-4:
|
||||
best_loss = valid_loss
|
||||
best_epoch = epoch
|
||||
stale_epochs = 0
|
||||
torch.save({
|
||||
"method": MOFE7_MLP,
|
||||
"config": MODEL_CONFIG,
|
||||
"dims": dims,
|
||||
"state_dict": model.state_dict(),
|
||||
"seed": seed,
|
||||
"best_epoch": epoch,
|
||||
}, checkpoint_path)
|
||||
else:
|
||||
stale_epochs += 1
|
||||
if stale_epochs >= patience:
|
||||
break
|
||||
|
||||
saved = torch.load(checkpoint_path, map_location=device, weights_only=False)
|
||||
model.load_state_dict(saved["state_dict"])
|
||||
model.eval()
|
||||
_write_csv(history_path, history)
|
||||
return model, best_epoch, history
|
||||
|
||||
|
||||
def _load_or_train_concat(
|
||||
train: Split,
|
||||
valid: Split,
|
||||
output_dir: Path,
|
||||
device: torch.device,
|
||||
seed: int,
|
||||
epochs: int,
|
||||
patience: int,
|
||||
batch_size: int,
|
||||
reuse_checkpoint: bool,
|
||||
) -> tuple[AlignedFusionModel, int, list[dict[str, Any]]]:
|
||||
checkpoint_path = output_dir / "model_best.pt"
|
||||
dims = tuple(int(x.shape[-1]) for x in train.x)
|
||||
if reuse_checkpoint and checkpoint_path.exists():
|
||||
saved = torch.load(checkpoint_path, map_location=device, weights_only=False)
|
||||
if saved.get("kind") != "concat" or tuple(saved.get("dims", ())) != dims or int(saved.get("seed", -1)) != seed:
|
||||
raise ValueError(f"cached EarlyConcat checkpoint does not match: {checkpoint_path}")
|
||||
model = AlignedFusionModel("concat", dims=dims).to(device)
|
||||
model.load_state_dict(saved["state_dict"])
|
||||
history = [
|
||||
{"method": EARLYCONCAT, "seed": seed, **{key: float(value) for key, value in row.items() if key in {"epoch", "train_loss", "valid_clean_loss"}}}
|
||||
for row in _read_csv(output_dir / "training_history.csv")
|
||||
]
|
||||
return model.eval(), int(saved.get("best_epoch", 0)), history
|
||||
|
||||
model, best_epoch, history = _train_one(
|
||||
"concat", train, valid, output_dir, device, seed, epochs, patience, batch_size
|
||||
)
|
||||
rows = [{"method": EARLYCONCAT, "seed": seed, **row} for row in history]
|
||||
return model.eval(), best_epoch, rows
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def _predict(
|
||||
model: nn.Module,
|
||||
split: Split,
|
||||
masks: np.ndarray,
|
||||
device: torch.device,
|
||||
batch_size: int,
|
||||
force_expert: str | None = None,
|
||||
) -> dict[str, np.ndarray]:
|
||||
fields = ["logits", "intensity"]
|
||||
if isinstance(model, MixtureOfFusionExperts):
|
||||
fields.extend(("alpha", "utility", "availability", "fallback"))
|
||||
chunks: dict[str, list[np.ndarray]] = {name: [] for name in fields}
|
||||
for start in range(0, split.n, batch_size):
|
||||
end = min(start + batch_size, split.n)
|
||||
xs = tuple(torch.as_tensor(x[start:end], dtype=torch.float32, device=device) for x in split.x)
|
||||
mask_batch = torch.as_tensor(masks[start:end], dtype=torch.bool, device=device)
|
||||
output = model(xs, mask_batch, force_expert=force_expert) if isinstance(model, MixtureOfFusionExperts) else model(xs, mask_batch)
|
||||
for name in fields:
|
||||
value = output[name]
|
||||
chunks[name].append(value.float().cpu().numpy())
|
||||
result = {name: np.concatenate(values, axis=0) for name, values in chunks.items()}
|
||||
result["intensity"] = np.clip(result["intensity"].reshape(-1), -3.0, 3.0)
|
||||
return result
|
||||
|
||||
|
||||
def _condition_row(
|
||||
method: str,
|
||||
seed: int,
|
||||
condition: str,
|
||||
rate: float,
|
||||
split: Split,
|
||||
prediction: dict[str, np.ndarray],
|
||||
) -> dict[str, Any]:
|
||||
return {
|
||||
"method": method,
|
||||
"seed": seed,
|
||||
"condition": condition,
|
||||
"missing_rate": rate,
|
||||
"n_valid": split.n,
|
||||
**_metric_dict(split.y_cls, split.y_reg, prediction["logits"], prediction["intensity"]),
|
||||
}
|
||||
|
||||
|
||||
def _diagnostics(
|
||||
seed: int,
|
||||
condition: str,
|
||||
masks: np.ndarray,
|
||||
prediction: dict[str, np.ndarray],
|
||||
) -> tuple[dict[str, Any], dict[str, Any]]:
|
||||
alpha = prediction["alpha"]
|
||||
availability = prediction["availability"].astype(bool)
|
||||
active = availability.any(axis=-1)
|
||||
active_alpha = alpha[active]
|
||||
if active_alpha.size:
|
||||
means = active_alpha.mean(axis=0)
|
||||
entropy = -(active_alpha * np.log(np.maximum(active_alpha, 1e-12))).sum(axis=-1) / np.log(len(EXPERT_NAMES))
|
||||
high_weight = (active_alpha.max(axis=-1) > 0.8).mean()
|
||||
else:
|
||||
means = np.zeros(len(EXPERT_NAMES), dtype=np.float64)
|
||||
entropy = np.zeros(0, dtype=np.float64)
|
||||
high_weight = 0.0
|
||||
route_row: dict[str, Any] = {
|
||||
"method": MOFE7_MLP,
|
||||
"seed": seed,
|
||||
"condition": condition,
|
||||
"active_position_fraction": float(active.mean()),
|
||||
"fallback_position_fraction": float((~active).mean()),
|
||||
"normalized_router_entropy": float(entropy.mean()) if entropy.size else 0.0,
|
||||
"fraction_active_positions_max_weight_over_0p8": float(high_weight),
|
||||
}
|
||||
for index, name in enumerate(EXPERT_NAMES):
|
||||
route_row[f"alpha_{name}_mean"] = float(means[index])
|
||||
utility = prediction["utility"]
|
||||
utility_row: dict[str, Any] = {"method": MOFE7_MLP, "seed": seed, "condition": condition}
|
||||
for modality, name in enumerate(MODALITIES):
|
||||
observed = masks[..., modality]
|
||||
utility_row[f"utility_{name}_mean"] = float(utility[..., modality][observed].mean()) if observed.any() else 0.0
|
||||
return route_row, utility_row
|
||||
|
||||
|
||||
def _summary_rows(rows: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
||||
summaries: list[dict[str, Any]] = []
|
||||
for method in (EARLYCONCAT, MOFE7_MLP):
|
||||
matching = [row for row in rows if row["method"] == method]
|
||||
seeds = sorted({int(row["seed"]) for row in matching})
|
||||
conditions = list(dict.fromkeys(row["condition"] for row in matching))
|
||||
by_seed_condition = {(int(row["seed"]), row["condition"]): row for row in matching}
|
||||
clean = [by_seed_condition[(seed, "clean")] for seed in seeds]
|
||||
corrupt_conditions = [condition for condition in conditions if condition != "clean"]
|
||||
corrupt_by_seed = {
|
||||
seed: [by_seed_condition[(seed, condition)] for condition in corrupt_conditions]
|
||||
for seed in seeds
|
||||
}
|
||||
condition_f1 = {
|
||||
condition: float(np.mean([by_seed_condition[(seed, condition)]["macro_f1"] for seed in seeds]))
|
||||
for condition in corrupt_conditions
|
||||
}
|
||||
worst_condition = min(condition_f1, key=condition_f1.get)
|
||||
row: dict[str, Any] = {"method": method, "n_seeds": len(seeds), "worst_condition": worst_condition}
|
||||
for metric in ("accuracy", "macro_f1", "mae", "pearson"):
|
||||
clean_values = [float(item[metric]) for item in clean]
|
||||
corrupt_values = [float(np.mean([item[metric] for item in corrupt_by_seed[seed]])) for seed in seeds]
|
||||
row[f"clean_{metric}"] = float(np.mean(clean_values))
|
||||
row[f"clean_{metric}_sd"] = float(np.std(clean_values, ddof=1)) if len(clean_values) > 1 else 0.0
|
||||
row[f"corrupt_{metric}_mean"] = float(np.mean(corrupt_values))
|
||||
row[f"corrupt_{metric}_sd"] = float(np.std(corrupt_values, ddof=1)) if len(corrupt_values) > 1 else 0.0
|
||||
row["worst_condition_macro_f1"] = condition_f1[worst_condition]
|
||||
row["worst_single_run_macro_f1"] = min(
|
||||
item["macro_f1"] for seed in seeds for item in corrupt_by_seed[seed]
|
||||
)
|
||||
text_30 = [by_seed_condition[(seed, "text_30")] for seed in seeds]
|
||||
row["text_30_macro_f1"] = float(np.mean([item["macro_f1"] for item in text_30]))
|
||||
row["text_30_macro_f1_sd"] = float(np.std([item["macro_f1"] for item in text_30], ddof=1)) if len(text_30) > 1 else 0.0
|
||||
for condition in ("audio_30", "vision_30", "audio_vision_30", "all_modalities_30"):
|
||||
values = [by_seed_condition[(seed, condition)]["macro_f1"] for seed in seeds]
|
||||
row[f"{condition}_macro_f1"] = float(np.mean(values))
|
||||
row[f"{condition}_macro_f1_sd"] = float(np.std(values, ddof=1)) if len(values) > 1 else 0.0
|
||||
row["corrupt_macro_f1"] = row["corrupt_macro_f1_mean"]
|
||||
row["corrupt_mae"] = row["corrupt_mae_mean"]
|
||||
row["corrupt_pearson"] = row["corrupt_pearson_mean"]
|
||||
summaries.append(row)
|
||||
return summaries
|
||||
|
||||
|
||||
def _bootstrap_distributions(
|
||||
method: str,
|
||||
predictions: dict[tuple[str, int, str], dict[str, np.ndarray]],
|
||||
valid: Split,
|
||||
seeds: list[int],
|
||||
conditions: list[str],
|
||||
group_counts: np.ndarray,
|
||||
) -> dict[str, np.ndarray]:
|
||||
group_names = sorted({sample_id.split("$_$", 1)[0] for sample_id in valid.ids})
|
||||
group_index = {name: index for index, name in enumerate(group_names)}
|
||||
row_group = np.asarray([group_index[sample_id.split("$_$", 1)[0]] for sample_id in valid.ids], dtype=np.int64)
|
||||
n_groups = len(group_names)
|
||||
n_slots = len(seeds) * len(conditions)
|
||||
confusion_by_group = np.zeros((n_groups, n_slots, 9), dtype=np.float64)
|
||||
regression_by_group = np.zeros((n_groups, n_slots, 7), dtype=np.float64)
|
||||
for seed_index, seed in enumerate(seeds):
|
||||
for condition_index, condition in enumerate(conditions):
|
||||
slot = seed_index * len(conditions) + condition_index
|
||||
pred = predictions[(method, seed, condition)]
|
||||
predicted_class = pred["logits"].argmax(axis=-1)
|
||||
code = valid.y_cls * 3 + predicted_class
|
||||
np.add.at(confusion_by_group[:, slot, :], (row_group, code), 1.0)
|
||||
intensity = np.clip(pred["intensity"].reshape(-1), -3.0, 3.0)
|
||||
values = np.stack((
|
||||
np.ones(valid.n),
|
||||
np.abs(valid.y_reg - intensity),
|
||||
valid.y_reg,
|
||||
valid.y_reg ** 2,
|
||||
intensity,
|
||||
intensity ** 2,
|
||||
valid.y_reg * intensity,
|
||||
), axis=-1)
|
||||
for statistic in range(values.shape[-1]):
|
||||
np.add.at(regression_by_group[:, slot, statistic], row_group, values[:, statistic])
|
||||
|
||||
weighted_confusion = np.einsum("rg,gsk->rsk", group_counts, confusion_by_group, optimize=True)
|
||||
cm = weighted_confusion.reshape(len(group_counts), len(seeds), len(conditions), 3, 3)
|
||||
true_count = cm.sum(axis=-1)
|
||||
predicted_count = cm.sum(axis=-2)
|
||||
true_positive = np.diagonal(cm, axis1=-2, axis2=-1)
|
||||
denominator = true_count + predicted_count
|
||||
class_f1 = np.divide(2.0 * true_positive, denominator, out=np.zeros_like(true_positive), where=denominator > 0)
|
||||
macro_f1 = class_f1.mean(axis=-1)
|
||||
|
||||
weighted_regression = np.einsum("rg,gsk->rsk", group_counts, regression_by_group, optimize=True)
|
||||
regression = weighted_regression.reshape(len(group_counts), len(seeds), len(conditions), 7)
|
||||
count = np.maximum(regression[..., 0], 1.0)
|
||||
mae = regression[..., 1] / count
|
||||
sum_y, sum_y2, sum_pred, sum_pred2, sum_yp = (regression[..., index] for index in range(2, 7))
|
||||
covariance = sum_yp - sum_y * sum_pred / count
|
||||
variance_y = np.maximum(sum_y2 - sum_y ** 2 / count, 0.0)
|
||||
variance_pred = np.maximum(sum_pred2 - sum_pred ** 2 / count, 0.0)
|
||||
denominator_corr = np.sqrt(variance_y * variance_pred)
|
||||
pearson = np.divide(covariance, denominator_corr, out=np.zeros_like(covariance), where=denominator_corr > 1e-12)
|
||||
text_30_index = conditions.index("text_30")
|
||||
return {
|
||||
"corrupt_macro_f1": macro_f1[:, :, 1:].mean(axis=(1, 2)),
|
||||
"worst_condition_macro_f1": macro_f1[:, :, 1:].mean(axis=1).min(axis=1),
|
||||
"text_30_macro_f1": macro_f1[:, :, text_30_index].mean(axis=1),
|
||||
"corrupt_mae": mae[:, :, 1:].mean(axis=(1, 2)),
|
||||
"corrupt_pearson": pearson[:, :, 1:].mean(axis=(1, 2)),
|
||||
}
|
||||
|
||||
|
||||
def _paired_bootstrap(
|
||||
predictions: dict[tuple[str, int, str], dict[str, np.ndarray]],
|
||||
valid: Split,
|
||||
seeds: list[int],
|
||||
conditions: list[str],
|
||||
reps: int,
|
||||
bootstrap_seed: int,
|
||||
summaries: list[dict[str, Any]],
|
||||
) -> list[dict[str, Any]]:
|
||||
groups = sorted({sample_id.split("$_$", 1)[0] for sample_id in valid.ids})
|
||||
rng = np.random.default_rng(bootstrap_seed)
|
||||
draws = rng.integers(0, len(groups), size=(reps, len(groups)))
|
||||
group_counts = np.zeros((reps, len(groups)), dtype=np.float64)
|
||||
for rep in range(reps):
|
||||
group_counts[rep] = np.bincount(draws[rep], minlength=len(groups))
|
||||
candidate = _bootstrap_distributions(MOFE7_MLP, predictions, valid, seeds, conditions, group_counts)
|
||||
reference = _bootstrap_distributions(EARLYCONCAT, predictions, valid, seeds, conditions, group_counts)
|
||||
summary_map = {row["method"]: row for row in summaries}
|
||||
point_keys = {
|
||||
"corrupt_macro_f1": "corrupt_macro_f1",
|
||||
"worst_condition_macro_f1": "worst_condition_macro_f1",
|
||||
"text_30_macro_f1": "text_30_macro_f1",
|
||||
"corrupt_mae": "corrupt_mae",
|
||||
"corrupt_pearson": "corrupt_pearson",
|
||||
}
|
||||
rows = []
|
||||
for metric in SUMMARY_METRICS:
|
||||
delta = candidate[metric] - reference[metric]
|
||||
key = point_keys[metric]
|
||||
rows.append({
|
||||
"comparison": "MoFE-7 MLP vs EarlyConcat",
|
||||
"candidate": MOFE7_MLP,
|
||||
"reference": EARLYCONCAT,
|
||||
"metric": metric,
|
||||
"delta_candidate_minus_reference": float(summary_map[MOFE7_MLP][key] - summary_map[EARLYCONCAT][key]),
|
||||
"bootstrap_ci_2p5": float(np.quantile(delta, 0.025)),
|
||||
"bootstrap_ci_97p5": float(np.quantile(delta, 0.975)),
|
||||
"bootstrap_probability_delta_gt_0": float(np.mean(delta > 0.0)),
|
||||
"bootstrap_replicates": reps,
|
||||
"resampling_unit": "source video id",
|
||||
"paired": True,
|
||||
"seed": bootstrap_seed,
|
||||
})
|
||||
return rows
|
||||
|
||||
|
||||
def _plot_summary(output: Path, summaries: list[dict[str, Any]]) -> None:
|
||||
import matplotlib
|
||||
matplotlib.use("Agg")
|
||||
import matplotlib.pyplot as plt
|
||||
|
||||
labels = ["EarlyConcat + BiGRU", "MoFE-7 + MLP Router"]
|
||||
by_method = {row["method"]: row for row in summaries}
|
||||
methods = (EARLYCONCAT, MOFE7_MLP)
|
||||
metrics = ("clean_macro_f1", "corrupt_macro_f1", "worst_condition_macro_f1")
|
||||
names = ("Clean", "Mean corrupted", "Worst condition")
|
||||
x = np.arange(len(names))
|
||||
width = 0.34
|
||||
fig, ax = plt.subplots(figsize=(8.6, 4.8), constrained_layout=True)
|
||||
for offset, method, label, color in (
|
||||
(-width / 2, methods[0], labels[0], "#4e79a7"),
|
||||
(width / 2, methods[1], labels[1], "#f28e2b"),
|
||||
):
|
||||
values = [by_method[method][metric] for metric in metrics]
|
||||
ax.bar(x + offset, values, width, label=label, color=color)
|
||||
ax.set_xticks(x, names)
|
||||
ax.set_ylabel("Macro-F1")
|
||||
ax.set_ylim(0, 1)
|
||||
ax.set_title("Q2 selected-model validation comparison")
|
||||
ax.legend(frameon=False)
|
||||
output.mkdir(parents=True, exist_ok=True)
|
||||
fig.savefig(output / "comparison_earlyconcat_mofe7.png", dpi=180)
|
||||
plt.close(fig)
|
||||
|
||||
|
||||
def _parameter_rows(dims: tuple[int, int, int], device: torch.device) -> list[dict[str, Any]]:
|
||||
models: dict[str, nn.Module] = {
|
||||
EARLYCONCAT: AlignedFusionModel("concat", dims=dims),
|
||||
MOFE7_MLP: MixtureOfFusionExperts(dims=dims, **MODEL_CONFIG),
|
||||
}
|
||||
baseline_count = sum(parameter.numel() for parameter in models[EARLYCONCAT].parameters() if parameter.requires_grad)
|
||||
rows = []
|
||||
for name, model in models.items():
|
||||
count = sum(parameter.numel() for parameter in model.parameters() if parameter.requires_grad)
|
||||
rows.append({
|
||||
"method": name,
|
||||
"trainable_parameters": count,
|
||||
"ratio_to_earlyconcat": count / baseline_count,
|
||||
"within_2x_earlyconcat": bool(count <= 2 * baseline_count),
|
||||
})
|
||||
return rows
|
||||
|
||||
|
||||
def _smoke_test(train: Split, output: Path, device: torch.device, seed: int) -> dict[str, Any]:
|
||||
seed_everything(seed)
|
||||
dims = tuple(int(x.shape[-1]) for x in train.x)
|
||||
count = min(4, train.n)
|
||||
xs = tuple(torch.as_tensor(x[:count], dtype=torch.float32, device=device) for x in train.x)
|
||||
masks = torch.as_tensor(train.mask[:count].copy(), dtype=torch.bool, device=device)
|
||||
masks[0] = True
|
||||
if count > 1:
|
||||
masks[1, 5:12, 0] = False
|
||||
if count > 2:
|
||||
masks[2, 18:23, :] = False
|
||||
target_class = torch.as_tensor(train.y_cls[:count], dtype=torch.long, device=device)
|
||||
target_intensity = torch.as_tensor(train.y_reg[:count], dtype=torch.float32, device=device)
|
||||
reports: dict[str, Any] = {}
|
||||
models: dict[str, nn.Module] = {
|
||||
EARLYCONCAT: AlignedFusionModel("concat", dims=dims).to(device),
|
||||
MOFE7_MLP: MixtureOfFusionExperts(dims=dims, **MODEL_CONFIG).to(device),
|
||||
}
|
||||
for name, model in models.items():
|
||||
model.train()
|
||||
result = model(xs, masks)
|
||||
loss = _loss(result, target_class, target_intensity)
|
||||
loss.backward()
|
||||
gradient = sum(float(p.grad.detach().abs().sum().cpu()) for p in model.parameters() if p.grad is not None)
|
||||
reports[name] = {
|
||||
"logits_shape": list(result["logits"].shape),
|
||||
"intensity_shape": list(result["intensity"].shape),
|
||||
"finite_loss": bool(torch.isfinite(loss).item()),
|
||||
"gradient_l1": gradient,
|
||||
}
|
||||
mofe_model = models[MOFE7_MLP]
|
||||
mo = mofe_model(xs, masks)
|
||||
active = mo["availability"].any(dim=-1)
|
||||
alpha_sums = mo["alpha"].sum(dim=-1)
|
||||
alpha_error = float((alpha_sums[active] - 1).abs().max().cpu()) if active.any() else 0.0
|
||||
unavailable_weights = float(mo["alpha"].masked_select(~mo["availability"]).abs().max().cpu()) if (~mo["availability"]).any() else 0.0
|
||||
expert_gradients = {
|
||||
name: sum(float(parameter.grad.detach().abs().sum().cpu()) for parameter in expert.parameters() if parameter.grad is not None)
|
||||
for name, expert in mofe_model.experts.items()
|
||||
}
|
||||
router_gradient = sum(float(parameter.grad.detach().abs().sum().cpu()) for parameter in mofe_model.router.parameters() if parameter.grad is not None)
|
||||
if alpha_error > 1e-6 or unavailable_weights > 1e-8:
|
||||
raise RuntimeError(f"MoFE routing mask invariant failed: sum_error={alpha_error}, unavailable={unavailable_weights}")
|
||||
if not all(value > 0 for value in expert_gradients.values()) or router_gradient <= 0:
|
||||
raise RuntimeError(f"MoFE expert/router gradients are incomplete: {expert_gradients}; router={router_gradient}")
|
||||
report = {
|
||||
"passed": all(item["finite_loss"] and item["gradient_l1"] > 0 for item in reports.values()),
|
||||
"seed": seed,
|
||||
"device": str(device),
|
||||
"cuda_device": torch.cuda.get_device_name(0) if device.type == "cuda" else None,
|
||||
"batch_size_checked": count,
|
||||
"steps": train.steps,
|
||||
"models": reports,
|
||||
"mofe_experts": list(EXPERT_NAMES),
|
||||
"mofe_alpha_shape": list(mo["alpha"].shape),
|
||||
"mofe_max_weight_sum_error": alpha_error,
|
||||
"mofe_max_weight_on_unavailable_experts": unavailable_weights,
|
||||
"mofe_expert_gradient_l1": expert_gradients,
|
||||
"mofe_router_gradient_l1": router_gradient,
|
||||
"parameter_count": {row["method"]: row["trainable_parameters"] for row in _parameter_rows(dims, device)},
|
||||
}
|
||||
output.mkdir(parents=True, exist_ok=True)
|
||||
(output / "smoke_test.json").write_text(json.dumps(report, indent=2), encoding="utf-8")
|
||||
return report
|
||||
|
||||
|
||||
def _run(args: argparse.Namespace) -> None:
|
||||
output = args.output_dir.resolve()
|
||||
output.mkdir(parents=True, exist_ok=True)
|
||||
device = _device_for(args.device)
|
||||
torch.set_num_threads(args.threads)
|
||||
torch.backends.cudnn.deterministic = True
|
||||
torch.backends.cudnn.benchmark = False
|
||||
|
||||
raw = load_aligned()
|
||||
computed_stats = fit_robust_stats(raw["train"])
|
||||
reference_stats_path = REFERENCE_OUTPUT / "aligned_robust_stats.npz"
|
||||
if reference_stats_path.exists():
|
||||
stats = RobustStats.load(reference_stats_path)
|
||||
scaler_diff = max(
|
||||
max(float(np.max(np.abs(a - b))) for a, b in zip(computed_stats.center, stats.center)),
|
||||
max(float(np.max(np.abs(a - b))) for a, b in zip(computed_stats.scale, stats.scale)),
|
||||
)
|
||||
else:
|
||||
stats = computed_stats
|
||||
scaler_diff = 0.0
|
||||
train = apply_robust_stats(raw["train"], stats)
|
||||
valid = apply_robust_stats(raw["valid"], stats)
|
||||
stats.save(output / "aligned_robust_stats.npz")
|
||||
dims = tuple(int(x.shape[-1]) for x in train.x)
|
||||
feature_path = ATTACHMENT2 / "aligned_50.pkl"
|
||||
if not feature_path.exists():
|
||||
raise FileNotFoundError(f"official aligned feature file not found: {feature_path}")
|
||||
|
||||
if args.phase == "smoke":
|
||||
report = _smoke_test(train, output, device, args.seeds[0])
|
||||
report["scaler_max_abs_difference_from_reference"] = scaler_diff
|
||||
(output / "smoke_test.json").write_text(json.dumps(report, indent=2), encoding="utf-8")
|
||||
print(f"selected-model smoke: passed={report['passed']} device={device}", flush=True)
|
||||
return
|
||||
|
||||
seeds = list(args.seeds)
|
||||
metrics_rows: list[dict[str, Any]] = []
|
||||
predictions: dict[tuple[str, int, str], dict[str, np.ndarray]] = {}
|
||||
router_rows: list[dict[str, Any]] = []
|
||||
utility_rows: list[dict[str, Any]] = []
|
||||
expert_rows: list[dict[str, Any]] = []
|
||||
history_rows: list[dict[str, Any]] = []
|
||||
best_epochs: dict[str, int] = {}
|
||||
condition_names: list[str] = []
|
||||
|
||||
for seed in seeds:
|
||||
baseline_dir = output / "models" / "baselines" / "concat" / f"seed_{seed}"
|
||||
baseline, baseline_epoch, baseline_history = _load_or_train_concat(
|
||||
train, valid, baseline_dir, device, seed, args.epochs, args.patience,
|
||||
args.batch_size, args.reuse_checkpoints and not args.force_retrain,
|
||||
)
|
||||
mofe_dir = output / "models" / MOFE7_MLP / f"seed_{seed}"
|
||||
mofe, mofe_epoch, mofe_history = _train_mofe(
|
||||
train, valid, mofe_dir, device, seed, args.epochs, args.patience,
|
||||
args.batch_size, args.reuse_checkpoints and not args.force_retrain,
|
||||
)
|
||||
best_epochs[f"{EARLYCONCAT}_seed_{seed}"] = baseline_epoch
|
||||
best_epochs[f"{MOFE7_MLP}_seed_{seed}"] = mofe_epoch
|
||||
history_rows.extend(baseline_history)
|
||||
history_rows.extend(mofe_history)
|
||||
|
||||
conditions = _conditions(valid, seed)
|
||||
names = [condition for condition, _, _ in conditions]
|
||||
if condition_names and names != condition_names:
|
||||
raise RuntimeError("validation condition ordering changed between seeds")
|
||||
condition_names = names
|
||||
for method, model in ((EARLYCONCAT, baseline), (MOFE7_MLP, mofe)):
|
||||
for condition, rate, masks in conditions:
|
||||
prediction = _predict(model, valid, masks, device, args.batch_size)
|
||||
predictions[(method, seed, condition)] = prediction
|
||||
metrics_rows.append(_condition_row(method, seed, condition, rate, valid, prediction))
|
||||
if method == MOFE7_MLP:
|
||||
route_row, utility_row = _diagnostics(seed, condition, masks, prediction)
|
||||
router_rows.append(route_row)
|
||||
utility_rows.append(utility_row)
|
||||
for expert in EXPERT_NAMES:
|
||||
forced = _predict(model, valid, masks, device, args.batch_size, force_expert=expert)
|
||||
observed = masks[..., list(SUBSETS[expert])].all(axis=-1)
|
||||
expert_rows.append({
|
||||
"method": MOFE7_MLP,
|
||||
"seed": seed,
|
||||
"condition": condition,
|
||||
"expert": expert,
|
||||
"available_position_fraction": float(observed.mean()),
|
||||
**_metric_dict(valid.y_cls, valid.y_reg, forced["logits"], forced["intensity"]),
|
||||
})
|
||||
print(f"evaluated {method}/seed{seed}", flush=True)
|
||||
del baseline, mofe
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
summaries = _summary_rows(metrics_rows)
|
||||
paired = _paired_bootstrap(
|
||||
predictions, valid, seeds, condition_names, args.bootstrap_reps,
|
||||
args.bootstrap_seed, summaries,
|
||||
) if args.bootstrap_reps > 0 else []
|
||||
parameter_rows = _parameter_rows(dims, device)
|
||||
output_rows = {
|
||||
"metrics_by_condition.csv": metrics_rows,
|
||||
"summary.csv": summaries,
|
||||
"paired_bootstrap.csv": paired,
|
||||
"parameter_count.csv": parameter_rows,
|
||||
"router_weights_by_condition.csv": router_rows,
|
||||
"routing_entropy.csv": router_rows,
|
||||
"modality_utility_by_condition.csv": utility_rows,
|
||||
"expert_condition_matrix.csv": expert_rows,
|
||||
"training_history.csv": history_rows,
|
||||
}
|
||||
for filename, rows in output_rows.items():
|
||||
_write_csv(output / filename, rows)
|
||||
_plot_summary(output / "figures", summaries)
|
||||
|
||||
manifest = {
|
||||
"experiment": "Q2 selected models: EarlyConcat + BiGRU and MoFE-7 + MLP Router",
|
||||
"created_unix": time.time(),
|
||||
"python_version": sys.version,
|
||||
"torch_version": torch.__version__,
|
||||
"numpy_version": np.__version__,
|
||||
"device": str(device),
|
||||
"cuda_device": torch.cuda.get_device_name(0) if device.type == "cuda" else None,
|
||||
"feature_file": str(feature_path),
|
||||
"feature_sha256": _sha256(feature_path),
|
||||
"feature_dimensions": dict(zip(MODALITIES, dims)),
|
||||
"sequence_length": train.steps,
|
||||
"representation_note": "official ordered 50-wordpiece positions; not 50 physical-time bins",
|
||||
"train_examples": train.n,
|
||||
"valid_examples": valid.n,
|
||||
"train_source_video_groups": len({sample_id.split("$_$", 1)[0] for sample_id in train.ids}),
|
||||
"valid_source_video_groups": len({sample_id.split("$_$", 1)[0] for sample_id in valid.ids}),
|
||||
"train_only_scaler": str(output / "aligned_robust_stats.npz"),
|
||||
"scaler_max_abs_difference_from_reference": scaler_diff,
|
||||
"test_labels_used": False,
|
||||
"seeds": seeds,
|
||||
"epochs_max": args.epochs,
|
||||
"patience": args.patience,
|
||||
"batch_size": args.batch_size,
|
||||
"optimizer": "AdamW(lr=1.5e-4, weight_decay=1e-4), gradient clip 1.0",
|
||||
"training_mask_augmentation": "same contiguous-block augment_masks protocol for both models",
|
||||
"validation_conditions": condition_names,
|
||||
"validation_corruption_seed": "seed + 13 + pattern_index*101 + int(rate*1000)",
|
||||
"loss": "cross_entropy + 0.5*SmoothL1(intensity/3, regression_label/3)",
|
||||
"models": {
|
||||
EARLYCONCAT: "project modalities independently, concatenate features and masks, then BiGRU",
|
||||
MOFE7_MLP: {
|
||||
"experts": list(EXPERT_NAMES),
|
||||
"router": "MLP over per-position observed values and local observation statistics",
|
||||
"availability": "hard mask; unavailable expert weights are zero",
|
||||
"shared_temporal_backbone": "one BiGRU after position-wise expert mixture",
|
||||
},
|
||||
},
|
||||
"best_epochs": best_epochs,
|
||||
"paired_bootstrap": {
|
||||
"replicates": args.bootstrap_reps,
|
||||
"seed": args.bootstrap_seed,
|
||||
"resampling_unit": "source video id",
|
||||
"paired": True,
|
||||
},
|
||||
}
|
||||
(output / "run_manifest.json").write_text(json.dumps(manifest, indent=2), encoding="utf-8")
|
||||
print(f"selected-model results saved to {output}", flush=True)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description="Train and compare the two retained Q2 models.")
|
||||
parser.add_argument("--phase", choices=("smoke", "full"), default="full")
|
||||
parser.add_argument("--epochs", type=int, default=32)
|
||||
parser.add_argument("--patience", type=int, default=6)
|
||||
parser.add_argument("--batch-size", type=int, default=64)
|
||||
parser.add_argument("--threads", type=int, default=4)
|
||||
parser.add_argument("--device", default="auto")
|
||||
parser.add_argument("--seeds", type=int, nargs="+", default=list(SEEDS))
|
||||
parser.add_argument("--bootstrap-reps", type=int, default=1000)
|
||||
parser.add_argument("--bootstrap-seed", type=int, default=20260924)
|
||||
parser.add_argument("--reuse-checkpoints", action="store_true")
|
||||
parser.add_argument("--force-retrain", action="store_true")
|
||||
parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT)
|
||||
_run(parser.parse_args())
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1 @@
|
||||
"""Mathematical Q2 model family and training driver."""
|
||||
@@ -0,0 +1,213 @@
|
||||
"""Restricted readers and split preparation for the official Q2 inputs."""
|
||||
from __future__ import annotations
|
||||
|
||||
import pickle
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
|
||||
from ...data_paths import ATTACHMENT2, ATTACHMENT3, DATA_ROOT, PROJECT_ROOT
|
||||
|
||||
ROOT = PROJECT_ROOT
|
||||
ATTACHMENT2_DIR = ATTACHMENT2
|
||||
ALIGNED_PATH = ATTACHMENT2 / "aligned_50.pkl"
|
||||
ATTACHMENT3_ALIGNED = ATTACHMENT3 / "对齐版本"
|
||||
ATTACHMENT3_UNALIGNED = ATTACHMENT3 / "未对齐版本"
|
||||
MODALITIES = ("text", "audio", "vision")
|
||||
EXPECTED_DIMS = {"text": 768, "audio": 74, "vision": 35}
|
||||
|
||||
|
||||
class RestrictedUnpickler(pickle.Unpickler):
|
||||
"""Allow only primitive containers and NumPy reconstruction primitives."""
|
||||
|
||||
_allowed = {
|
||||
("builtins", name): getattr(__import__("builtins"), name)
|
||||
for name in ("set", "frozenset", "slice", "complex", "bytearray")
|
||||
}
|
||||
_allowed.update({
|
||||
("collections", "OrderedDict"): __import__("collections").OrderedDict,
|
||||
("numpy", "ndarray"): np.ndarray,
|
||||
("numpy", "dtype"): np.dtype,
|
||||
("numpy", "asarray"): np.asarray,
|
||||
("numpy.core.multiarray", "_reconstruct"): np.core.multiarray._reconstruct,
|
||||
("numpy.core.multiarray", "scalar"): np.core.multiarray.scalar,
|
||||
("numpy._core.multiarray", "_reconstruct"): np.core.multiarray._reconstruct,
|
||||
("numpy._core.multiarray", "scalar"): np.core.multiarray.scalar,
|
||||
})
|
||||
if hasattr(np.core.numeric, "_frombuffer"):
|
||||
_allowed[("numpy.core.numeric", "_frombuffer")] = np.core.numeric._frombuffer
|
||||
_allowed[("numpy._core.numeric", "_frombuffer")] = np.core.numeric._frombuffer
|
||||
|
||||
def find_class(self, module: str, name: str) -> Any:
|
||||
try:
|
||||
return self._allowed[(module, name)]
|
||||
except KeyError as exc:
|
||||
raise pickle.UnpicklingError(f"blocked pickle global: {module}.{name}") from exc
|
||||
|
||||
|
||||
def restricted_load(path: Path) -> Any:
|
||||
with path.open("rb") as stream:
|
||||
return RestrictedUnpickler(stream).load()
|
||||
|
||||
|
||||
def _decode(value: Any) -> str:
|
||||
if isinstance(value, bytes):
|
||||
return value.decode("utf-8", errors="replace")
|
||||
if isinstance(value, np.bytes_):
|
||||
return bytes(value).decode("utf-8", errors="replace")
|
||||
if isinstance(value, np.ndarray) and value.shape == ():
|
||||
return _decode(value.item())
|
||||
return str(value)
|
||||
|
||||
|
||||
def _one_dim(value: Any, dtype: Any | None = None) -> np.ndarray:
|
||||
out = np.asarray(value)
|
||||
if out.ndim > 1 and out.shape[-1] == 1:
|
||||
out = out.reshape(-1)
|
||||
elif out.ndim > 1 and out.shape[0] == 1:
|
||||
out = out.reshape(-1)
|
||||
else:
|
||||
out = out.reshape(-1)
|
||||
return out.astype(dtype) if dtype is not None else out
|
||||
|
||||
|
||||
@dataclass
|
||||
class SplitData:
|
||||
name: str
|
||||
x: dict[str, np.ndarray]
|
||||
mask: np.ndarray
|
||||
class_y: np.ndarray | None
|
||||
regression_y: np.ndarray | None
|
||||
ids: list[str]
|
||||
groups: np.ndarray
|
||||
alignment_audit: dict[str, Any] | None = None
|
||||
|
||||
@property
|
||||
def n(self) -> int:
|
||||
return len(self.ids)
|
||||
|
||||
|
||||
def _extract_split(
|
||||
name: str, obj: dict[str, Any], with_labels: bool,
|
||||
mask_override: np.ndarray | None = None,
|
||||
alignment_audit: dict[str, Any] | None = None,
|
||||
) -> SplitData:
|
||||
raw: dict[str, np.ndarray] = {}
|
||||
masks = []
|
||||
for modality in MODALITIES:
|
||||
arr = np.asarray(obj[modality])
|
||||
if arr.ndim != 3 or arr.shape[1] != 50 or arr.shape[2] != EXPECTED_DIMS[modality]:
|
||||
raise ValueError(f"{name}.{modality}: unexpected feature shape {arr.shape}")
|
||||
arr = arr.astype(np.float32)
|
||||
if not np.isfinite(arr).all():
|
||||
raise ValueError(f"{name}.{modality}: non-finite feature values; refusing to reinterpret them as missing")
|
||||
# The dataset documentation defines all-zero aligned rows as missing.
|
||||
observed = np.any(arr != 0.0, axis=-1)
|
||||
raw[modality] = arr
|
||||
masks.append(observed)
|
||||
mask = np.stack(masks, axis=-1)
|
||||
if mask_override is not None:
|
||||
override = np.asarray(mask_override, bool)
|
||||
if override.shape != mask.shape:
|
||||
raise ValueError(f"{name}: projected mask shape {override.shape} differs from {mask.shape}")
|
||||
mask = override
|
||||
ids = [_decode(v) for v in _one_dim(obj["id"])]
|
||||
if len(ids) != len(mask):
|
||||
raise ValueError(f"{name}: id count differs from feature count")
|
||||
if len(set(ids)) != len(ids):
|
||||
raise ValueError(f"{name}: duplicate video$_$clip primary keys")
|
||||
malformed = [sample_id for sample_id in ids if "$_$" not in sample_id or not all(sample_id.split("$_$", 1))]
|
||||
if malformed:
|
||||
raise ValueError(f"{name}: malformed video$_$clip keys: {malformed[:5]}")
|
||||
groups = np.asarray([sample_group(v) for v in ids], dtype=str)
|
||||
if with_labels:
|
||||
class_y = _one_dim(obj["classification_labels"], np.int64)
|
||||
regression_y = _one_dim(obj["regression_labels"], np.float32)
|
||||
if len(class_y) != len(ids) or len(regression_y) != len(ids):
|
||||
raise ValueError(f"{name}: label count differs from feature count")
|
||||
if not np.isfinite(regression_y).all() or np.any(np.abs(regression_y) > 3.0):
|
||||
raise ValueError(f"{name}: regression labels must be finite and within [-3,3]")
|
||||
if not np.isin(class_y, [0, 1, 2]).all():
|
||||
raise ValueError(f"{name}: expected class labels in 0,1,2")
|
||||
expected_class = np.where(regression_y < 0.0, 0, np.where(regression_y == 0.0, 1, 2))
|
||||
mismatch = np.flatnonzero(class_y != expected_class)
|
||||
if len(mismatch):
|
||||
examples = [(ids[int(i)], int(class_y[i]), float(regression_y[i])) for i in mismatch[:5]]
|
||||
raise ValueError(f"{name}: polarity/regression label mismatch (sample, class, score): {examples}")
|
||||
else:
|
||||
class_y = regression_y = None
|
||||
return SplitData(name, raw, mask, class_y, regression_y, ids, groups, alignment_audit)
|
||||
|
||||
|
||||
def sample_group(sample_id: str) -> str:
|
||||
"""Official ids are video$_$clip; group on the source video only."""
|
||||
return sample_id.split("$_$", 1)[0]
|
||||
|
||||
|
||||
def load_official_splits(path: Path = ALIGNED_PATH, *, version: str = "aligned_50") -> dict[str, SplitData]:
|
||||
if version not in {"aligned_50", "unaligned_50"}:
|
||||
raise ValueError(f"unsupported feature version: {version}")
|
||||
obj = restricted_load(path)
|
||||
required = {"train", "valid", "test"}
|
||||
if not isinstance(obj, dict) or not required.issubset(obj):
|
||||
raise ValueError(f"{path.name} must contain train, valid, and test dictionaries")
|
||||
if version == "unaligned_50":
|
||||
from ...adapter import adapt_official_split
|
||||
|
||||
splits = {}
|
||||
for name in ("train", "valid", "test"):
|
||||
projected, mask, audit = adapt_official_split(obj[name])
|
||||
fields = {**obj[name], **projected}
|
||||
splits[name] = _extract_split(name, fields, with_labels=True,
|
||||
mask_override=mask, alignment_audit=audit)
|
||||
else:
|
||||
splits = {name: _extract_split(name, obj[name], with_labels=True) for name in ("train", "valid", "test")}
|
||||
del obj
|
||||
return splits
|
||||
|
||||
|
||||
def load_attachment3_case(path: Path) -> dict[str, np.ndarray]:
|
||||
obj = restricted_load(path)
|
||||
case = obj.get("test", obj)
|
||||
text_bert = np.asarray(case["text_bert"])
|
||||
audio = np.asarray(case["audio"])
|
||||
vision = np.asarray(case["vision"])
|
||||
if text_bert.ndim == 3 and text_bert.shape[0] == 1:
|
||||
text_bert = text_bert[0]
|
||||
if text_bert.shape != (3, 50):
|
||||
raise ValueError(f"{path.name}: expected text_bert (1,3,50), got {np.asarray(case['text_bert']).shape}")
|
||||
result = {"input_ids": text_bert[0].astype(np.int64), "attention_mask": text_bert[1].astype(bool), "token_type_ids": text_bert[2].astype(np.int64)}
|
||||
for name, arr, dim in (("audio", audio, 74), ("vision", vision, 35)):
|
||||
if arr.ndim == 3 and arr.shape[0] == 1:
|
||||
arr = arr[0]
|
||||
if arr.shape != (50, dim):
|
||||
raise ValueError(f"{path.name}: expected {name} (1,50,{dim}), got {np.asarray(case[name]).shape}")
|
||||
arr = arr.astype(np.float32)
|
||||
if not np.isfinite(arr).all():
|
||||
raise ValueError(f"{path.name}: {name} contains non-finite features")
|
||||
result[name] = arr
|
||||
return result
|
||||
|
||||
|
||||
def fit_preprocessor(train: SplitData) -> dict[str, dict[str, np.ndarray]]:
|
||||
"""Fit per-dimension mean/std on observed training rows only."""
|
||||
fitted: dict[str, dict[str, np.ndarray]] = {}
|
||||
for j, name in enumerate(MODALITIES):
|
||||
rows = train.x[name][train.mask[:, :, j]]
|
||||
mean = rows.mean(axis=0, dtype=np.float64).astype(np.float32)
|
||||
std = rows.std(axis=0, dtype=np.float64).astype(np.float32)
|
||||
std[std < 1e-5] = 1.0
|
||||
fitted[name] = {"mean": mean, "std": std}
|
||||
return fitted
|
||||
|
||||
|
||||
def transform_split(split: SplitData, fitted: dict[str, dict[str, np.ndarray]]) -> dict[str, np.ndarray]:
|
||||
output = {}
|
||||
for j, name in enumerate(MODALITIES):
|
||||
arr = (split.x[name] - fitted[name]["mean"]) / fitted[name]["std"]
|
||||
arr = np.clip(arr, -10.0, 10.0)
|
||||
arr[~split.mask[:, :, j]] = 0.0
|
||||
output[name] = arr.astype(np.float32)
|
||||
return output
|
||||
@@ -0,0 +1,142 @@
|
||||
"""Create the method-comparison figures from Q2's fixed evaluation outputs."""
|
||||
from __future__ import annotations
|
||||
|
||||
import csv
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
import matplotlib.pyplot as plt
|
||||
import numpy as np
|
||||
|
||||
from .train import RESULTS
|
||||
|
||||
|
||||
def read_csv(path: Path) -> list[dict[str, str]]:
|
||||
with path.open("r", encoding="utf-8-sig", newline="") as stream:
|
||||
return list(csv.DictReader(stream))
|
||||
|
||||
|
||||
def main() -> None:
|
||||
controlled = read_csv(RESULTS / "controlled_missingness.csv")
|
||||
test = read_csv(RESULTS / "test_predictions.csv")
|
||||
gates = read_csv(RESULTS / "test_gate_diagnostics.csv")
|
||||
metrics = json.loads((RESULTS / "test_metrics.json").read_text(encoding="utf-8"))
|
||||
selected = str(metrics["selected_model"])
|
||||
|
||||
figure, axes = plt.subplots(2, 3, figsize=(17, 10), constrained_layout=True)
|
||||
rate_axis, modality_axis, location_axis, confusion_axis, scatter_axis, interval_axis = axes.flat
|
||||
|
||||
comparison_models = ("C0", "C3", "C4", "C5", "C6", "C7_distill", "C7_group")
|
||||
for model in comparison_models:
|
||||
subset = [row for row in controlled if row["model"] == model and
|
||||
(row["mask_pattern"] == "none" or row["mask_pattern"] == "single")]
|
||||
subset.sort(key=lambda row: float(row["rate_realized_additional_global"]))
|
||||
if subset:
|
||||
rate_axis.plot([float(row["rate_realized_additional_global"]) for row in subset],
|
||||
[float(row["regression_mae"]) for row in subset], marker="o", label=model)
|
||||
rate_axis.set(title="MAE by realized additional missing rate", xlabel="Additional missing rate (equal T/A/V)", ylabel="MAE")
|
||||
rate_axis.legend(fontsize=8, ncol=2)
|
||||
rate_axis.grid(alpha=0.25)
|
||||
|
||||
modality_labels = ("T", "A", "V", "TA", "TV", "AV", "TAV")
|
||||
modality_scenarios = {f"0.3/modality_{label}": label for label in modality_labels}
|
||||
modality_models = ("C0", "C5", "C6", "C7_distill", "C7_group")
|
||||
modality_values = np.full((len(modality_models), len(modality_labels)), np.nan)
|
||||
for i, model in enumerate(modality_models):
|
||||
for j, (scenario, _) in enumerate(modality_scenarios.items()):
|
||||
pattern = scenario.split("/", 1)[1]
|
||||
row = next((r for r in controlled if r["model"] == model and r["mask_pattern"] == pattern), None)
|
||||
if row is not None:
|
||||
modality_values[i, j] = float(row["regression_mae"])
|
||||
image = modality_axis.imshow(modality_values, aspect="auto", cmap="viridis")
|
||||
modality_axis.set(title="Modality combination control: MAE", xticks=range(len(modality_labels)),
|
||||
xticklabels=modality_labels, yticks=range(len(modality_models)), yticklabels=modality_models)
|
||||
modality_axis.tick_params(axis="x", rotation=35)
|
||||
figure.colorbar(image, ax=modality_axis, fraction=0.046, pad=0.04)
|
||||
|
||||
position_values = np.full((3, 3), np.nan)
|
||||
for i, modality in enumerate(("T", "A", "V")):
|
||||
for j, location in enumerate(("start", "middle", "end")):
|
||||
scenario = f"0.3/location_{location}_{modality}"
|
||||
pattern = scenario.split("/", 1)[1]
|
||||
row = next((r for r in controlled if r["model"] == selected and r["mask_pattern"] == pattern), None)
|
||||
if row is not None:
|
||||
position_values[i, j] = float(row["regression_mae"])
|
||||
image = location_axis.imshow(position_values, aspect="auto", cmap="magma")
|
||||
location_axis.set(title=f"Selected model {selected}: location MAE", xticks=range(3),
|
||||
xticklabels=("start", "middle", "end"), yticks=range(3), yticklabels=("T", "A", "V"))
|
||||
figure.colorbar(image, ax=location_axis, fraction=0.046, pad=0.04)
|
||||
|
||||
confusion = np.zeros((3, 3), dtype=np.int64)
|
||||
for row in test:
|
||||
confusion[int(row["true_class"]), int(row["predicted_class"])] += 1
|
||||
image = confusion_axis.imshow(confusion, cmap="Blues")
|
||||
for i in range(3):
|
||||
for j in range(3):
|
||||
confusion_axis.text(j, i, str(confusion[i, j]), ha="center", va="center")
|
||||
confusion_axis.set(title=f"Test confusion matrix: {selected}", xlabel="Predicted", ylabel="True",
|
||||
xticks=range(3), xticklabels=("negative", "neutral", "positive"),
|
||||
yticks=range(3), yticklabels=("negative", "neutral", "positive"))
|
||||
figure.colorbar(image, ax=confusion_axis, fraction=0.046, pad=0.04)
|
||||
|
||||
true_score = np.asarray([float(row["true_sentiment"]) for row in test])
|
||||
predicted_score = np.asarray([float(row["predicted_sentiment"]) for row in test])
|
||||
scatter_axis.scatter(true_score, predicted_score, alpha=0.55, s=18)
|
||||
scatter_axis.plot([-3, 3], [-3, 3], "k--", linewidth=1)
|
||||
scatter_axis.set(title=f"Test sentiment: MAE={metrics['regression_mae']:.3f}",
|
||||
xlabel="True sentiment", ylabel="Predicted sentiment", xlim=(-3, 3), ylim=(-3, 3))
|
||||
scatter_axis.grid(alpha=0.2)
|
||||
|
||||
lower = np.asarray([float(row["interval_90_lower"]) for row in test])
|
||||
upper = np.asarray([float(row["interval_90_upper"]) for row in test])
|
||||
width = upper - lower
|
||||
covered = (true_score >= lower) & (true_score <= upper)
|
||||
order = np.argsort(width)
|
||||
bins = np.array_split(order, min(10, len(order)))
|
||||
interval_axis.plot([width[idx].mean() for idx in bins], [covered[idx].mean() for idx in bins], marker="o")
|
||||
interval_axis.axhline(0.9, color="black", linestyle="--", linewidth=1, label="nominal 90%")
|
||||
interval_axis.set(title="Test interval coverage by width decile", xlabel="Mean interval width", ylabel="Empirical coverage", ylim=(0, 1))
|
||||
interval_axis.legend()
|
||||
interval_axis.grid(alpha=0.2)
|
||||
figure.suptitle("Q2 validation controls and official-test diagnostics", fontsize=15)
|
||||
figure.savefig(RESULTS / "q2_diagnostics.png", dpi=160)
|
||||
plt.close(figure)
|
||||
|
||||
steps = 50
|
||||
modalities = ("text", "audio", "vision")
|
||||
weight_sum = np.zeros((steps, len(modalities)), dtype=np.float64)
|
||||
reliability_sum = np.zeros_like(weight_sum)
|
||||
count = np.zeros_like(weight_sum)
|
||||
for row in gates:
|
||||
if row["modality"] not in modalities:
|
||||
continue
|
||||
t, m = int(row["step"]), modalities.index(row["modality"])
|
||||
weight_sum[t, m] += float(row["fusion_weight_mean_over_paths"])
|
||||
reliability_sum[t, m] += float(row["reliability"])
|
||||
count[t, m] += 1
|
||||
weights = weight_sum / np.maximum(count, 1.0)
|
||||
reliabilities = reliability_sum / np.maximum(count, 1.0)
|
||||
gate_figure, gate_axis = plt.subplots(figsize=(12, 5), constrained_layout=True)
|
||||
for m, modality in enumerate(modalities):
|
||||
gate_axis.plot(range(steps), weights[:, m], label=f"{modality} fusion weight")
|
||||
gate_axis.set(title=f"Test mean fusion gates by position: {selected}", xlabel="Aligned step", ylabel="Mean fusion weight")
|
||||
gate_axis.legend(ncol=3)
|
||||
gate_axis.grid(alpha=0.25)
|
||||
reliability_axis = gate_axis.twinx()
|
||||
for m, modality in enumerate(modalities):
|
||||
reliability_axis.plot(range(steps), reliabilities[:, m], linestyle=":", alpha=0.7, label=f"{modality} reliability")
|
||||
reliability_axis.set_ylabel("Mean reliability proxy")
|
||||
handles, labels = gate_axis.get_legend_handles_labels()
|
||||
right_handles, right_labels = reliability_axis.get_legend_handles_labels()
|
||||
gate_axis.legend(handles + right_handles, labels + right_labels, ncol=3, fontsize=8)
|
||||
gate_figure.savefig(RESULTS / "q2_gate_positions.png", dpi=160)
|
||||
plt.close(gate_figure)
|
||||
manifest_path = RESULTS / "run_manifest.json"
|
||||
manifest = json.loads(manifest_path.read_text(encoding="utf-8"))
|
||||
manifest["diagnostic_figures"] = ["q2_diagnostics.png", "q2_gate_positions.png"]
|
||||
manifest_path.write_text(json.dumps(manifest, ensure_ascii=False, indent=2), encoding="utf-8")
|
||||
print(f"Wrote Q2 figures for selected model {selected} to {RESULTS}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,90 @@
|
||||
"""Plot the aligned-data rate sweep and matched missing-type response."""
|
||||
from __future__ import annotations
|
||||
|
||||
import csv
|
||||
from pathlib import Path
|
||||
|
||||
import matplotlib.pyplot as plt
|
||||
import numpy as np
|
||||
|
||||
|
||||
RESULTS = Path(__file__).resolve().parents[2] / "experiments" / "q2" / "math_current"
|
||||
|
||||
|
||||
def read_csv(name: str) -> list[dict[str, str]]:
|
||||
with (RESULTS / name).open(encoding="utf-8-sig", newline="") as stream:
|
||||
return list(csv.DictReader(stream))
|
||||
|
||||
|
||||
def main() -> None:
|
||||
rate_rows = read_csv("controlled_missingness.csv")
|
||||
rate_bootstrap = read_csv("controlled_group_bootstrap.csv")
|
||||
type_bootstrap = read_csv("matched_missing_type_bootstrap.csv")
|
||||
|
||||
colors = {"single": "#3b82f6", "sync": "#dc2626", "partial": "#16a34a", "async": "#9333ea"}
|
||||
labels = {"single": "Single modality", "sync": "Synchronous", "partial": "Partial overlap", "async": "Asynchronous"}
|
||||
fig, (ax_rate, ax_type) = plt.subplots(1, 2, figsize=(12.4, 4.8), gridspec_kw={"width_ratios": [1.35, 1.0]})
|
||||
|
||||
baseline = next(row for row in rate_rows if row["model"] == "C5" and row["mask_pattern"] == "none")
|
||||
baseline_ci = next(row for row in rate_bootstrap if row["model"] == "C5" and row["scenario"] == "0.0/none" and row["metric"] == "mae")
|
||||
for mode in ("single", "sync", "partial", "async"):
|
||||
rows = [baseline] + sorted(
|
||||
(row for row in rate_rows if row["model"] == "C5" and row["mask_pattern"] == mode),
|
||||
key=lambda row: float(row["rate_requested_per_selected_source"]),
|
||||
)
|
||||
x, y, lower, upper = [], [], [], []
|
||||
for row in rows:
|
||||
if row["mask_pattern"] == "none":
|
||||
ci = baseline_ci
|
||||
scenario = "0.0/none"
|
||||
else:
|
||||
scenario = f"{float(row['rate_requested_per_selected_source']):.1f}/{mode}"
|
||||
ci = next(item for item in rate_bootstrap if item["model"] == "C5" and item["scenario"] == scenario and item["metric"] == "mae")
|
||||
x.append(float(row["rate_realized_additional_global"]))
|
||||
y.append(float(row["regression_mae"]))
|
||||
lower.append(float(ci["ci_2_5"]))
|
||||
upper.append(float(ci["ci_97_5"]))
|
||||
ax_rate.errorbar(
|
||||
x, y, yerr=[np.asarray(y) - np.asarray(lower), np.asarray(upper) - np.asarray(y)],
|
||||
color=colors[mode], marker="o", linewidth=1.7, markersize=4.5,
|
||||
capsize=2.5, label=labels[mode], alpha=0.95,
|
||||
)
|
||||
ax_rate.set_title("C5 performance across missing rates")
|
||||
ax_rate.set_xlabel("Added missing rate (paper definition)")
|
||||
ax_rate.set_ylabel("Regression MAE (95% group-bootstrap CI)")
|
||||
ax_rate.grid(axis="both", color="#d1d5db", linewidth=0.7, alpha=0.65)
|
||||
ax_rate.legend(frameon=False, fontsize=8.5, loc="upper left")
|
||||
|
||||
type_order = ("T", "A", "V", "TA", "TV", "AV", "TAV")
|
||||
point, low, high = [], [], []
|
||||
for label in type_order:
|
||||
scenario = f"matched_type_{label}"
|
||||
boot = next(row for row in type_bootstrap if row["model"] == "C5" and row["scenario"] == scenario and row["metric"] == "mae")
|
||||
point.append(float(boot["delta_to_natural"]))
|
||||
low.append(float(boot["delta_to_natural_ci_2_5"]))
|
||||
high.append(float(boot["delta_to_natural_ci_97_5"]))
|
||||
positions = np.arange(len(type_order))
|
||||
ax_type.errorbar(
|
||||
positions, point, yerr=[np.asarray(point) - low, high - np.asarray(point)],
|
||||
fmt="o", color="#2563eb", ecolor="#2563eb", capsize=3, linewidth=1.4,
|
||||
markersize=5,
|
||||
)
|
||||
ax_type.axhline(0, color="#374151", linewidth=1, linestyle="--")
|
||||
ax_type.set_xticks(positions, type_order)
|
||||
ax_type.set_title("Matched missing-modality types")
|
||||
ax_type.set_xlabel("Hidden modality set")
|
||||
ax_type.set_ylabel("MAE change from natural condition")
|
||||
ax_type.grid(axis="y", color="#d1d5db", linewidth=0.7, alpha=0.65)
|
||||
ax_type.text(
|
||||
0.02, 0.02, "Same added feature-row count per sample and type",
|
||||
transform=ax_type.transAxes, fontsize=7.5, color="#4b5563",
|
||||
)
|
||||
|
||||
fig.tight_layout(pad=1.2)
|
||||
output = RESULTS / "aligned_missingness_effects.png"
|
||||
fig.savefig(output, dpi=200, bbox_inches="tight", facecolor="white")
|
||||
print(output)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,103 @@
|
||||
"""Run the saved Q2 student on the aligned, unlabeled attachment-3 cases."""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import time
|
||||
import csv
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from ...model.crg import INPUT_DIMS, MODALITIES, StructuredGaussianImputer
|
||||
from .train import RESULTS, _make_variant, infer_attachment3, reencode_attachment3, validate_attachment3_predictions, write_csv
|
||||
|
||||
|
||||
def main() -> None:
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
manifest_path = RESULTS / "run_manifest.json"
|
||||
manifest = json.loads(manifest_path.read_text(encoding="utf-8"))
|
||||
calibration = json.loads((RESULTS / "validation_metrics.json").read_text(encoding="utf-8"))
|
||||
selected = calibration.get("selected_model", manifest.get("selected_model"))
|
||||
if not selected:
|
||||
raise ValueError("run_manifest.json does not identify a selected model")
|
||||
|
||||
imputer = StructuredGaussianImputer(INPUT_DIMS).to(device)
|
||||
imputer_state = torch.load(RESULTS / "structured_imputer.pt", map_location=device, weights_only=True)
|
||||
imputer.load_state_dict(imputer_state)
|
||||
model = _make_variant(selected, imputer).to(device)
|
||||
state = torch.load(RESULTS / "crg_student.pt", map_location=device, weights_only=True)
|
||||
model.load_state_dict(state)
|
||||
|
||||
with np.load(RESULTS / "preprocessor.npz", allow_pickle=False) as archive:
|
||||
fitted = {m: {k: archive[f"{m}_{k}"].copy() for k in ("mean", "std")} for m in MODALITIES}
|
||||
priors = manifest["attachment3_low_information_priors"]
|
||||
temperature = float(calibration["temperature"])
|
||||
class_prior = np.asarray(priors["class_probability_values"], dtype=np.float64)
|
||||
magnitude_priors = np.asarray((priors["negative_beta"], priors["positive_beta"]), dtype=np.float32)
|
||||
|
||||
cases, source_audit = reencode_attachment3(device)
|
||||
predictions, inference_audit = infer_attachment3(
|
||||
model, cases, fitted, device, temperature, class_prior, magnitude_priors,
|
||||
)
|
||||
validate_attachment3_predictions([case["case_id"] for case in cases], predictions)
|
||||
inference_by_id = {row["case_id"]: row for row in inference_audit}
|
||||
write_csv(RESULTS / "attachment3_predictions.csv", predictions)
|
||||
write_csv(RESULTS / "attachment3_audit.csv", [
|
||||
{**source, **inference_by_id[source["case_id"]]} for source in source_audit
|
||||
])
|
||||
|
||||
# The training script can finish and persist all labeled-evaluation outputs
|
||||
# before an unlabeled attachment export fails. Reconcile the manifest from
|
||||
# those completed artifacts so the standalone export is safely rerunnable.
|
||||
group_risk_rows = list(csv.DictReader((RESULTS / "group_risk_tuning.csv").open(encoding="utf-8-sig", newline="")))
|
||||
selected_risk = next((row for row in group_risk_rows if row.get("selected", "").lower() == "true"), None)
|
||||
reliability_rows = list(csv.DictReader((RESULTS / "reliability_hparam_tuning.csv").open(encoding="utf-8-sig", newline="")))
|
||||
# split_calibration's generic internal names are canonicalized in train.py;
|
||||
# repair artifacts from runs produced before that naming fix as well.
|
||||
for row in group_risk_rows:
|
||||
if row.get("selection_split") == "fit":
|
||||
row["selection_split"] = "reliability_validation"
|
||||
for row in reliability_rows:
|
||||
if row.get("selection_split") == "fit":
|
||||
row["selection_split"] = "reliability_validation"
|
||||
write_csv(RESULTS / "group_risk_tuning.csv", group_risk_rows)
|
||||
write_csv(RESULTS / "reliability_hparam_tuning.csv", reliability_rows)
|
||||
if selected_risk:
|
||||
risk_values = (float(selected_risk["lambda_group"]), float(selected_risk["group_temperature"]))
|
||||
manifest["group_risk_hyperparameters"]["selected"] = list(risk_values)
|
||||
manifest["loss"]["selected_group_risk"] = list(risk_values)
|
||||
manifest["group_risk_hyperparameters"]["selection_split"] = "reliability_validation"
|
||||
manifest["reliability_hyperparameters"]["selected_by_model"] = {
|
||||
row["model"]: [float(row[key]) for key in ("rho_imp", "lambda_u", "lambda_gap", "lambda_span")]
|
||||
for row in reliability_rows
|
||||
if row.get("selected", "").lower() == "true"
|
||||
and (not row.get("risk_candidate_selected") or row["risk_candidate_selected"].lower() == "true")
|
||||
}
|
||||
test_metrics = json.loads((RESULTS / "test_metrics.json").read_text(encoding="utf-8"))
|
||||
manifest["selected_model"] = selected
|
||||
manifest["final_test_metrics"] = test_metrics
|
||||
manifest["calibration"]["temperature"] = temperature
|
||||
manifest["calibration"]["valid_used_for_selection"] = True
|
||||
manifest["calibration"]["test_used_for_selection_or_calibration"] = False
|
||||
manifest["training_configuration"].update({
|
||||
"student_epoch_limit": 12,
|
||||
"imputer_epochs": 8,
|
||||
"batch_size": 64,
|
||||
"early_stopping_patience": 3,
|
||||
})
|
||||
manifest["imputer"]["epochs"] = 8
|
||||
manifest.update({
|
||||
"completed_utc": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()),
|
||||
"attachment3_cases": len(cases),
|
||||
"attachment3_prediction_file": "attachment3_predictions.csv",
|
||||
"attachment3_audit_file": "attachment3_audit.csv",
|
||||
"attachment3_labeled_metrics": None,
|
||||
"quality_flags": {m: "unavailable; q*=1 fallback for visible rows, unknown flag retained" for m in MODALITIES},
|
||||
"neutral_output": "exact zero when neutral is the predicted class; no near-zero threshold",
|
||||
})
|
||||
manifest_path.write_text(json.dumps(manifest, ensure_ascii=False, indent=2), encoding="utf-8")
|
||||
print(f"Wrote {len(predictions)} unlabeled attachment-3 predictions to {RESULTS}", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,369 @@
|
||||
"""Matched-volume missing-modality evaluation for the official validation split.
|
||||
|
||||
This complements the standard requested-rate sweep. Every type condition hides
|
||||
the same number of originally observed feature rows in each validation sample;
|
||||
the affected rows are placed in one contiguous span per selected modality.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import hashlib
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from sklearn.metrics import accuracy_score, f1_score, mean_absolute_error
|
||||
|
||||
from . import train
|
||||
from ...model.crg import INPUT_DIMS, StructuredGaussianImputer
|
||||
from .data import ALIGNED_PATH, fit_preprocessor, load_official_splits, transform_split
|
||||
|
||||
|
||||
RESULTS = Path(__file__).resolve().parent / "results"
|
||||
MODALITY_SETS = (
|
||||
((0,), "T"), ((1,), "A"), ((2,), "V"),
|
||||
((0, 1), "TA"), ((0, 2), "TV"), ((1, 2), "AV"), ((0, 1, 2), "TAV"),
|
||||
)
|
||||
MODALITY_NAMES = ("text", "audio", "vision")
|
||||
|
||||
|
||||
def make_matched_type_masks(
|
||||
split: Any, seed: int, per_sample_cap: int = 15,
|
||||
) -> tuple[dict[str, np.ndarray], list[dict[str, Any]]]:
|
||||
original = np.asarray(split.mask, dtype=bool)
|
||||
counts = original.sum(axis=1).astype(np.int64)
|
||||
keep_minimum = np.maximum(1, np.ceil(0.2 * counts).astype(np.int64))
|
||||
capacity = np.maximum(0, counts - keep_minimum)
|
||||
# Match the same feasible volume per sample for every modality set. The
|
||||
# least observed of audio/vision determines the cap, so no condition can
|
||||
# gain an advantage by applying its mask to a different subset of samples.
|
||||
budget = np.minimum(per_sample_cap, np.minimum(capacity[:, 1], capacity[:, 2]))
|
||||
masks: dict[str, np.ndarray] = {"0.0/none": original.copy()}
|
||||
audit: list[dict[str, Any]] = []
|
||||
|
||||
for selected, label in MODALITY_SETS:
|
||||
key = f"matched_type_{label}"
|
||||
current = original.copy()
|
||||
for row_index, sample_id in enumerate(split.ids):
|
||||
total = int(budget[row_index])
|
||||
base, remainder = divmod(total, len(selected))
|
||||
sample_seed = int.from_bytes(
|
||||
hashlib.sha256(f"{seed}:{sample_id}:{key}".encode("utf-8")).digest()[:8],
|
||||
"little",
|
||||
)
|
||||
rng = np.random.default_rng(sample_seed)
|
||||
allocation = np.full(len(selected), base, dtype=np.int64)
|
||||
if remainder:
|
||||
allocation[rng.permutation(len(selected))[:remainder]] += 1
|
||||
starts: list[str] = []
|
||||
ends: list[str] = []
|
||||
hidden_by_modality = np.zeros(3, dtype=np.int64)
|
||||
for modality, amount_value in zip(selected, allocation):
|
||||
amount = int(amount_value)
|
||||
if amount == 0:
|
||||
starts.append("")
|
||||
ends.append("")
|
||||
continue
|
||||
interval = train._best_interval(
|
||||
original[row_index, :, modality], amount,
|
||||
int(capacity[row_index, modality]), "random", rng,
|
||||
)
|
||||
if interval is None:
|
||||
raise RuntimeError(f"no feasible interval for {sample_id}/{label}/{MODALITY_NAMES[modality]}")
|
||||
left, right = interval
|
||||
positions = np.flatnonzero(original[row_index, left:right + 1, modality]) + left
|
||||
if len(positions) != amount:
|
||||
raise RuntimeError(f"matched interval hid {len(positions)} rows, expected {amount}")
|
||||
current[row_index, positions, modality] = False
|
||||
hidden_by_modality[modality] = amount
|
||||
starts.append(str(int(left)))
|
||||
ends.append(str(int(right)))
|
||||
actual_total = int(np.sum(original[row_index] & ~current[row_index]))
|
||||
if actual_total != total:
|
||||
raise RuntimeError(f"matched volume differs for {sample_id}: {actual_total} != {total}")
|
||||
audit.append({
|
||||
"scenario": key,
|
||||
"sample_id": sample_id,
|
||||
"source_video_id": str(split.groups[row_index]),
|
||||
"selected_modalities": json.dumps([MODALITY_NAMES[m] for m in selected]),
|
||||
"base_mask_seed": int(seed),
|
||||
"sample_mask_seed": sample_seed,
|
||||
"matched_added_rows_target": total,
|
||||
"matched_added_rows_actual": actual_total,
|
||||
"hidden_text_rows": int(hidden_by_modality[0]),
|
||||
"hidden_audio_rows": int(hidden_by_modality[1]),
|
||||
"hidden_vision_rows": int(hidden_by_modality[2]),
|
||||
"span_start_by_selected_modality": json.dumps(starts),
|
||||
"span_end_by_selected_modality": json.dumps(ends),
|
||||
})
|
||||
if not np.array_equal(np.sum(original & ~current, axis=(1, 2)), budget):
|
||||
raise RuntimeError(f"per-sample matched-volume invariant failed for {label}")
|
||||
masks[key] = current
|
||||
|
||||
total_masked = int(budget.sum())
|
||||
if len({int(np.sum(original & ~mask)) for key, mask in masks.items() if key != "0.0/none"}) != 1:
|
||||
raise RuntimeError("matched modality scenarios do not have identical total missing volume")
|
||||
print(
|
||||
f"matched type masks: samples={split.n}, added_rows_per_scenario={total_masked}, "
|
||||
f"mean_per_sample={budget.mean():.3f}, zero_budget_samples={int(np.sum(budget == 0))}",
|
||||
flush=True,
|
||||
)
|
||||
return masks, audit
|
||||
|
||||
|
||||
def metric_values(split: Any, prediction: dict[str, np.ndarray], indices: np.ndarray) -> dict[str, float]:
|
||||
return {
|
||||
"accuracy": float(accuracy_score(split.class_y[indices], prediction["predicted_class"][indices])),
|
||||
"macro_f1": float(f1_score(
|
||||
split.class_y[indices], prediction["predicted_class"][indices],
|
||||
labels=[0, 1, 2], average="macro", zero_division=0,
|
||||
)),
|
||||
"mae": float(mean_absolute_error(
|
||||
split.regression_y[indices], prediction["predicted_score"][indices],
|
||||
)),
|
||||
}
|
||||
|
||||
|
||||
def evaluate_matched_masks(
|
||||
model_name: str,
|
||||
split: Any,
|
||||
arrays: dict[str, np.ndarray],
|
||||
masks: dict[str, np.ndarray],
|
||||
temperature: float,
|
||||
*,
|
||||
model: Any | None = None,
|
||||
c0_state: dict[str, Any] | None = None,
|
||||
device: torch.device | None = None,
|
||||
seed: int = 0,
|
||||
) -> tuple[list[dict[str, Any]], dict[str, dict[str, np.ndarray]]]:
|
||||
rows = []
|
||||
predictions = {}
|
||||
for scenario, mask in masks.items():
|
||||
if model_name == "C0":
|
||||
if c0_state is None:
|
||||
raise ValueError("C0 state is required")
|
||||
metrics, prediction = train.evaluate_c0(c0_state, split, arrays, temperature, mask)
|
||||
else:
|
||||
if model is None or device is None:
|
||||
raise ValueError("neural model and device are required")
|
||||
scenario_seed = train._scenario_seed(seed, split.name, scenario)
|
||||
with train.fixed_torch_seed(scenario_seed, device):
|
||||
metrics, prediction = train.evaluate(
|
||||
model, arrays, split, device, 64, masks=mask, temperature=temperature,
|
||||
)
|
||||
predictions[scenario] = prediction
|
||||
rates = train._missing_rate_summary(split.mask, mask)
|
||||
row = {
|
||||
"model": model_name,
|
||||
"scenario": scenario,
|
||||
"rate_realized_additional_global": rates["additional_global"],
|
||||
"rate_realized_additional_by_modality": json.dumps(
|
||||
[None if not np.isfinite(value) else float(value)
|
||||
for value in np.nanmean(rates["additional_by_modality"], axis=0)]
|
||||
),
|
||||
"natural_missing_rate_global": rates["natural_global"],
|
||||
"natural_missing_rate_by_modality": json.dumps(
|
||||
np.mean(rates["natural_by_modality"], axis=0).tolist()
|
||||
),
|
||||
"rate_final_total_missing_global": rates["final_global"],
|
||||
"rate_final_total_missing_by_modality": json.dumps(
|
||||
np.mean(rates["final_by_modality"], axis=0).tolist()
|
||||
),
|
||||
"synchronous_no_observation_rate": float(np.mean(rates["synchronous_no_observation"])),
|
||||
"matched_added_rows_total": int(np.sum(split.mask & ~mask)),
|
||||
**metrics,
|
||||
}
|
||||
rows.append(row)
|
||||
return rows, predictions
|
||||
|
||||
|
||||
def paired_source_video_bootstrap(
|
||||
split: Any,
|
||||
predictions: dict[str, dict[str, dict[str, np.ndarray]]],
|
||||
repeats: int,
|
||||
seed: int,
|
||||
) -> list[dict[str, Any]]:
|
||||
scenarios = list(predictions["C0"])
|
||||
groups = np.unique(split.groups)
|
||||
group_indices = {group: np.flatnonzero(split.groups == group) for group in groups}
|
||||
names = tuple(predictions)
|
||||
point = {
|
||||
(model, scenario, metric): value
|
||||
for model in names
|
||||
for scenario in scenarios
|
||||
for metric, value in metric_values(split, predictions[model][scenario], np.arange(split.n)).items()
|
||||
}
|
||||
draws = {key: [] for key in point}
|
||||
within_natural = {
|
||||
(model, scenario, metric): []
|
||||
for model in names for scenario in scenarios if scenario != "0.0/none"
|
||||
for metric in ("accuracy", "macro_f1", "mae")
|
||||
}
|
||||
model_deltas = {
|
||||
(scenario, metric): []
|
||||
for scenario in scenarios for metric in ("accuracy", "macro_f1", "mae")
|
||||
}
|
||||
rng = np.random.default_rng(seed)
|
||||
for _ in range(repeats):
|
||||
chosen = rng.choice(groups, size=len(groups), replace=True)
|
||||
indices = np.concatenate([group_indices[group] for group in chosen])
|
||||
replicate = {}
|
||||
for model in names:
|
||||
for scenario in scenarios:
|
||||
for metric, value in metric_values(split, predictions[model][scenario], indices).items():
|
||||
replicate[(model, scenario, metric)] = value
|
||||
draws[(model, scenario, metric)].append(value)
|
||||
for model in names:
|
||||
for scenario in scenarios:
|
||||
if scenario == "0.0/none":
|
||||
continue
|
||||
for metric in ("accuracy", "macro_f1", "mae"):
|
||||
within_natural[(model, scenario, metric)].append(
|
||||
replicate[(model, scenario, metric)] - replicate[(model, "0.0/none", metric)]
|
||||
)
|
||||
for scenario in scenarios:
|
||||
for metric in ("accuracy", "macro_f1", "mae"):
|
||||
model_deltas[(scenario, metric)].append(
|
||||
replicate[("C5", scenario, metric)] - replicate[("C0", scenario, metric)]
|
||||
)
|
||||
|
||||
def interval(values: list[float]) -> tuple[float, float, float]:
|
||||
values_np = np.asarray(values, dtype=np.float64)
|
||||
return (float(np.median(values_np)), float(np.percentile(values_np, 2.5)),
|
||||
float(np.percentile(values_np, 97.5)))
|
||||
|
||||
rows = []
|
||||
for model in names:
|
||||
for scenario in scenarios:
|
||||
for metric in ("accuracy", "macro_f1", "mae"):
|
||||
median, lower, upper = interval(draws[(model, scenario, metric)])
|
||||
row: dict[str, Any] = {
|
||||
"model": model, "scenario": scenario, "metric": metric,
|
||||
"estimate": point[(model, scenario, metric)],
|
||||
"bootstrap_median": median, "ci_2_5": lower, "ci_97_5": upper,
|
||||
"replicates": repeats, "unit": "paired source-video group resample",
|
||||
}
|
||||
if scenario != "0.0/none":
|
||||
delta = point[(model, scenario, metric)] - point[(model, "0.0/none", metric)]
|
||||
d_median, d_lower, d_upper = interval(within_natural[(model, scenario, metric)])
|
||||
row.update({
|
||||
"delta_to_natural": delta,
|
||||
"delta_to_natural_bootstrap_median": d_median,
|
||||
"delta_to_natural_ci_2_5": d_lower,
|
||||
"delta_to_natural_ci_97_5": d_upper,
|
||||
})
|
||||
model_delta = point[("C5", scenario, metric)] - point[("C0", scenario, metric)]
|
||||
md_median, md_lower, md_upper = interval(model_deltas[(scenario, metric)])
|
||||
row.update({
|
||||
"C5_minus_C0": model_delta,
|
||||
"C5_minus_C0_bootstrap_median": md_median,
|
||||
"C5_minus_C0_ci_2_5": md_lower,
|
||||
"C5_minus_C0_ci_97_5": md_upper,
|
||||
})
|
||||
rows.append(row)
|
||||
return rows
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--seed", type=int, default=20260924 + 1209)
|
||||
parser.add_argument("--per-sample-cap", type=int, default=15)
|
||||
parser.add_argument("--bootstrap-repeats", type=int, default=1000)
|
||||
parser.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
|
||||
args = parser.parse_args()
|
||||
train.seed_everything(args.seed)
|
||||
device = torch.device(args.device)
|
||||
|
||||
official = load_official_splits()
|
||||
fit, heldout_train = train.split_calibration(official["train"], 20260924)
|
||||
_, temperature_calibration = train.split_calibration(heldout_train, 20260925, fraction=0.5)
|
||||
fitted = fit_preprocessor(fit)
|
||||
transformed = {name: transform_split(split, fitted) for name, split in official.items()}
|
||||
transformed["fit"] = transform_split(fit, fitted)
|
||||
transformed["temperature_calibration"] = transform_split(temperature_calibration, fitted)
|
||||
|
||||
imputer = StructuredGaussianImputer(INPUT_DIMS)
|
||||
imputer.load_state_dict(torch.load(RESULTS / "structured_imputer.pt", map_location="cpu", weights_only=True))
|
||||
with (RESULTS / "validation_metrics.json").open(encoding="utf-8") as stream:
|
||||
validation_metadata = json.load(stream)
|
||||
if validation_metadata.get("selected_model") != "C5":
|
||||
raise RuntimeError(f"expected selected C5 model, found {validation_metadata.get('selected_model')}")
|
||||
temperature = float(validation_metadata["temperature"])
|
||||
selected_reliability = (0.3, 0.05, 0.05, 0.05)
|
||||
model = train._make_variant("C5", imputer, selected_reliability).to(device)
|
||||
model.load_state_dict(torch.load(RESULTS / "crg_student.pt", map_location="cpu", weights_only=True))
|
||||
model.eval()
|
||||
|
||||
_, c0_state = train.fit_c0(fit, official["valid"], transformed)
|
||||
train.calibrate_c0_interval(c0_state, temperature_calibration, transformed["temperature_calibration"])
|
||||
_, c0_calibration = train.evaluate_c0(c0_state, temperature_calibration, transformed["temperature_calibration"])
|
||||
c0_temperature = train.fit_temperature(c0_calibration["probabilities"], temperature_calibration.class_y)
|
||||
|
||||
masks, audit = make_matched_type_masks(official["valid"], args.seed, args.per_sample_cap)
|
||||
metrics_c0, pred_c0 = evaluate_matched_masks(
|
||||
"C0", official["valid"], transformed["valid"], masks, c0_temperature,
|
||||
c0_state=c0_state,
|
||||
)
|
||||
metrics_c5, pred_c5 = evaluate_matched_masks(
|
||||
"C5", official["valid"], transformed["valid"], masks, temperature,
|
||||
model=model, device=device, seed=args.seed + 1,
|
||||
)
|
||||
total_masked = int(np.sum(official["valid"].mask & ~masks["matched_type_T"]))
|
||||
per_sample_budget_mean = total_masked / official["valid"].n
|
||||
zero_budget_samples = sum(
|
||||
1 for row in audit if row["scenario"] == "matched_type_T" and row["matched_added_rows_target"] == 0
|
||||
)
|
||||
summary_rows = []
|
||||
for row in metrics_c0 + metrics_c5:
|
||||
row["matched_added_rows_mean_per_sample"] = per_sample_budget_mean
|
||||
row["matched_zero_budget_samples"] = zero_budget_samples
|
||||
summary_rows.append(row)
|
||||
train.write_csv(RESULTS / "matched_missing_type.csv", summary_rows)
|
||||
train.write_csv(RESULTS / "matched_missing_type_audit.csv", audit)
|
||||
bootstrap_rows = paired_source_video_bootstrap(
|
||||
official["valid"], {"C0": pred_c0, "C5": pred_c5},
|
||||
args.bootstrap_repeats, args.seed + 2,
|
||||
)
|
||||
train.write_csv(RESULTS / "matched_missing_type_bootstrap.csv", bootstrap_rows)
|
||||
manifest = {
|
||||
"input": str(ALIGNED_PATH.relative_to(train.ROOT)),
|
||||
"input_sha256": train.sha256(ALIGNED_PATH),
|
||||
"evaluation_split": "official validation",
|
||||
"validation_samples": official["valid"].n,
|
||||
"source_video_groups": int(len(np.unique(official["valid"].groups))),
|
||||
"models": ["C0", "C5"],
|
||||
"modalities": {"T": "text", "A": "audio", "V": "vision"},
|
||||
"matched_type_sets": [label for _, label in MODALITY_SETS],
|
||||
"mask_rule": "per-sample target=min(per_sample_cap, audio_hide_capacity, vision_hide_capacity); split target evenly across selected modalities; continuous intervals",
|
||||
"per_sample_cap_rows": args.per_sample_cap,
|
||||
"total_added_feature_rows_per_type": total_masked,
|
||||
"mean_added_feature_rows_per_sample": per_sample_budget_mean,
|
||||
"zero_budget_samples": zero_budget_samples,
|
||||
"mask_seed": args.seed,
|
||||
"C5_evaluation_seed": args.seed + 1,
|
||||
"bootstrap_seed": args.seed + 2,
|
||||
"bootstrap_repeats": args.bootstrap_repeats,
|
||||
"C5_temperature": temperature,
|
||||
"C0_temperature": c0_temperature,
|
||||
"test_split_used": False,
|
||||
}
|
||||
(RESULTS / "matched_missing_type_manifest.json").write_text(
|
||||
json.dumps(manifest, indent=2, ensure_ascii=False), encoding="utf-8",
|
||||
)
|
||||
print(f"C0 temperature={c0_temperature:.6f}; C5 temperature={temperature:.6f}", flush=True)
|
||||
for row in summary_rows:
|
||||
if row["model"] != "C5" or row["scenario"] == "0.0/none":
|
||||
continue
|
||||
print(
|
||||
f"{row['model']} {row['scenario']}: added={row['rate_realized_additional_global']:.4f} "
|
||||
f"final={row['rate_final_total_missing_global']:.4f} "
|
||||
f"Acc={row['accuracy']:.4f} MacroF1={row['macro_f1']:.4f} "
|
||||
f"MAE={row['regression_mae']:.4f}",
|
||||
flush=True,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,316 @@
|
||||
"""Numerical checks for the Q2 state posterior, joint sampling, and decoder."""
|
||||
from __future__ import annotations
|
||||
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from scipy.special import betainc as scipy_betainc
|
||||
|
||||
from ...model.crg import CRG, ReliabilityGRU, StructuredGaussianImputer
|
||||
from .train import (
|
||||
_decode_mixture,
|
||||
_calibrated_mixture_moments,
|
||||
_group_ids,
|
||||
_missing_rate_summary,
|
||||
_predictive_intervals,
|
||||
_trajectory_variance_components,
|
||||
continuous_mask,
|
||||
controlled_group_bootstrap,
|
||||
gate_diagnostic_rows,
|
||||
regularized_beta,
|
||||
smooth_group_risk,
|
||||
validate_attachment3_predictions,
|
||||
)
|
||||
|
||||
|
||||
class StructuredGaussianTests(unittest.TestCase):
|
||||
def test_filter_nll_matches_dense_marginal_gaussian(self) -> None:
|
||||
torch.manual_seed(73)
|
||||
model = StructuredGaussianImputer((2, 2, 2)).double()
|
||||
xs = [torch.randn(1, 2, 2, dtype=torch.float64) for _ in range(3)]
|
||||
observed = torch.ones(1, 2, 3, dtype=torch.bool)
|
||||
got = model.observed_nll(xs, observed)[0]
|
||||
|
||||
with torch.no_grad():
|
||||
transition = model._transition()
|
||||
p0, q = model._covariances()
|
||||
emissions = model.emissions()
|
||||
emission = torch.cat(emissions, dim=0)
|
||||
noise = torch.block_diag(*[torch.diag(torch.nn.functional.softplus(raw) + 1e-4) for raw in model.r_raw])
|
||||
offset = torch.cat(list(model.biases))
|
||||
state_mean = torch.cat((model.mu0, transition @ model.mu0))
|
||||
p01 = p0 @ transition.T
|
||||
p11 = transition @ p0 @ transition.T + q
|
||||
state_cov = torch.cat((torch.cat((p0, p01), dim=1), torch.cat((p01.T, p11), dim=1)), dim=0)
|
||||
observation_map = torch.block_diag(emission, emission)
|
||||
observation_cov = observation_map @ state_cov @ observation_map.T + torch.block_diag(noise, noise)
|
||||
observation_mean = torch.cat((offset + emission @ model.mu0,
|
||||
offset + emission @ (transition @ model.mu0)))
|
||||
values = torch.cat((torch.cat([xs[m][0, 0] for m in range(3)]),
|
||||
torch.cat([xs[m][0, 1] for m in range(3)])))
|
||||
residual = values - observation_mean
|
||||
expected = 0.5 * (
|
||||
residual @ torch.linalg.solve(observation_cov, residual)
|
||||
+ torch.linalg.slogdet(observation_cov).logabsdet
|
||||
+ len(values) * np.log(2.0 * np.pi)
|
||||
)
|
||||
torch.testing.assert_close(got, expected, rtol=2e-4, atol=2e-4)
|
||||
|
||||
def test_joint_trajectory_draws_retain_temporal_dependence(self) -> None:
|
||||
torch.manual_seed(19)
|
||||
model = StructuredGaussianImputer((2, 2, 2))
|
||||
with torch.no_grad():
|
||||
for emission in model.emission_raw:
|
||||
emission.zero_()
|
||||
model.emission_raw[1][0, 0] = 1.0
|
||||
xs = [torch.zeros(1, 2, 2) for _ in range(3)]
|
||||
observed = torch.zeros(1, 2, 3, dtype=torch.bool)
|
||||
draws, _ = model.complete(xs, observed, 1600, joint_draws=True)
|
||||
temporal_correlation = float(np.corrcoef(draws[1][:, 0, 0, 0].cpu(), draws[1][:, 0, 1, 0].cpu())[0, 1])
|
||||
self.assertGreater(temporal_correlation, 0.15)
|
||||
|
||||
|
||||
class LossAndMaskTests(unittest.TestCase):
|
||||
def test_beta_cdf_matches_scipy(self) -> None:
|
||||
a = torch.tensor([0.7, 2.0, 5.0])
|
||||
b = torch.tensor([1.3, 3.0, 2.5])
|
||||
x = torch.tensor([0.2, 0.8, 0.55])
|
||||
actual = regularized_beta(x, a, b).detach().cpu().numpy()
|
||||
expected = scipy_betainc(a.numpy(), b.numpy(), x.numpy())
|
||||
np.testing.assert_allclose(actual, expected, rtol=2e-5, atol=2e-6)
|
||||
|
||||
def test_mask_is_contiguous_and_preserves_each_selected_source(self) -> None:
|
||||
original = np.ones((50, 3), dtype=bool)
|
||||
for mode in ("single", "sync", "partial", "async"):
|
||||
masked = continuous_mask(original, 0.5, mode, np.random.default_rng(101))
|
||||
hidden = original & ~masked
|
||||
for modality in range(3):
|
||||
positions = np.flatnonzero(hidden[:, modality])
|
||||
if len(positions):
|
||||
self.assertEqual(int(positions[-1] - positions[0] + 1), len(positions))
|
||||
self.assertGreaterEqual(int(masked[:, modality].sum()), 10)
|
||||
|
||||
def test_point_mask_keeps_rate_but_breaks_contiguous_span(self) -> None:
|
||||
original = np.ones((50, 3), dtype=bool)
|
||||
masked = continuous_mask(original, 0.3, "single", np.random.default_rng(887),
|
||||
modalities=(1,), kind="point")
|
||||
hidden = np.flatnonzero(original[:, 1] & ~masked[:, 1])
|
||||
self.assertEqual(len(hidden), 15)
|
||||
self.assertGreaterEqual(int(masked[:, 1].sum()), 10)
|
||||
runs = np.split(hidden, np.flatnonzero(np.diff(hidden) > 1) + 1)
|
||||
self.assertGreater(len([run for run in runs if len(run)]), 1)
|
||||
|
||||
def test_position_and_gap_structure_controls_hold_total_missing_fixed(self) -> None:
|
||||
original = np.ones((50, 3), dtype=bool)
|
||||
counts = []
|
||||
for location in ("start", "middle", "end"):
|
||||
masked = continuous_mask(
|
||||
original, 0.3, "single", np.random.default_rng(22),
|
||||
modalities=(0,), location=location,
|
||||
)
|
||||
hidden = np.flatnonzero(original[:, 0] & ~masked[:, 0])
|
||||
counts.append(len(hidden))
|
||||
if location == "start":
|
||||
self.assertEqual(int(hidden[0]), 0)
|
||||
elif location == "end":
|
||||
self.assertEqual(int(hidden[-1]), 49)
|
||||
else:
|
||||
self.assertLessEqual(abs(float(hidden.mean()) - 24.5), 1.0)
|
||||
self.assertEqual(counts, [15, 15, 15])
|
||||
|
||||
long = continuous_mask(
|
||||
original, 0.3, "single", np.random.default_rng(22),
|
||||
modalities=(0,), span_structure="long",
|
||||
)
|
||||
short = continuous_mask(
|
||||
original, 0.3, "single", np.random.default_rng(22),
|
||||
modalities=(0,), span_structure="multi_short",
|
||||
)
|
||||
long_hidden = np.flatnonzero(original[:, 0] & ~long[:, 0])
|
||||
short_hidden = np.flatnonzero(original[:, 0] & ~short[:, 0])
|
||||
self.assertEqual(len(long_hidden), len(short_hidden))
|
||||
short_runs = np.split(short_hidden, np.flatnonzero(np.diff(short_hidden) > 1) + 1)
|
||||
self.assertGreaterEqual(len([run for run in short_runs if len(run)]), 2)
|
||||
|
||||
def test_group_id_uses_any_newly_hidden_source(self) -> None:
|
||||
original = np.ones((2, 50, 3), dtype=bool)
|
||||
current = original.copy()
|
||||
current[0, 10:20, 1] = False
|
||||
current[1, 15:25, 2] = False
|
||||
groups = _group_ids(original, current)
|
||||
self.assertNotEqual(int(groups[0]), int(groups[1]))
|
||||
|
||||
def test_missing_rates_follow_equal_modality_pdf_denominators(self) -> None:
|
||||
original = np.asarray([
|
||||
[1, 1, 0], [1, 1, 0], [1, 0, 0], [1, 0, 0],
|
||||
], dtype=bool)
|
||||
current = original.copy()
|
||||
current[0, 0] = False
|
||||
rates = _missing_rate_summary(original, current)
|
||||
np.testing.assert_allclose(rates["natural_by_modality"], [0.0, 0.5, 1.0])
|
||||
self.assertAlmostEqual(rates["natural_global"], 0.5)
|
||||
np.testing.assert_allclose(rates["final_by_modality"], [0.25, 0.5, 1.0])
|
||||
self.assertAlmostEqual(rates["final_global"], 7.0 / 12.0)
|
||||
np.testing.assert_allclose(rates["additional_by_modality"][:2], [0.25, 0.0])
|
||||
self.assertTrue(np.isnan(rates["additional_by_modality"][2]))
|
||||
|
||||
def test_smooth_group_risk_matches_prior_weighted_formula(self) -> None:
|
||||
losses = torch.tensor([1.0, 2.0, 4.0], dtype=torch.float64)
|
||||
group_ids = np.asarray([0, 0, 1])
|
||||
lambda_group, tau = 0.2, 0.5
|
||||
group_losses = torch.tensor([1.5, 4.0], dtype=torch.float64)
|
||||
priors = torch.tensor([2 / 3, 1 / 3], dtype=torch.float64)
|
||||
expected = ((1 - lambda_group) * (priors * group_losses).sum()
|
||||
+ lambda_group * tau * torch.logsumexp(priors.log() + group_losses / tau, dim=0))
|
||||
actual = smooth_group_risk(losses, group_ids, lambda_group, tau)
|
||||
torch.testing.assert_close(actual, expected)
|
||||
|
||||
def test_controlled_group_bootstrap_is_paired_and_reports_aurc(self) -> None:
|
||||
split = SimpleNamespace(
|
||||
n=4,
|
||||
class_y=np.asarray([0, 0, 1, 2]),
|
||||
regression_y=np.asarray([-1.0, -0.5, 0.0, 1.0]),
|
||||
groups=np.asarray(["v1", "v1", "v2", "v3"]),
|
||||
mask=np.ones((4, 50, 3), dtype=bool),
|
||||
)
|
||||
scenarios = ["0.0/none"] + [f"{rate:.1f}/{mode}" for mode in ("single", "sync", "partial", "async")
|
||||
for rate in (0.1, 0.3, 0.5, 0.7)]
|
||||
scenario_masks = {}
|
||||
for scenario in scenarios:
|
||||
mask = split.mask.copy()
|
||||
rate_name, pattern = scenario.split("/", 1)
|
||||
rate = float(rate_name)
|
||||
if rate > 0:
|
||||
modality = {"single": 0, "sync": 0, "partial": 1, "async": 2}[pattern]
|
||||
count = int(round(rate * 50))
|
||||
mask[:, :count, modality] = False
|
||||
scenario_masks[scenario] = mask
|
||||
predictions = {}
|
||||
for model, shift in (("C0", 0.0), ("C1", 0.1)):
|
||||
predictions[model] = {}
|
||||
for index, scenario in enumerate(scenarios):
|
||||
predictions[model][scenario] = {
|
||||
"predicted_class": np.asarray([0, 1, 1, 2]),
|
||||
"predicted_score": split.regression_y + shift + index * 0.01,
|
||||
}
|
||||
rows = controlled_group_bootstrap(split, predictions, scenario_masks, repeats=20, seed=29)
|
||||
self.assertTrue(any(row["metric"] == "AURC_MAE" and row["model"] == "C1" for row in rows))
|
||||
paired = next(row for row in rows if row["model"] == "C1" and row["scenario"] == "0.3/single" and row["metric"] == "mae")
|
||||
self.assertAlmostEqual(paired["delta_estimate"], 0.1)
|
||||
self.assertAlmostEqual(paired["delta_to_natural_mae"], 0.02)
|
||||
self.assertEqual(paired["replicates"], 20)
|
||||
|
||||
def test_attachment3_submission_invariants(self) -> None:
|
||||
rows = [
|
||||
{"case_id": "case-a", "predicted_class": 0, "predicted_sentiment": -0.2,
|
||||
"p_negative": 0.5, "p_neutral": 0.3, "p_positive": 0.2,
|
||||
"interval_90_lower": -1.0, "interval_90_upper": 0.5},
|
||||
{"case_id": "case-b", "predicted_class": 1, "predicted_sentiment": 0.0,
|
||||
"p_negative": 0.2, "p_neutral": 0.6, "p_positive": 0.2,
|
||||
"interval_90_lower": -0.5, "interval_90_upper": 0.5},
|
||||
]
|
||||
validate_attachment3_predictions(["case-a", "case-b"], rows)
|
||||
rows[1]["predicted_sentiment"] = 1e-9
|
||||
with self.assertRaisesRegex(ValueError, "polarity mismatch"):
|
||||
validate_attachment3_predictions(["case-a", "case-b"], rows)
|
||||
|
||||
def test_gate_diagnostic_rows_keep_sample_position_and_modality(self) -> None:
|
||||
split = SimpleNamespace(ids=["v1$_$c1"], groups=np.asarray(["v1"]),
|
||||
mask=np.ones((1, 2, 3), dtype=bool))
|
||||
scalar = np.zeros((1, 2, 3), dtype=np.float32)
|
||||
predictions = {
|
||||
"fusion_weights": np.full((1, 2, 3), 0.2, dtype=np.float32),
|
||||
"null_weights": np.full((1, 2), 0.4, dtype=np.float32),
|
||||
"time_pool_weights": np.full((1, 2), 0.5, dtype=np.float32),
|
||||
"reliability": np.ones((1, 2, 3), dtype=np.float32),
|
||||
"imputation_uncertainty": scalar,
|
||||
"gap": scalar,
|
||||
"span": scalar,
|
||||
"distance_before": scalar,
|
||||
"distance_after": scalar,
|
||||
}
|
||||
rows = gate_diagnostic_rows(split, predictions)
|
||||
self.assertEqual(len(rows), 6)
|
||||
self.assertEqual(rows[0]["sample_id"], "v1$_$c1")
|
||||
self.assertEqual(rows[-1]["modality"], "vision")
|
||||
|
||||
def test_decoder_uses_neutral_priority_and_exact_zero(self) -> None:
|
||||
probabilities = np.asarray([[[1 / 3, 1 / 3, 1 / 3]], [[1 / 3, 1 / 3, 1 / 3]]], dtype=np.float32)
|
||||
beta = np.full((2, 1, 2, 2), 2.0, dtype=np.float32)
|
||||
_, classes, scores = _decode_mixture(probabilities, beta)
|
||||
self.assertEqual(int(classes[0]), 1)
|
||||
self.assertEqual(float(scores[0]), 0.0)
|
||||
|
||||
def test_calibrated_signed_mixture_interval_and_variance_components(self) -> None:
|
||||
probabilities = np.asarray(
|
||||
[[[0.25, 0.5, 0.25]], [[0.4, 0.2, 0.4]]], dtype=np.float64,
|
||||
)
|
||||
beta = np.full((2, 1, 2, 2), 2.0, dtype=np.float64)
|
||||
low, high = _predictive_intervals(probabilities, beta, temperature=1.5)
|
||||
self.assertLess(float(low[0]), 0.0)
|
||||
self.assertGreater(float(high[0]), 0.0)
|
||||
self.assertLess(float(low[0]), float(high[0]))
|
||||
total, within, between = _trajectory_variance_components(probabilities, beta)
|
||||
np.testing.assert_allclose(total, within + between, rtol=1e-6, atol=1e-7)
|
||||
mean_cold, variance_cold = _calibrated_mixture_moments(probabilities, beta, temperature=0.5)
|
||||
mean_warm, variance_warm = _calibrated_mixture_moments(probabilities, beta, temperature=2.0)
|
||||
self.assertTrue(np.isfinite(mean_cold).all() and np.isfinite(variance_cold).all())
|
||||
self.assertGreater(abs(float(variance_cold[0] - variance_warm[0])), 1e-5)
|
||||
|
||||
|
||||
class RecurrentAndVariantTests(unittest.TestCase):
|
||||
def test_gru_reset_gate_is_applied_before_candidate_recurrent_map(self) -> None:
|
||||
model = ReliabilityGRU(input_dim=1, hidden=1)
|
||||
with torch.no_grad():
|
||||
model.x_proj.weight.zero_()
|
||||
model.x_proj.bias.copy_(torch.tensor([10.0, 0.0, 1.0]))
|
||||
model.h_proj.weight.zero_()
|
||||
model.candidate_h.weight.fill_(2.0)
|
||||
x = torch.zeros(1, 2, 1)
|
||||
rho = torch.ones(1, 2)
|
||||
distance = torch.zeros(1, 2)
|
||||
actual = model._one_direction(x, rho, distance, reverse=False, reliability_update=False)
|
||||
z = torch.sigmoid(torch.tensor(10.0))
|
||||
first = z * torch.tanh(torch.tensor(1.0))
|
||||
reset = torch.sigmoid(torch.tensor(0.0))
|
||||
candidate = torch.tanh(torch.tensor(1.0) + 2.0 * reset * first)
|
||||
expected = (1.0 - z) * first + z * candidate
|
||||
torch.testing.assert_close(actual[0, 1, 0], expected)
|
||||
|
||||
def test_all_ablation_architectures_forward_and_backward(self) -> None:
|
||||
torch.manual_seed(9)
|
||||
options = {
|
||||
"C1": dict(use_imputer=False, use_joint_draws=False, use_final_gate=False, use_source_attention=False, reliability_update=False, use_low_rank=False),
|
||||
"C2": dict(use_imputer=True, use_joint_draws=False, use_final_gate=False, use_source_attention=False, reliability_update=False, use_low_rank=False),
|
||||
"C3": dict(use_imputer=True, use_joint_draws=True, use_final_gate=True, use_source_attention=False, reliability_update=False, use_low_rank=False),
|
||||
"C4": dict(use_imputer=True, use_joint_draws=True, use_final_gate=True, use_source_attention=True, reliability_update=False, use_low_rank=False),
|
||||
"C5": dict(use_imputer=True, use_joint_draws=True, use_final_gate=True, use_source_attention=True, reliability_update=True, use_low_rank=False),
|
||||
"C6": dict(use_imputer=True, use_joint_draws=True, use_final_gate=True, use_source_attention=True, reliability_update=True, use_low_rank=True),
|
||||
}
|
||||
xs = [torch.randn(1, 4, width) for width in (3, 2, 2)]
|
||||
observed = torch.ones(1, 4, 3, dtype=torch.bool)
|
||||
observed[:, 1:3, 1] = False
|
||||
for name, flags in options.items():
|
||||
with self.subTest(model=name):
|
||||
model = CRG(input_dims=(3, 2, 2), **flags)
|
||||
output = model(xs, observed, paths=2, joint_draws=flags["use_joint_draws"])
|
||||
loss = output["class_logits"].sum() + output["beta_params"].sum()
|
||||
loss.backward()
|
||||
expected_paths = 2 if flags["use_imputer"] else 1
|
||||
self.assertEqual(tuple(output["class_probs"].shape), (1, 3))
|
||||
self.assertEqual(tuple(output["fusion_weights_by_path"].shape), (expected_paths, 1, 4, 3))
|
||||
self.assertEqual(tuple(output["null_weights_by_path"].shape), (expected_paths, 1, 4))
|
||||
self.assertEqual(tuple(output["time_pool_weights_by_path"].shape), (expected_paths, 1, 4))
|
||||
torch.testing.assert_close(
|
||||
output["fusion_weights_by_path"].sum(dim=-1) + output["null_weights_by_path"],
|
||||
torch.ones((expected_paths, 1, 4)),
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
output["time_pool_weights_by_path"].sum(dim=-1), torch.ones((expected_paths, 1)),
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,158 @@
|
||||
"""Train the predeclared C5 Q2 architecture on Q1's exploratory index view."""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import time
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from . import train as q2_train
|
||||
from ...model.crg import INPUT_DIMS, StructuredGaussianImputer
|
||||
from .data import DATA_ROOT, fit_preprocessor, load_official_splits, transform_split
|
||||
from .train import (
|
||||
RESULTS as ALIGNED_RESULTS,
|
||||
SEED,
|
||||
_fit_neural,
|
||||
_make_variant,
|
||||
assert_group_disjoint,
|
||||
evaluate,
|
||||
fit_imputer,
|
||||
fit_temperature,
|
||||
group_bootstrap,
|
||||
label_resolution_from_train,
|
||||
make_reliability_scenarios,
|
||||
seed_everything,
|
||||
sha256,
|
||||
split_calibration,
|
||||
tune_reliability_hparams,
|
||||
write_csv,
|
||||
)
|
||||
|
||||
RESULTS = ALIGNED_RESULTS.parent / "results_unaligned"
|
||||
SOURCE = DATA_ROOT / "附件2-数据集特征文件" / "unaligned_50.pkl"
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--epochs", type=int, default=12)
|
||||
parser.add_argument("--imputer-epochs", type=int, default=8)
|
||||
parser.add_argument("--batch-size", type=int, default=64)
|
||||
parser.add_argument("--patience", type=int, default=3)
|
||||
parser.add_argument("--bootstrap-repeats", type=int, default=300)
|
||||
parser.add_argument("--seed", type=int, default=SEED)
|
||||
parser.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
|
||||
args = parser.parse_args()
|
||||
seed_everything(args.seed)
|
||||
device = torch.device(args.device)
|
||||
RESULTS.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
official = load_official_splits(SOURCE, version="unaligned_50")
|
||||
overlap = assert_group_disjoint(official)
|
||||
fit, heldout = split_calibration(official["train"], args.seed)
|
||||
reliability_validation, temperature_calibration = split_calibration(heldout, args.seed + 1, fraction=0.5)
|
||||
reliability_validation.name = "reliability_validation"
|
||||
temperature_calibration.name = "temperature_calibration"
|
||||
q2_train.DELTA_U = label_resolution_from_train(fit.regression_y)
|
||||
fitted = fit_preprocessor(fit)
|
||||
transformed = {name: transform_split(split, fitted) for name, split in official.items()}
|
||||
transformed["fit"] = transform_split(fit, fitted)
|
||||
transformed["reliability_validation"] = transform_split(reliability_validation, fitted)
|
||||
transformed["temperature_calibration"] = transform_split(temperature_calibration, fitted)
|
||||
np.savez_compressed(RESULTS / "preprocessor.npz", **{
|
||||
f"{modality}_{stat}": value
|
||||
for modality, values in fitted.items() for stat, value in values.items()
|
||||
})
|
||||
scenarios = make_reliability_scenarios(reliability_validation, args.seed + 906)
|
||||
print("official split sizes:", {k: v.n for k, v in official.items()}, flush=True)
|
||||
|
||||
imputer = StructuredGaussianImputer(INPUT_DIMS).to(device)
|
||||
imputer_history = fit_imputer(imputer, transformed["fit"], fit, device,
|
||||
args.imputer_epochs, args.batch_size, args.seed + 1)
|
||||
torch.save({k: v.detach().cpu() for k, v in imputer.state_dict().items()},
|
||||
RESULTS / "structured_imputer.pt")
|
||||
model = _make_variant("C5", imputer)
|
||||
model, history = _fit_neural(
|
||||
model, "C5", fit, official["valid"], transformed, device,
|
||||
args.epochs, args.batch_size, args.patience, np.random.default_rng(args.seed + 303),
|
||||
selection_split=reliability_validation,
|
||||
selection_arrays=transformed["reliability_validation"],
|
||||
selection_scenarios=scenarios,
|
||||
)
|
||||
reliability, tuning_rows = tune_reliability_hparams(
|
||||
model, transformed["reliability_validation"], reliability_validation,
|
||||
scenarios, device, args.batch_size, "C5", seed=args.seed + 551,
|
||||
)
|
||||
_, calibration_prediction = evaluate(model, transformed["temperature_calibration"],
|
||||
temperature_calibration, device, args.batch_size)
|
||||
temperature = fit_temperature(calibration_prediction["probabilities"],
|
||||
temperature_calibration.class_y)
|
||||
valid_metrics, _ = evaluate(model, transformed["valid"], official["valid"],
|
||||
device, args.batch_size, temperature=temperature)
|
||||
test_metrics, test_prediction = evaluate(model, transformed["test"], official["test"],
|
||||
device, args.batch_size, temperature=temperature)
|
||||
torch.save({k: v.detach().cpu() for k, v in model.state_dict().items()}, RESULTS / "crg_student.pt")
|
||||
(RESULTS / "validation_metrics.json").write_text(
|
||||
json.dumps({**valid_metrics, "model": "C5", "temperature": temperature}, indent=2), encoding="utf-8")
|
||||
(RESULTS / "test_metrics.json").write_text(
|
||||
json.dumps({**test_metrics, "model": "C5", "temperature": temperature}, indent=2), encoding="utf-8")
|
||||
write_csv(RESULTS / "reliability_hparam_tuning.csv", tuning_rows)
|
||||
write_csv(RESULTS / "training_history.csv", imputer_history + history)
|
||||
write_csv(RESULTS / "group_bootstrap_ci.csv",
|
||||
group_bootstrap(official["test"], test_prediction, args.bootstrap_repeats, args.seed + 44))
|
||||
rows = []
|
||||
for i, sample_id in enumerate(official["test"].ids):
|
||||
p = test_prediction["probabilities"][i]
|
||||
rows.append({
|
||||
"sample_id": sample_id,
|
||||
"source_video_id": official["test"].groups[i],
|
||||
"true_class": int(official["test"].class_y[i]),
|
||||
"predicted_class": int(test_prediction["predicted_class"][i]),
|
||||
"true_sentiment": float(official["test"].regression_y[i]),
|
||||
"predicted_sentiment": float(test_prediction["predicted_score"][i]),
|
||||
"p_negative": float(p[0]), "p_neutral": float(p[1]), "p_positive": float(p[2]),
|
||||
})
|
||||
write_csv(RESULTS / "test_predictions.csv", rows)
|
||||
manifest = {
|
||||
"scope": "exploratory unaligned_50 relative-index projection and prespecified C5 training",
|
||||
"physical_time_alignment": False,
|
||||
"input": str(SOURCE),
|
||||
"input_sha256": sha256(SOURCE),
|
||||
"text_encoder": "official precomputed text field; revision not supplied",
|
||||
"q1_adapter_audit": {name: split.alignment_audit for name, split in official.items()},
|
||||
"official_group_overlap": overlap,
|
||||
"internal_splits": {"fit": fit.n, "reliability_validation": reliability_validation.n,
|
||||
"temperature_calibration": temperature_calibration.n},
|
||||
"model": "C5 fixed before this run; no unaligned architecture selection",
|
||||
"quality": "no quality scores in official file; q*=1 for visible rows and J_Q=0",
|
||||
"seed": args.seed,
|
||||
"device": str(device),
|
||||
"device_name": torch.cuda.get_device_name(device) if device.type == "cuda" else "CPU",
|
||||
"torch_version": torch.__version__,
|
||||
"epochs_limit": args.epochs,
|
||||
"trained_c5_epochs": len(history),
|
||||
"selected_c5_epoch": int(min(history, key=lambda row: row["inner_selection_nll"])["epoch"]),
|
||||
"imputer_epochs": args.imputer_epochs,
|
||||
"batch_size": args.batch_size,
|
||||
"patience": args.patience,
|
||||
"selected_reliability": reliability,
|
||||
"temperature": temperature,
|
||||
"bootstrap_repeats": args.bootstrap_repeats,
|
||||
"attachment3": "not inferred: unaligned files lack numerical text and trusted lengths",
|
||||
"validation_metrics": valid_metrics,
|
||||
"test_metrics": test_metrics,
|
||||
"completed_utc": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()),
|
||||
}
|
||||
(RESULTS / "run_manifest.json").write_text(json.dumps(manifest, ensure_ascii=False, indent=2), encoding="utf-8")
|
||||
stale_teacher = RESULTS / "teacher.pt"
|
||||
if stale_teacher.exists():
|
||||
stale_teacher.unlink()
|
||||
print("C5 unaligned complete:", json.dumps({
|
||||
"accuracy": test_metrics["accuracy"], "macro_f1": test_metrics["macro_f1"],
|
||||
"mae": test_metrics["regression_mae"], "temperature": temperature,
|
||||
}), flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user