Prepare minimum submission bundle

This commit is contained in:
2026-09-26 16:36:33 +08:00
parent 9cdd604117
commit 411f0f97e5
172 changed files with 12565 additions and 0 deletions
+705
View File
@@ -0,0 +1,705 @@
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": "final.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()