706 lines
29 KiB
Python
706 lines
29 KiB
Python
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()
|