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
+1
View File
@@ -0,0 +1 @@
"""Q3 interpretable emotion-recognition pipeline."""
+46
View File
@@ -0,0 +1,46 @@
# ATI–HO 训练与评估
ATI–HO 是当前 Q3 方案。模型定义位于 `final/model/ati_ho.py` 和 `final/model/ati_ho_config.py`。
## 输入准备
设置 `FINAL_DATA_DIR` 指向官方数据根目录。完整训练和评估需要附件 2 的 `unaligned_50.pkl`、附件 4 的未对齐特征文件和训练集 robust scaler:
`final/experiments/q2/unaligned_deep_two_b128/unaligned_50_robust_stats.npz`
评估时附件 4 视频文件仅用于来源核验,不是模型输入。模型只读取特征文件。所有未对齐输入由统一 adapter 投影到 50 个 Relative-Progress 槽,不能解释为物理时间同步。
## 训练命令
从仓库根目录执行:
```bash
export FINAL_DATA_DIR="/path/to/E题数据"
python -m final.q3.ati_ho.train --phase all --device auto
```
`--phase stage1` 运行 seed 42 的 A0/A1/A2/A3 和 D0 结构审计;`--phase stage2` 使用 Stage I 写出的 provisional candidate 运行三 seed 基线和关键消融。默认最多 12 轮,固定验证情景任务损失早停。训练记录与检查点保存在 `final/experiments/q3/ati_ho/`。只有明确要覆盖检查点时才加 `--force`。
## 评估命令
```bash
export FINAL_DATA_DIR="/path/to/E题数据"
python -m final.q3.ati_ho.evaluate --device auto
```
评估会比较 ATI 消融、重算官方验证指标和按来源视频组 Bootstrap、审计解析/精确 Shapley、运行附件 4 局部 Owen 和删除/保留诊断,并生成 Q3 输出。最终 ATI 方案按三 seed、四个固定验证情景的平均任务损失选出。附件 4 标签不会读取或用于报告;附件 4 输出没有准确率。
若完整评估已写完 CSV,但报告阶段中断,可运行:
```bash
python -m final.q3.ati_ho.evaluate --reports-only
```
若只需补算训练 seed 与 1% 输入扰动下的 attribution 稳定性:
```bash
export FINAL_DATA_DIR="/path/to/E题数据"
python -m final.q3.ati_ho.evaluate --device auto --stability-only
```
题目交付文件写入 `final/output/q3/ati_ho/`。完整结果表、审计和论文式记录位于 `final/experiments/q3/ati_ho/results/ati_ho/`。验证指标解释和结果边界见 `final/REPORTS.md`。
+6
View File
@@ -0,0 +1,6 @@
"""ATI–HO: anchored temporal interactions with hierarchical Owen attribution."""
from ...model.ati_ho import ATIHOModel
from ...model.ati_ho_config import ATIConfig, CONFIGS
__all__ = ["ATIConfig", "CONFIGS", "ATIHOModel"]
+166
View File
@@ -0,0 +1,166 @@
from __future__ import annotations
import itertools
import math
from typing import Any, Sequence
import numpy as np
import torch
from torch import nn
def ensemble_forward(
models: Sequence[nn.Module],
xs: tuple[torch.Tensor, torch.Tensor, torch.Tensor],
masks: torch.Tensor,
*,
details: bool = True,
) -> dict[str, Any]:
"""Average additive parameters first, then decode the ensemble prediction."""
outputs = []
for model in models:
try:
outputs.append(model(xs, masks, return_details=details))
except TypeError:
outputs.append(model(xs, masks))
if "params" not in outputs[0]:
logits = torch.stack([output["logits"] for output in outputs], dim=0).mean(dim=0)
probabilities = torch.stack(
[torch.softmax(output["logits"], dim=-1) for output in outputs], dim=0
).mean(dim=0)
intensity = torch.stack([output["intensity"] for output in outputs], dim=0).mean(dim=0)
result = {
"logits": logits,
"probabilities": probabilities,
"predicted_class": logits.argmax(dim=-1),
"intensity": intensity,
}
if "utility" in outputs[0]:
result["utility"] = torch.stack(
[output["utility"] for output in outputs], dim=0
).mean(dim=0)
return result
result: dict[str, Any] = {}
averaged = ("params", "baseline", "main_effects", "pair_effects", "mask_logits")
for key in averaged:
if key in outputs[0]:
result[key] = torch.stack([output[key] for output in outputs], dim=0).mean(dim=0)
params = result["params"]
logits = params[:, :3]
probabilities = torch.softmax(logits, dim=-1)
nu_negative = 3.0 * torch.sigmoid(params[:, 3])
nu_positive = 3.0 * torch.sigmoid(params[:, 4])
predicted_class = logits.argmax(dim=-1)
intensity = torch.where(
predicted_class == 0,
-nu_negative,
torch.where(predicted_class == 2, nu_positive, torch.zeros_like(nu_positive)),
)
result.update(
{
"logits": logits,
"probabilities": probabilities,
"predicted_class": predicted_class,
"intensity": intensity,
"nu_negative": nu_negative,
"nu_positive": nu_positive,
"soft_intensity": probabilities[:, 2] * nu_positive - probabilities[:, 0] * nu_negative,
}
)
return result
def shapley_from_eight(values: np.ndarray) -> np.ndarray:
"""Exact three-player Shapley values from the eight coalition values."""
values = np.asarray(values, dtype=np.float64)
if values.shape[-1] != 8:
raise ValueError("the three-modality game requires exactly eight coalition values")
result = np.zeros((*values.shape[:-1], 3), dtype=np.float64)
factorial = math.factorial
for modality in range(3):
others = [i for i in range(3) if i != modality]
for size in range(3):
weight = factorial(size) * factorial(2 - size) / factorial(3)
for subset in itertools.combinations(others, size):
before = sum(1 << item for item in subset)
after = before | (1 << modality)
result[..., modality] += weight * (values[..., after] - values[..., before])
return result
def analytic_class_shapley(details: dict[str, torch.Tensor], target: torch.Tensor, other: torch.Tensor) -> np.ndarray:
"""Closed form for the fixed target-vs-runner-up logit margin."""
main = details["main_effects"]
pairs = details["pair_effects"]
delta = torch.zeros((main.shape[0], 3), dtype=main.dtype, device=main.device)
rows = torch.arange(main.shape[0], device=main.device)
delta[rows, target] = 1.0
delta[rows, other] = -1.0
contributions = main.clone()
pair_modalities = ((0, 1), (0, 2), (1, 2))
for pair_idx, (left, right) in enumerate(pair_modalities):
contributions[:, left] = contributions[:, left] + 0.5 * pairs[:, pair_idx]
contributions[:, right] = contributions[:, right] + 0.5 * pairs[:, pair_idx]
values = torch.einsum("bi,bmi->bm", delta, contributions[..., :3])
return values.detach().cpu().numpy().astype(np.float64)
@torch.inference_mode()
def exact_shapley_audit(
models: Sequence[nn.Module],
xs: tuple[torch.Tensor, torch.Tensor, torch.Tensor],
masks: torch.Tensor,
*,
batch_size: int = 128,
) -> dict[str, np.ndarray]:
"""Compare analytic output-parameter Shapley with exact 8-coalition values.
The class target and runner-up are fixed from each sample's full-input
prediction. A second exact game is computed for the decoded hard intensity;
that value is nonlinear and is not compared with the analytic formula.
"""
n = masks.shape[0]
full = ensemble_forward(models, xs, masks, details=True)
logits = full["logits"]
target = logits.argmax(dim=-1)
ranked = logits.argsort(dim=-1, descending=True)
other = ranked[:, 1]
analytic = analytic_class_shapley(full, target, other)
margin_values = np.zeros((n, 8), dtype=np.float64)
intensity_values = np.zeros((n, 8), dtype=np.float64)
for coalition in range(8):
for start in range(0, n, batch_size):
end = min(n, start + batch_size)
current_mask = masks[start:end].clone()
for modality in range(3):
if not coalition & (1 << modality):
current_mask[..., modality] = False
current_xs = tuple(x[start:end] for x in xs)
output = ensemble_forward(models, current_xs, current_mask, details=False)
local_rows = torch.arange(end - start, device=logits.device)
local_target = target[start:end]
local_other = other[start:end]
margin = (
output["logits"][local_rows, local_target]
- output["logits"][local_rows, local_other]
)
margin_values[start:end, coalition] = margin.detach().cpu().numpy()
intensity_values[start:end, coalition] = output["intensity"].detach().cpu().numpy()
exact = shapley_from_eight(margin_values)
exact_intensity = shapley_from_eight(intensity_values)
error = np.abs(analytic - exact)
tolerance = 1e-6 + 1e-5 * np.abs(exact)
return {
"analytic_class": analytic,
"exact_class": exact,
"class_abs_error": error,
"class_pass": error <= tolerance,
"exact_intensity": exact_intensity,
"coalition_margin": margin_values,
"coalition_intensity": intensity_values,
"target_class": target.detach().cpu().numpy(),
"runner_up_class": other.detach().cpu().numpy(),
"full_output": full,
}
+68
View File
@@ -0,0 +1,68 @@
from __future__ import annotations
from typing import Any
import torch
from torch import nn
PAIR_MODES = ((0, 1), (0, 2), (1, 2))
@torch.inference_mode()
def structural_audit(
model: nn.Module,
xs: tuple[torch.Tensor, torch.Tensor, torch.Tensor],
masks: torch.Tensor,
*,
atol: float = 1e-6,
) -> dict[str, Any]:
model.eval()
output = model(xs, masks, return_details=True)
reconstructed = output["baseline"] + output["main_effects"].sum(dim=1) + output["pair_effects"].sum(dim=1)
additive_residual = torch.max(torch.abs(reconstructed - output["params"])).item()
absent_xs = tuple(torch.zeros_like(x) for x in xs)
absent_mask = torch.zeros_like(masks, dtype=torch.bool)
absent = model(absent_xs, absent_mask, return_details=True)
main_zero_residual = torch.max(torch.abs(absent["main_effects"])).item()
full_zero_residual = torch.max(torch.abs(absent["params"] - absent["baseline"])).item()
pair_zero_residual = torch.max(torch.abs(absent["pair_effects"])).item()
pair_anchor_residuals: dict[str, float] = {}
for pair_index, (left, right) in enumerate(PAIR_MODES):
maxima = []
for hidden in (left, right):
altered = masks.clone()
altered[..., hidden] = False
out = model(xs, altered, return_details=True)
maxima.append(torch.max(torch.abs(out["pair_effects"][:, pair_index])).item())
pair_anchor_residuals[f"{left}{right}"] = max(maxima)
finite_count = 0
nonfinite_count = 0
for key in ("params", "logits", "intensity", "main_effects", "pair_effects"):
tensor = output[key]
finite_count += int(torch.isfinite(tensor).sum().item())
nonfinite_count += int((~torch.isfinite(tensor)).sum().item())
anchored = bool(getattr(getattr(model, "config", None), "anchored", False))
main_additive_pass = max(additive_residual, main_zero_residual) <= atol
pair_pass = max(pair_anchor_residuals.values(), default=0.0) <= atol
baseline_pass = full_zero_residual <= atol and pair_zero_residual <= atol
checks_pass = main_additive_pass and nonfinite_count == 0 and (baseline_pass if anchored else True)
return {
"anchored": anchored,
"main_effect_zero_anchor_max_abs": main_zero_residual,
"pair_effect_zero_anchor_max_abs": pair_zero_residual,
"pair_single_missing_anchor_max_abs": pair_anchor_residuals,
"additive_reconstruction_max_abs": additive_residual,
"all_modalities_missing_equals_baseline_max_abs": full_zero_residual,
"finite_value_count": finite_count,
"nonfinite_value_count": nonfinite_count,
"main_and_additivity_pass": main_additive_pass,
"full_baseline_anchor_pass": baseline_pass if anchored else None,
"pair_anchor_pass": pair_pass if anchored else None,
"unanchored_control_detected_leakage": ((not pair_pass) or not baseline_pass) if not anchored else False,
"checks_pass": checks_pass,
}
File diff suppressed because it is too large Load Diff
+240
View File
@@ -0,0 +1,240 @@
from __future__ import annotations
from typing import Any, Sequence
import numpy as np
import torch
from torch import nn
from .attribution import ensemble_forward
def _margin_values(
models: Sequence[nn.Module],
xs: tuple[torch.Tensor, torch.Tensor, torch.Tensor],
masks: np.ndarray,
target: int,
other: int,
*,
batch_size: int = 64,
) -> np.ndarray:
device = xs[0].device
values: list[np.ndarray] = []
with torch.inference_mode():
for start in range(0, len(masks), batch_size):
end = min(len(masks), start + batch_size)
current_mask = torch.as_tensor(masks[start:end], dtype=torch.bool, device=device)
repeated = tuple(x.expand(end - start, -1, -1).contiguous() for x in xs)
output = ensemble_forward(models, repeated, current_mask, details=False)
margin = output["logits"][:, target] - output["logits"][:, other]
values.append(margin.detach().cpu().numpy().astype(np.float64))
return np.concatenate(values) if values else np.empty(0, dtype=np.float64)
def _top_stability(previous: np.ndarray, current: np.ndarray, k: int = 5) -> float:
old = set(np.argsort(-np.abs(previous), kind="stable")[:k].tolist())
new = set(np.argsort(-np.abs(current), kind="stable")[:k].tolist())
return float(len(old & new) / max(1, len(old | new)))
@torch.inference_mode()
def hierarchical_owen_one(
models: Sequence[nn.Module],
xs: tuple[torch.Tensor, torch.Tensor, torch.Tensor],
mask: torch.Tensor,
*,
seed: int,
bins_per_modality: int = 10,
start_permutations: int = 8,
max_permutations: int = 64,
batch_size: int = 64,
) -> dict[str, Any]:
"""Estimate a three-group Owen allocation over relative-progress segments.
The outer permutation orders modalities. Each modality's 10 temporal bins
are then added in a random inner permutation. This is a sampled Owen value,
not a perturbation of arbitrary individual feature dimensions.
"""
model_output = ensemble_forward(models, xs, mask, details=False)
logits = model_output["logits"][0]
target = int(logits.argmax().item())
other = int(logits.argsort(descending=True)[1].item())
full_margin = float((logits[target] - logits[other]).item())
no_features = torch.zeros_like(mask, dtype=torch.bool)
baseline = ensemble_forward(models, xs, no_features, details=False)["logits"][0]
base_margin = float((baseline[target] - baseline[other]).item())
base_mask = mask[0].detach().cpu().numpy().astype(bool)
steps = base_mask.shape[0]
boundaries = np.linspace(0, steps, bins_per_modality + 1).round().astype(int)
slices = [(int(boundaries[k]), int(boundaries[k + 1])) for k in range(bins_per_modality)]
rng = np.random.default_rng(seed)
draws: list[np.ndarray] = []
previous_mean: np.ndarray | None = None
stability = float("nan")
schedule = [start_permutations]
while schedule[-1] < max_permutations:
schedule.append(min(max_permutations, schedule[-1] * 2))
next_target = schedule[0]
stopping_status = "max_permutations"
while len(draws) < max_permutations:
needed = min(8, max_permutations - len(draws))
all_masks: list[np.ndarray] = []
player_order: list[list[tuple[int, int]]] = []
for _ in range(needed):
current = np.zeros_like(base_mask, dtype=bool)
order: list[tuple[int, int]] = []
outer = rng.permutation(3)
for modality in outer:
for bin_index in rng.permutation(bins_per_modality):
left, right = slices[int(bin_index)]
current[left:right, int(modality)] = base_mask[left:right, int(modality)]
order.append((int(modality), int(bin_index)))
all_masks.append(current.copy())
player_order.append(order)
scores = _margin_values(
models,
xs,
np.stack(all_masks),
target,
other,
batch_size=batch_size,
)
cursor = 0
for order in player_order:
previous_score = base_margin
draw = np.zeros((3, bins_per_modality), dtype=np.float64)
for modality, bin_index in order:
current_score = float(scores[cursor])
cursor += 1
draw[modality, bin_index] = current_score - previous_score
previous_score = current_score
draws.append(draw)
if len(draws) >= next_target:
current_mean = np.mean(np.stack(draws), axis=0)
flat_mean = current_mean.reshape(-1)
standard_error = np.std(np.stack(draws), axis=0, ddof=1) / np.sqrt(len(draws))
se_mean = float(np.mean(standard_error))
if previous_mean is not None:
stability = _top_stability(previous_mean, flat_mean)
se_limit = max(0.02, 0.10 * abs(full_margin - base_margin))
if stability >= 0.8 and se_mean <= se_limit:
stopping_status = "stable"
break
previous_mean = flat_mean
next_idx = next((i for i, value in enumerate(schedule) if value > len(draws)), None)
if next_idx is None:
break
next_target = schedule[next_idx]
draw_array = np.stack(draws)
mean = draw_array.mean(axis=0)
standard_error = draw_array.std(axis=0, ddof=1) / np.sqrt(len(draws)) if len(draws) > 1 else np.full_like(mean, np.nan)
conservation = float(mean.sum() - (full_margin - base_margin))
return {
"target_class": target,
"runner_up_class": other,
"full_margin": full_margin,
"baseline_margin": base_margin,
"contribution": mean,
"standard_error": standard_error,
"permutations": len(draws),
"stopping_status": stopping_status,
"top5_jaccard_last_check": stability,
"local_conservation_residual": conservation,
"bin_slices": slices,
}
@torch.inference_mode()
def fidelity_audit_one(
models: Sequence[nn.Module],
xs: tuple[torch.Tensor, torch.Tensor, torch.Tensor],
mask: torch.Tensor,
contribution: np.ndarray,
*,
sample_id: str,
seed: int,
bins_per_modality: int = 10,
random_replicates: int = 20,
) -> list[dict[str, Any]]:
"""Deletion/retention at 10/20/30%, matched by modality and segment count."""
base_mask = mask[0].detach().cpu().numpy().astype(bool)
steps = base_mask.shape[0]
boundaries = np.linspace(0, steps, bins_per_modality + 1).round().astype(int)
slices = [(int(boundaries[k]), int(boundaries[k + 1])) for k in range(bins_per_modality)]
full = ensemble_forward(models, xs, mask, details=False)
ranked = full["logits"][0].argsort(descending=True)
target, other = int(ranked[0].item()), int(ranked[1].item())
full_margin = float((full["logits"][0, target] - full["logits"][0, other]).item())
rng = np.random.default_rng(seed)
candidates = [
[index for index, (left, right) in enumerate(slices) if base_mask[left:right, m].any()]
for m in range(3)
]
masks_to_score: list[np.ndarray] = []
row_specs: list[tuple[float, str, int]] = []
for rate in (0.1, 0.2, 0.3):
counts = [max(1, int(np.ceil(rate * len(indices)))) if indices else 0 for indices in candidates]
method_choices: list[tuple[str, list[list[int]]]] = []
top_choices: list[list[int]] = []
random_choices: list[list[int]] = []
for modality in range(3):
available = candidates[modality]
count = min(counts[modality], len(available))
score_order = sorted(available, key=lambda index: (-abs(contribution[modality, index]), index))
top_choices.append(score_order[:count])
random_choices.append(list(rng.choice(available, size=count, replace=False)) if count else [])
method_choices.append(("owen", top_choices))
method_choices.append(("matched_random", random_choices))
for method_name, choices in method_choices:
reps = 1 if method_name == "owen" else random_replicates
for rep in range(reps):
if method_name == "matched_random" and rep > 0:
choices = [
list(rng.choice(candidates[m], size=min(counts[m], len(candidates[m])), replace=False))
if counts[m]
else []
for m in range(3)
]
deleted = base_mask.copy()
retained = np.zeros_like(base_mask, dtype=bool)
for modality, selected in enumerate(choices):
for bin_index in selected:
left, right = slices[bin_index]
deleted[left:right, modality] = False
retained[left:right, modality] = base_mask[left:right, modality]
masks_to_score.extend((deleted, retained))
row_specs.extend(((rate, method_name, rep), (rate, method_name, rep)))
score_values = _margin_values(models, xs, np.stack(masks_to_score), target, other)
output: list[dict[str, Any]] = []
cursor = 0
grouped: dict[tuple[float, str], list[tuple[float, float]]] = {}
for rate, method_name, _rep in row_specs[::2]:
deletion_margin = float(score_values[cursor])
retention_margin = float(score_values[cursor + 1])
cursor += 2
grouped.setdefault((rate, method_name), []).append(
(full_margin - deletion_margin, retention_margin)
)
for (rate, method_name), values in sorted(grouped.items()):
arr = np.asarray(values, dtype=np.float64)
output.append(
{
"sample_id": sample_id,
"budget": rate,
"method": method_name,
"replicates": len(values),
"deletion_margin_drop_mean": float(arr[:, 0].mean()),
"retention_margin_mean": float(arr[:, 1].mean()),
"retention_margin_drop_mean": float((full_margin - arr[:, 1]).mean()),
"full_margin": full_margin,
}
)
return output
@@ -0,0 +1,86 @@
from __future__ import annotations
import argparse
import csv
from pathlib import Path
from typing import Any
import numpy as np
import torch
from ...data_paths import PROJECT_ROOT
from ...q2.deep_learning.q2.data import RobustStats
from ..run_experiments import _read_attachment4
from .evaluate import (
MODEL_SEEDS,
_attachment_predictions_and_explanations,
_attachment_split,
_load_ensemble,
)
from .owen import hierarchical_owen_one
SCALER_PATH = PROJECT_ROOT / "experiments" / "q2" / "unaligned_deep_two_b128" / "unaligned_50_robust_stats.npz"
DEFAULT_OUTPUT = PROJECT_ROOT / "output" / "q3" / "ati_ho"
def _write_csv(path: Path, rows: list[dict[str, Any]]) -> None:
if not rows:
raise ValueError(f"no rows to write: {path}")
path.parent.mkdir(parents=True, exist_ok=True)
with path.open("w", encoding="utf-8-sig", newline="") as stream:
writer = csv.DictWriter(stream, fieldnames=list(rows[0]))
writer.writeheader()
writer.writerows(rows)
def main() -> None:
parser = argparse.ArgumentParser(description="Regenerate the official Q3 Attachment 4 predictions and explanations.")
parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT)
parser.add_argument("--device", choices=("auto", "cpu", "cuda"), default="auto")
args = parser.parse_args()
if args.device == "cuda" and not torch.cuda.is_available():
parser.error("CUDA was requested but is not available")
device_name = "cuda" if args.device == "auto" and torch.cuda.is_available() else args.device
if device_name == "auto":
device_name = "cpu"
device = torch.device(device_name)
cases, _ = _read_attachment4("unaligned_50")
stats = RobustStats.load(SCALER_PATH)
attachment = _attachment_split(cases, stats)
dims = tuple(int(values.shape[-1]) for values in attachment.x)
models = _load_ensemble("A0", MODEL_SEEDS, dims, device)
predictions, explanations, _, _ = _attachment_predictions_and_explanations(
"A0", models, cases, attachment, device
)
local_rows: list[dict[str, Any]] = []
for index, case in enumerate(cases):
xs = tuple(torch.as_tensor(values[index:index + 1], dtype=torch.float32, device=device) for values in attachment.x)
mask = torch.as_tensor(attachment.mask[index:index + 1], dtype=torch.bool, device=device)
result = hierarchical_owen_one(
models, xs, mask, seed=20260926 + index, start_permutations=8, max_permutations=64
)
for modality, label in enumerate(("T", "A", "V")):
for bin_index, (left, right) in enumerate(result["bin_slices"]):
local_rows.append({
"case_id": case["case_id"],
"modality": label,
"relative_bin": bin_index,
"relative_position_start": left / 50.0,
"relative_position_end": right / 50.0,
"local_owen_margin_contribution": float(result["contribution"][modality, bin_index]),
"owen_standard_error": float(result["standard_error"][modality, bin_index]),
"permutations": result["permutations"],
"stopping_status": result["stopping_status"],
"physical_time_alignment": False,
})
_write_csv(args.output_dir / "attachment4_predictions.csv", predictions)
_write_csv(args.output_dir / "attachment4_explanations.csv", explanations)
_write_csv(args.output_dir / "attachment4_local_evidence.csv", local_rows)
print(f"Wrote {len(predictions)} predictions, {len(explanations)} explanations, and {len(local_rows)} local-evidence rows to {args.output_dir}")
if __name__ == "__main__":
main()
+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()
File diff suppressed because it is too large Load Diff