from __future__ import annotations import argparse import csv import hashlib import json import math import platform import random import subprocess import time from collections import Counter 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, confusion_matrix, f1_score, mean_absolute_error, mean_squared_error, recall_score from torch import nn from data_paths import ATTACHMENT2, PROJECT_ROOT from adapter import adapt_official_split from q2.deep_learning.q2.data import ( MODALITIES, RobustStats, Split, _ids_and_targets, _unpickle, apply_robust_stats, fit_robust_stats, ) from q2.deep_learning.q2.evaluate_math_protocol import ( SCENARIO_SEED, continuous_mask, make_scenarios, scenario_seed, ) from q2.deep_learning.q2.models import AlignedFusionModel from q2.deep_learning.q2.mofe import MixtureOfFusionExperts from q2.deep_learning.q2.train_compare import _loss as baseline_loss from q2.deep_learning.q2.train_mofe import EARLYCONCAT, MODEL_CONFIG, MOFE7_MLP from model.ati_ho import ATIHOModel, task_loss from model.ati_ho_config import ATIConfig, CONFIGS from .attribution import exact_shapley_audit from .audit import structural_audit Q3_ROOT = Path(__file__).resolve().parents[1] EXPERIMENT_ROOT = PROJECT_ROOT / "experiments" / "q3" / "ati_ho" SCALER_PATH = PROJECT_ROOT / "experiments" / "q2" / "unaligned_deep_two_b128" / "unaligned_50_robust_stats.npz" TRAIN_MASK_SEED = 20261227 TRAIN_RATES = (0.0, 0.1, 0.3, 0.5, 0.7) TRAIN_MODES = ("single", "sync", "partial", "async") SELECTION_SCENARIOS = ("0.0/none", "0.3/single", "0.3/sync", "0.5/async") MODEL_SEEDS = (42, 3407, 2026) BATCH_SIZE = 64 EPOCH_LIMIT = 12 PATIENCE = 3 LEARNING_RATE = 3e-4 WEIGHT_DECAY = 1e-3 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 torch.set_num_threads(4) def _sha256(path: Path) -> str: digest = hashlib.sha256() with path.open("rb") as stream: for chunk in iter(lambda: stream.read(1024 * 1024), b""): digest.update(chunk) return digest.hexdigest() def _group_count(ids: list[str]) -> int: return len({sample_id.split("$_$", 1)[0] for sample_id in ids}) def load_training_data() -> tuple[Split, Split, RobustStats, dict[str, Any]]: feature_path = ATTACHMENT2 / "unaligned_50.pkl" if not feature_path.is_file(): raise FileNotFoundError(f"missing official unaligned_50.pkl: {feature_path}") if not SCALER_PATH.is_file(): raise FileNotFoundError(f"missing Q2 train-only robust scaler: {SCALER_PATH}") source = _unpickle(feature_path) raw_splits: dict[str, Split] = {} adapter_audit: dict[str, Any] = {} group_sets: dict[str, set[str]] = {} sample_counts: dict[str, int] = {} for name in ("train", "valid", "test"): part = source[name] ids, y_cls, y_reg = _ids_and_targets(part) sample_counts[name] = len(ids) group_sets[name] = {sample_id.split("$_$", 1)[0] for sample_id in ids} if name == "test": continue arrays, mask, audit = adapt_official_split(part) raw_splits[name] = Split(tuple(arrays[m] for m in MODALITIES), mask, y_cls, y_reg, ids) adapter_audit[name] = audit overlap = { f"{first}/{second}": len(group_sets[first] & group_sets[second]) for first, second in (("train", "valid"), ("train", "test"), ("valid", "test")) } if any(overlap.values()): raise ValueError(f"official source-video groups overlap: {overlap}") train_raw, valid_raw = raw_splits["train"], raw_splits["valid"] del source expected_train_stats = fit_robust_stats(train_raw) stats = RobustStats.load(SCALER_PATH) deltas = [ float(np.max(np.abs(expected_train_stats.center[i] - stats.center[i]))) for i in range(3) ] + [ float(np.max(np.abs(expected_train_stats.scale[i] - stats.scale[i]))) for i in range(3) ] if max(deltas) > 2e-4: raise ValueError( "Q2 baseline scaler does not match the train-only scaler recomputed from the official split; " f"maximum coordinate difference={max(deltas):.6g}" ) train = apply_robust_stats(train_raw, stats) valid = apply_robust_stats(valid_raw, stats) if train.steps != 50 or valid.steps != 50: raise ValueError("ATI–HO requires the frozen 50-slot Relative-Progress interface") metadata = { "feature_file": str(feature_path), "feature_sha256": _sha256(feature_path), "scaler_file": str(SCALER_PATH), "scaler_max_abs_difference_from_train_only_recompute": max(deltas), "representation": "Q1 adapter Relative-Progress projection; 50 slots; not physical-time alignment", "adapter": "adapter.adapt_official_split; shared train-only robust scaler retained from Q2 V2", "dimensions": [int(x.shape[-1]) for x in train.x], "train_samples": train.n, "valid_samples": valid.n, "train_source_video_groups": len(group_sets["train"]), "valid_source_video_groups": len(group_sets["valid"]), "test_samples": sample_counts["test"], "test_source_video_groups": len(group_sets["test"]), "source_video_overlap_counts": overlap, "adapter_audit": adapter_audit, } return train, valid, stats, metadata def build_baseline(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 baseline: {method}") def _training_masks(split: Split, seed: int, epoch: int) -> tuple[np.ndarray, Counter[str]]: 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)) counts[f"{rate:.1f}/{mode}"] += 1 rows.append(continuous_mask(observed, rate, mode, rng)) return np.stack(rows), counts def _metric_row( split: Split, logits: np.ndarray, intensity: np.ndarray, probabilities: np.ndarray, *, method: str, seed: int, scenario: str, ) -> dict[str, Any]: pred_class = logits.argmax(axis=-1) y_cls = split.y_cls y_reg = split.y_reg confidence = probabilities.max(axis=-1) correct = (pred_class == y_cls).astype(np.float64) ece = 0.0 for left in np.linspace(0.0, 1.0, 16)[:-1]: right = left + 1.0 / 15.0 hit = (confidence >= left) & (confidence < right if right < 1.0 else confidence <= right) if hit.any(): ece += float(hit.mean() * abs(confidence[hit].mean() - correct[hit].mean())) one_hot = np.eye(3, dtype=np.float64)[y_cls] pearson = float(np.corrcoef(y_reg, intensity)[0, 1]) if np.std(y_reg) > 0 and np.std(intensity) > 0 else 0.0 cm = confusion_matrix(y_cls, pred_class, labels=[0, 1, 2]).tolist() return { "method": method, "seed": seed, "scenario": scenario, "samples": len(y_cls), "accuracy": float(accuracy_score(y_cls, pred_class)), "macro_f1": float(f1_score(y_cls, pred_class, labels=[0, 1, 2], average="macro", zero_division=0)), "weighted_f1": float(f1_score(y_cls, pred_class, average="weighted", zero_division=0)), "negative_recall": float(recall_score(y_cls, pred_class, labels=[0, 1, 2], average=None, zero_division=0)[0]), "neutral_recall": float(recall_score(y_cls, pred_class, labels=[0, 1, 2], average=None, zero_division=0)[1]), "positive_recall": float(recall_score(y_cls, pred_class, labels=[0, 1, 2], average=None, zero_division=0)[2]), "mae": float(mean_absolute_error(y_reg, intensity)), "rmse": float(math.sqrt(mean_squared_error(y_reg, intensity))), "pearson": pearson, "ece_15bin": ece, "brier_multiclass": float(np.mean(np.sum((probabilities - one_hot) ** 2, axis=-1))), "confusion_matrix_0_1_2": json.dumps(cm), } @torch.inference_mode() def _predict_arrays( model: nn.Module, split: Split, masks: np.ndarray, device: torch.device, *, batch_size: int = BATCH_SIZE, ati: bool, details: bool = False, ) -> dict[str, np.ndarray]: model.eval() outputs: dict[str, list[np.ndarray]] = {"logits": [], "intensity": [], "probabilities": []} if details: outputs.update({"params": [], "baseline": [], "main_effects": [], "pair_effects": []}) for start in range(0, split.n, batch_size): end = min(split.n, start + batch_size) xs = tuple(torch.as_tensor(x[start:end], dtype=torch.float32, device=device) for x in split.x) mask = torch.as_tensor(masks[start:end], dtype=torch.bool, device=device) result = model(xs, mask, return_details=details) if ati else model(xs, mask) logits = result["logits"] if ati: probs = result["probabilities"] intensity = result["intensity"] if details: for key in ("params", "baseline", "main_effects", "pair_effects"): outputs[key].append(result[key].detach().cpu().numpy()) else: probs = torch.softmax(logits, dim=-1) intensity = result["intensity"].clamp(-3.0, 3.0) outputs["logits"].append(logits.detach().cpu().numpy()) outputs["probabilities"].append(probs.detach().cpu().numpy()) outputs["intensity"].append(intensity.detach().cpu().numpy()) return {key: np.concatenate(values, axis=0) for key, values in outputs.items()} def _loss_on_masks( model: nn.Module, split: Split, masks: np.ndarray, device: torch.device, *, ati: bool, lambda_interaction: float, lambda_mask: float, ) -> float: model.eval() losses: list[float] = [] counts: list[int] = [] with torch.inference_mode(): for start in range(0, split.n, BATCH_SIZE): end = min(split.n, start + BATCH_SIZE) 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) output = model(xs, mb, return_details=False) if ati else model(xs, mb) if ati: loss, _ = task_loss( output, y_cls, y_reg, lambda_interaction=lambda_interaction, lambda_mask=0.0, ) else: loss = baseline_loss(output, y_cls, y_reg) losses.append(float(loss.item())) counts.append(end - start) return float(np.average(losses, weights=counts)) def _selection_loss( model: nn.Module, valid: Split, scenarios: dict[str, np.ndarray], device: torch.device, *, ati: bool, config: ATIConfig | None, ) -> float: return float( np.mean( [ _loss_on_masks( model, valid, scenarios[key], device, ati=ati, lambda_interaction=config.lambda_interaction if config else 0.0, lambda_mask=0.0, ) for key in SELECTION_SCENARIOS ] ) ) def _save_csv(path: Path, rows: list[dict[str, Any]], *, append: bool = False) -> 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)) write_header = not (append and path.exists() and path.stat().st_size > 0) mode = "a" if append else "w" with path.open(mode, newline="", encoding="utf-8-sig") as stream: writer = csv.DictWriter(stream, fieldnames=fields, extrasaction="ignore") if write_header: writer.writeheader() writer.writerows(rows) def _train_one( method: str, seed: int, train: Split, valid: Split, valid_scenarios: dict[str, np.ndarray], device: torch.device, *, epochs: int, force: bool, ) -> tuple[nn.Module, dict[str, Any]]: ati = method in CONFIGS config = CONFIGS[method] if ati else None run_dir = EXPERIMENT_ROOT / "models" / method / f"seed_{seed}" run_dir.mkdir(parents=True, exist_ok=True) checkpoint_path = run_dir / "model_best.pt" if checkpoint_path.is_file() and not force: saved = torch.load(checkpoint_path, map_location=device, weights_only=False) if saved.get("method") != method or int(saved.get("seed", -1)) != seed: raise ValueError(f"stale or mismatched checkpoint: {checkpoint_path}") model = ATIHOModel(tuple(x.shape[-1] for x in train.x), config).to(device) if ati else build_baseline(method, tuple(x.shape[-1] for x in train.x), device) model.load_state_dict(saved["state_dict"]) model.eval() return model, saved _seed_everything(seed) dims = tuple(int(x.shape[-1]) for x in train.x) model = ATIHOModel(dims, config).to(device) if ati else build_baseline(method, dims, device) optimizer = torch.optim.AdamW(model.parameters(), lr=LEARNING_RATE, weight_decay=WEIGHT_DECAY) train_x = tuple(torch.as_tensor(x, dtype=torch.float32, device=device) for x in train.x) train_cls = torch.as_tensor(train.y_cls, dtype=torch.long, device=device) train_reg = torch.as_tensor(train.y_reg, dtype=torch.float32, device=device) train_base_mask = torch.as_tensor(train.mask, dtype=torch.bool, device=device) order_rng = np.random.default_rng(seed + 809) orders = [order_rng.permutation(train.n) for _ in range(epochs)] history: list[dict[str, Any]] = [] mask_counts: Counter[str] = Counter() best_selection = math.inf best_epoch = 0 stale = 0 start_time = time.perf_counter() for epoch in range(1, epochs + 1): model.train() current_masks, current_counts = _training_masks(train, seed, epoch) mask_counts.update(current_counts) losses: list[float] = [] order = orders[epoch - 1] for start in range(0, train.n, BATCH_SIZE): index_np = order[start : start + BATCH_SIZE] index = torch.as_tensor(index_np, dtype=torch.long, device=device) mb = torch.as_tensor(current_masks[index_np], dtype=torch.bool, device=device) xs = tuple(x.index_select(0, index) for x in train_x) output = model(xs, mb, return_details=False) if ati else model(xs, mb) if ati: loss, loss_parts = task_loss( output, train_cls.index_select(0, index), train_reg.index_select(0, index), lambda_interaction=config.lambda_interaction, lambda_mask=config.lambda_mask, mask_target=train_base_mask.index_select(0, index), ) else: loss = baseline_loss(output, train_cls.index_select(0, index), train_reg.index_select(0, index)) loss_parts = {"total": loss} optimizer.zero_grad(set_to_none=True) loss.backward() nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) optimizer.step() losses.append(float(loss.detach().item())) selection = _selection_loss( model, valid, valid_scenarios, device, ati=ati, config=config ) clean = _loss_on_masks( model, valid, valid.mask, device, ati=ati, lambda_interaction=config.lambda_interaction if config else 0.0, lambda_mask=0.0, ) row = { "method": method, "seed": seed, "epoch": epoch, "train_loss": float(np.mean(losses)), "valid_selection_loss": selection, "valid_clean_loss": clean, "lambda_interaction": config.lambda_interaction if config else 0.0, "lambda_mask": config.lambda_mask if config else 0.0, } history.append(row) print( f"[{method} seed={seed}] epoch={epoch:02d} train={row['train_loss']:.4f} " f"valid_selection={selection:.4f} clean={clean:.4f}", flush=True, ) if selection < best_selection - 1e-4: best_selection = selection best_epoch = epoch stale = 0 state = { "method": method, "seed": seed, "dims": dims, "steps": train.steps, "config": config.to_dict() if config else None, "state_dict": model.state_dict(), "best_epoch": best_epoch, "best_selection_loss": best_selection, "protocol": "official Q2 V2 unaligned_50 Relative-Progress; train-only robust scaler; video-disjoint validation", } torch.save(state, 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() _save_csv(run_dir / "training_history.csv", history) (run_dir / "training_manifest.json").write_text( json.dumps( { "method": method, "seed": seed, "best_epoch": best_epoch, "best_selection_loss": best_selection, "elapsed_seconds": time.perf_counter() - start_time, "batch_size": BATCH_SIZE, "epoch_limit": epochs, "patience": PATIENCE, "optimizer": "AdamW", "learning_rate": LEARNING_RATE, "weight_decay": WEIGHT_DECAY, "gradient_clip_norm": 1.0, "training_mask_rates": list(TRAIN_RATES), "training_mask_patterns": list(TRAIN_MODES), "training_mask_seed_base": TRAIN_MASK_SEED, "same_orders_and_masks_across_methods_for_same_seed": True, "config": config.to_dict() if config else {"model_config": MODEL_CONFIG}, "history": history, "mask_counts": dict(mask_counts), }, ensure_ascii=False, indent=2, ), encoding="utf-8", ) return model, saved def _evaluate_job( method: str, seed: int, model: nn.Module, valid: Split, scenarios: dict[str, np.ndarray], device: torch.device, ) -> list[dict[str, Any]]: ati = method in CONFIGS rows: list[dict[str, Any]] = [] selected = {"clean": valid.mask} selected.update({key: scenarios[key] for key in SELECTION_SCENARIOS if key != "0.0/none"}) for name, masks in selected.items(): predictions = _predict_arrays(model, valid, masks, device, ati=ati) rows.append( _metric_row( valid, predictions["logits"], predictions["intensity"], predictions["probabilities"], method=method, seed=seed, scenario=name, ) ) return rows def _write_root_manifest(data_meta: dict[str, Any], device: torch.device, epochs: int) -> None: EXPERIMENT_ROOT.mkdir(parents=True, exist_ok=True) try: git_sha = subprocess.check_output( ["git", "rev-parse", "HEAD"], cwd=PROJECT_ROOT, text=True, stderr=subprocess.DEVNULL ).strip() except Exception: git_sha = None info: dict[str, Any] = { "experiment": "ATI–HO Q3 staged training and structural attribution audit", "created_utc": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()), "device": str(device), "torch_version": torch.__version__, "cuda_available": torch.cuda.is_available(), "cuda_version": torch.version.cuda, "gpu": torch.cuda.get_device_name(0) if torch.cuda.is_available() else None, "python": platform.python_version(), "seeds": list(MODEL_SEEDS), "epochs_max": epochs, "training_protocol": { "batch_size": BATCH_SIZE, "early_stopping_patience": PATIENCE, "optimizer": "AdamW", "learning_rate": LEARNING_RATE, "weight_decay": WEIGHT_DECAY, "gradient_clip_norm": 1.0, "training_mask_rates": list(TRAIN_RATES), "training_mask_patterns": list(TRAIN_MODES), "validation_selection_scenarios": list(SELECTION_SCENARIOS), "held_out_attachment4_touched_during_training": False, }, "ati_output": { "parameter_vector": "3 centered class logits + r_negative + r_positive", "intensity": "negative/positive magnitudes are 3*sigmoid(r); neutral class is exactly zero", "loss": "cross entropy + conditional magnitude SmoothL1 + 0.2*Huber(delta=0.25) + configured regularizers", "baseline_checkpoint_reuse": "No: retrain B0 and B1 on the fixed ATI split/mask schedule because existing Q2 checkpoints differ in seeds, batch size, and schedule.", "calibration_temperature": 1.0, }, "data": data_meta, } (EXPERIMENT_ROOT / "run_manifest.json").write_text( json.dumps(info, ensure_ascii=False, indent=2), encoding="utf-8" ) def run_stage1(train: Split, valid: Split, scenarios: dict[str, np.ndarray], device: torch.device, epochs: int, force: bool) -> None: (EXPERIMENT_ROOT / "stage1_complete.json").unlink(missing_ok=True) rows: list[dict[str, Any]] = [] models: dict[str, nn.Module] = {} for method in ("A0", "A1", "A2", "A3", "D0"): model, saved = _train_one(method, 42, train, valid, scenarios, device, epochs=epochs, force=force) models[method] = model rows.extend(_evaluate_job(method, 42, model, valid, scenarios, device)) print(f"[{method}] best_epoch={saved['best_epoch']} selected_loss={saved['best_selection_loss']:.5f}", flush=True) _save_csv(EXPERIMENT_ROOT / "validation_results.csv", rows) candidates = [] for method in ("A0", "A1", "A2", "A3"): checkpoint = torch.load(EXPERIMENT_ROOT / "models" / method / "seed_42" / "model_best.pt", map_location="cpu", weights_only=False) candidates.append({ "method": method, "best_selection_loss": float(checkpoint["best_selection_loss"]), "best_epoch": int(checkpoint["best_epoch"]), }) candidates.sort(key=lambda row: row["best_selection_loss"]) choice = candidates[0]["method"] candidate_doc = { "stage1_candidates": candidates, "provisional_selected_candidate": choice, "selection_rule": "lowest fixed four-scenario ATI task loss on the locked official validation split; seed 42 only in Stage I", } (EXPERIMENT_ROOT / "provisional_candidate.json").write_text( json.dumps(candidate_doc, ensure_ascii=False, indent=2), encoding="utf-8" ) smoke_count = min(16, valid.n) smoke_xs = tuple(torch.as_tensor(x[:smoke_count], dtype=torch.float32, device=device) for x in valid.x) smoke_mask = torch.as_tensor(valid.mask[:smoke_count], dtype=torch.bool, device=device) structural_rows = [] for method, model in models.items(): report = structural_audit(model, smoke_xs, smoke_mask) row = {"method": method, "seed": 42, "samples": smoke_count, **report} row["pair_single_missing_anchor_max_abs"] = json.dumps( report["pair_single_missing_anchor_max_abs"], sort_keys=True ) structural_rows.append(row) if not report["checks_pass"]: raise RuntimeError(f"Stage I structural audit failed for {method}: {report}") if method == "D0" and not report["unanchored_control_detected_leakage"]: raise RuntimeError("D0 unanchored diagnostic did not expose the expected missing-modality leakage") _save_csv(EXPERIMENT_ROOT / "structural_audit.csv", structural_rows) shapley = exact_shapley_audit([models[choice]], smoke_xs, smoke_mask, batch_size=32) if not np.asarray(shapley["class_pass"]).all(): raise RuntimeError( f"Stage I analytic-vs-exact Shapley audit failed for {choice}: " f"max_abs={float(np.max(shapley['class_abs_error'])):.8g}" ) shapley_rows = [] for index in range(smoke_count): shapley_rows.append( { "sample_index": index, "method": choice, "target_class": int(shapley["target_class"][index]), "runner_up_class": int(shapley["runner_up_class"][index]), "analytic_T": float(shapley["analytic_class"][index, 0]), "analytic_A": float(shapley["analytic_class"][index, 1]), "analytic_V": float(shapley["analytic_class"][index, 2]), "exact_T": float(shapley["exact_class"][index, 0]), "exact_A": float(shapley["exact_class"][index, 1]), "exact_V": float(shapley["exact_class"][index, 2]), "max_abs_error": float(shapley["class_abs_error"][index].max()), "pass": bool(shapley["class_pass"][index].all()), } ) _save_csv(EXPERIMENT_ROOT / "shapley_audit_seed42_smoke.csv", shapley_rows) (EXPERIMENT_ROOT / "stage1_complete.json").write_text( json.dumps( { "models": ["A0", "A1", "A2", "A3", "D0"], "seed": 42, "structural_audit_pass": True, "analytic_vs_exact_shapley_pass": True, "shapley_smoke_samples": smoke_count, "selected_candidate": choice, }, indent=2, ), encoding="utf-8", ) def run_stage2(train: Split, valid: Split, scenarios: dict[str, np.ndarray], device: torch.device, epochs: int, force: bool) -> None: provisional_path = EXPERIMENT_ROOT / "provisional_candidate.json" if not provisional_path.is_file(): raise FileNotFoundError("run Stage I before Stage II; provisional_candidate.json is missing") selected = json.loads(provisional_path.read_text(encoding="utf-8"))["provisional_selected_candidate"] key_ablations = { "A0": ["A1", "A2"], "A1": ["A0", "A2"], "A2": ["A0", "A1"], "A3": ["A2", "A1"], }[selected] methods = list(dict.fromkeys([selected, *key_ablations])) rows: list[dict[str, Any]] = [] for method in (EARLYCONCAT, MOFE7_MLP, *methods): for seed in MODEL_SEEDS: model, saved = _train_one(method, seed, train, valid, scenarios, device, epochs=epochs, force=force) rows.extend(_evaluate_job(method, seed, model, valid, scenarios, device)) print( f"[{method} seed={seed}] best_epoch={saved['best_epoch']} " f"selected_loss={saved['best_selection_loss']:.5f}", flush=True, ) _save_csv(EXPERIMENT_ROOT / "validation_results.csv", rows, append=True) stage2 = { "selected_candidate": selected, "key_ablations": key_ablations, "baseline_methods": [EARLYCONCAT, MOFE7_MLP], "seeds": list(MODEL_SEEDS), "baseline_checkpoints_retrained": True, "all_selection_uses_locked_validation_only": True, } (EXPERIMENT_ROOT / "stage2_complete.json").write_text( json.dumps(stage2, ensure_ascii=False, indent=2), encoding="utf-8" ) def main() -> None: parser = argparse.ArgumentParser(description="Train ATI–HO and compatible Q3 baselines.") parser.add_argument("--phase", choices=("stage1", "stage2", "all"), default="all") parser.add_argument("--device", default="auto") parser.add_argument("--epochs", type=int, default=EPOCH_LIMIT) parser.add_argument("--force", action="store_true") args = parser.parse_args() device = torch.device("cuda" if args.device == "auto" and torch.cuda.is_available() else "cpu" if args.device == "auto" else args.device) train, valid, _stats, data_meta = load_training_data() scenarios = make_scenarios(valid, SCENARIO_SEED) _write_root_manifest(data_meta, device, args.epochs) print( f"ATI–HO protocol: train={train.n} valid={valid.n} groups=" f"{data_meta['train_source_video_groups']}/{data_meta['valid_source_video_groups']} " f"dims={data_meta['dimensions']} device={device}", flush=True, ) if args.phase in {"stage1", "all"}: run_stage1(train, valid, scenarios, device, args.epochs, args.force) if args.phase in {"stage2", "all"}: run_stage2(train, valid, scenarios, device, args.epochs, args.force) if __name__ == "__main__": main()