"""Retrain the two maintained Q2 models under the math/Q2 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 math/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": "math/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) -> None: 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") feature_path = ATTACHMENT2 / "aligned_50.pkl" raw_splits = load_splits(feature_path) 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 / "aligned_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 math/Q2 V2-compatible 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 Q1 physical-time bins", "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": "math/Q2 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: 按 math/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) arguments = parser.parse_args() run(device_name=arguments.device, output_dir=arguments.output_dir)