Complete standalone final deliverable and unaligned Q2 results

This commit is contained in:
2026-09-25 22:22:37 +08:00
parent c6b018e5d0
commit adc9c2064b
267 changed files with 15479 additions and 7976 deletions
+1
View File
@@ -0,0 +1 @@
"""Q2 training, comparison, and inference pipelines."""
+1
View File
@@ -0,0 +1 @@
"""Maintained Q2 deep-learning schemes."""
+1
View File
@@ -0,0 +1 @@
"""Q2 multimodal emotion-recognition experiments."""
+214
View File
@@ -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)
+4
View File
@@ -0,0 +1,4 @@
"""Compatibility import for the maintained model registry."""
from ....model.early_concat import AlignedFusionModel
__all__ = ["AlignedFusionModel"]
+4
View File
@@ -0,0 +1,4 @@
"""Compatibility import for the maintained model registry."""
from ....model.mofe import EXPERT_NAMES, SUBSETS, MixtureOfFusionExperts
__all__ = ["EXPERT_NAMES", "SUBSETS", "MixtureOfFusionExperts"]
+490
View File
@@ -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)
+787
View File
@@ -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()
+1
View File
@@ -0,0 +1 @@
"""Mathematical Q2 model family and training driver."""
+213
View File
@@ -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
+142
View File
@@ -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()
+90
View File
@@ -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()
+103
View File
@@ -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()
+369
View File
@@ -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()
+316
View File
@@ -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
+158
View File
@@ -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()