Train ATI-HO and finalize project outputs

This commit is contained in:
2026-09-26 16:05:44 +08:00
parent a86560da64
commit 9cdd604117
358 changed files with 10540 additions and 173 deletions
+33 -45
View File
@@ -1,67 +1,55 @@
# Q3:分层反事实证据归因
# Q3:ATI–HO 锚定交互与分层 Owen 归因
## 第一轮比较
当前 Q3 方案采用 ATI–HO。训练和评估入口位于 `q3/ati_ho/`,模型定义集中在 `model/ati_ho.py` 与 `model/ati_ho_config.py`。附件 4 只用于最终推理和解释,不参与训练、选型或性能指标计算。
第一轮直接复用 Q2 官方 unaligned_50 实验中的 EarlyConcat + BiGRU、MoFE-7 + MLP Router 权重和 train-only robust scaler。这样 E0/E1/E2 使用固定预测器,主要比较解释方式。
## 数据与输入
| 方案 | 预测器 | 解释 |
|---|---|---|
| E0 | EarlyConcat + BiGRU | 三模态精确 Shapley、配对交互、多尺度局部遮蔽 |
| E1 | MoFE-7 + MLP Router | Router 模态/位置权重;用反事实删除检验其是否 faithful |
| E2 | 与 E1 相同的 MoFE 检查点 | 三模态精确 Shapley、配对交互、多尺度局部遮蔽 |
将题目附件放在 `final/data/`,或设置 `FINAL_DATA_DIR` 指向包含官方附件目录的根路径。运行需要附件 2 的 `unaligned_50.pkl`、附件 4 的未对齐特征和视频,以及 Q2 训练集拟合的 scaler:`experiments/q2/unaligned_deep_two_b128/unaligned_50_robust_stats.npz`。
E1 与 E2 的预测逐样本相同。E1 的路由权重只描述融合机制,只有通过删除检验后才能说明它在这些样本上是否与预测行为一致。
训练、验证、测试按官方来源视频组隔离。所有模态经统一 Q1 adapter 投影到 50 个 Relative-Progress 槽;这统一的是序列内部进度,不代表物理时间同步。Scaler 仅使用训练集统计量。
## 运行
## 训练
从项目根目录执行,附件目录由 FINAL_DATA_DIR 指定。它应包含 附件2-数据集特征文件/unaligned_50.pkl 和 附件4-可解释专项视频样本与特征文件/。
从项目根目录执行。首次完整训练依次完成 Stage I 和 Stage II:
~~~bash
export FINAL_DATA_DIR="/path/to/task-data"
python -m final.q3.run_experiments \
--output-dir final/output/q3/first_round \
--device auto
~~~
```bash
export FINAL_DATA_DIR="/path/to/E题数据"
python -m final.q3.ati_ho.train --phase all --device auto
```
检查点与 scaler 默认读取 final/experiments/q2/unaligned_deep_two_b128/。完整运行会计算官方验证集预测及误差归因。--skip-validation 跳过这一步;--no-frames 跳过候选帧抽取。每轮实验使用新的空输出目录;若上次运行中断且留下部分文件,可对该目录增加 --resume 重新生成结果。
Stage I 使用 seed 42 训练 A0/A1/A2/A3 与未锚定诊断 D0,执行结构和 8 联盟 Shapley 审计并确定 provisional candidate。Stage II 使用 seeds 42、3407、2026 训练 EarlyConcat + BiGRU、MoFE-7 + MLP Router、初选 ATI 方案和两项关键消融。若需分阶段运行,可用 `--phase stage1` 或 `--phase stage2`;Stage II 需要 Stage I 的选择记录。默认最多 12 个 epoch,按锁定验证情景损失早停。已存在的检查点会复用;需要重训时显式加 `--force`。
## 方法定义
训练输出写入 `experiments/q3/ati_ho/`,含检查点、训练历史、逐情景验证指标、阶段状态、scaler 和 adapter 元数据。附件 4 不会在训练程序中加载。
令三个模态为 (M={T,A,V})。对每个附件 4 样本完整计算 8 个 coalition。分类价值函数使用完整输入预测类别的 logit,并在所有 coalition 上固定该类别;回归价值函数使用模型情感强度输出。绝对 Shapley 贡献除以三模态绝对贡献之和,得到模态作用比例;带符号值保留支持或反对预测的方向。配对交互采用标准 Shapley interaction index 系数。
## 评估与附件 4 输出
局部证据对每个可见模态位置计算窗口宽度 (win{1,3,5}) 的遮蔽前后差值,三个尺度等权平均。每个模态选取局部贡献绝对值最高的 10% 位置,合并相邻位置,并输出代表证据段。
完成两阶段训练后运行:
Faithfulness 使用固定预测类别 logit。Comprehensiveness 比较完整输入与删去高排名证据后的分数;sufficiency 比较完整输入与只保留高排名证据后的分数;deletion AUC 汇总删除 0% 到 70% 的分数下降。另记录宽度 1/3/5 局部图之间的 Spearman 相关作为尺度稳定性诊断。MoFE Router 与精确 Shapley 按样本比较 Spearman 排序相关和主导模态一致率。
```bash
export FINAL_DATA_DIR="/path/to/E题数据"
python -m final.q3.ati_ho.evaluate --device auto
```
## 输出
评估按 A0/A1/A2 三 seed 固定四场景验证损失的均值确定最终 ATI 方案,随后生成配对来源视频组 Bootstrap、最终模型结构审计、全验证集解析/精确 Shapley 审计、Owen 稳定性和删除/保留诊断,以及模型成本表和结果图。评估不基于附件 4 标签计算指标。
- attachment4_predictions.csv:E0、E1、E2 对 20 个样本的完整预测与类别概率。
- attachment4_modal_shapley.csv:分类与强度的有符号贡献、绝对比例和完备性残差。
- attachment4_pairwise_interactions.csv:Text–Audio、Text–Vision、Audio–Vision 交互。
- attachment4_local_evidence.csv:E0/E2 的 1/3/5-bin 遮蔽差值和来源行。
- attachment4_router_profiles.csv、attachment4_router_local_evidence.csv:MoFE 的专家权重和逐位置 Router utility。
- attachment4_evidence_segments.csv:稀疏代表证据段、来源行、转写文本和候选视频时间。
- faithfulness_by_sample.csv、q3_method_comparison.csv:逐样本与方案汇总的删除/保留检查。
- validation_predictions.csv、validation_errors.csv、validation_error_attribution.csv:官方验证集指标、误差样本和分类 margin 归因。
- explanation_cards/、typical_explanation_card.md、evidence_profiles/:逐样本解释卡、代表卡和 3×50 热图。
- evidence_frames/:从原始视频抽取的候选帧。
- run_manifest.json:检查点、scaler、输入模式、公式、运行边界与产物清单。
题目输出放在 `output/q3/ati_ho/`:
## Router 热力图样例
- `attachment4_predictions.csv`:20 个样本的类别、情感强度、类别概率和各模态可见槽数量。
- `attachment4_explanations.csv`:五个参数的主效应、pairwise 参数项、解析与精确分类 Shapley,以及精确强度 Shapley。
- `attachment4_local_evidence.csv`:每个样本按模态和 10 个相对进度片段排列的 Owen logit-margin 贡献与标准误。
- `attachment4_prediction_manifest.json`:adapter/scaler、类别顺序、输入与检查点 SHA-256、行数和无标签推理声明。
- `README.md`:上述交付件说明与坐标限制。
在完成 Q3 第一轮后,还可以绘制附件 4 样本 02、03 的输入遮蔽与 Router 权重对照图。每个样本分别展示原始完整输入和受控残缺输入:样本 02 遮蔽 30% 文本与音频,样本 03 遮蔽 30% 视觉。蓝色斜线只标输入中被遮蔽的位置;紧接着一行把七个 expert(T、A、V、TA、TV、AV、TAV)横向排列,色块/数字显示样本平均 Router 权重,下方窄条显示 50 个 bin 上的 α[t,e]。
完整实验产物在 `experiments/q3/ati_ho/results/ati_ho/`,包括 `ATI_HO_RESULTS.md`、论文式报告、验收摘要、CSV 审计表、图和运行清单。验证集逐样本结果及训练权重也保存在 `experiments/q3/ati_ho/`。实验归因文件可供复核,不属于精简的题目输出目录。
~~~bash
export FINAL_DATA_DIR="/path/to/task-data"
python -m final.q3.plot_router_heatmaps --output-dir final/output/q3
~~~
## 模型定义与解释边界
这两个残缺样例是人为遮蔽的对照,不是附件 4 的原生缺失数据。文本 token 按序列顺序显示,蓝色斜线标出对应遮蔽词段;音频波形和视频帧按相对进程显示蓝色遮蔽区。每个 bin 的七个 expert 权重在可用专家集合内归一化,未满足模态条件的 expert 权重为 0;色条使用原始 0–1 数值并通过平方根归一化提高低权重的可读性。样例顶部另列 T/A/V 的总体 Router exposure share。视频帧和音轨只按归一化进程投影,不能当作精确词/帧同步。Router 权重反映融合路由,不是预测贡献或情绪因果解释。
模型输出由 3 个居中的类别 logit、负向强度参数和正向强度参数组成。ATI 主效应以空输入前向作零锚定;候选 pairwise 分支只读取对应的两种模态并对缺失模态基线作锚定。A0 是只有主效应的基线,A1 增加秩 4 的 pairwise 分支,A2 再加入一层 4 头交叉注意力,A3 增加可见性掩码去噪辅助目标;D0 是未锚定诊断。未加入三阶项。
图表与原始数据输出为 `mofe_router_heatmap_examples.png`、`mofe_router_heatmap_examples.pdf`、`mofe_router_heatmap_scores.csv`(逐 bin 七个 α 分数及模态可见掩码)、`mofe_router_heatmap_summary.csv` 和 `mofe_router_heatmap_manifest.json`。
分类解释固定完整输入的预测类别与次高类别,以 logit margin 为目标;解析 Shapley 与完整枚举 8 个模态联盟的结果比较。情感强度经过类别选择和 sigmoid 解码,是非线性输出,因此单独对 8 个联盟精确枚举强度 Shapley。局部 Owen 将三模态作外层组、每模态 10 个五槽片段作内层组,按 8、16、32、64 个随机排列检查稳定性。
## 回溯边界
位置均为相对进度槽,不是秒数。遮挡测试描述模型对输入可见性的响应,不是人类解释准确率或情绪因果效应。附件 4 没有真实标签,所以只输出预测与模型解释,不声称其预测精度。
附件 4 的文本、音频和视觉序列为未对齐特征,没有逐词、逐音频帧或逐视频帧的真实时间戳。输出会从 adapter 的稀疏投影权重记录来源特征行和归一化进程。视频候选秒数由归一化进程乘视频时长估算,只供人工回看;它不是真实物理时间对齐。文本 token 需要本地缓存 google-bert/bert-base-uncased tokenizer;缓存不存在时,解释卡仍保留完整转写和来源行。
## 早期 Q3 文件
这些解释衡量的是当前预测器对输入遮蔽的响应,不是现实情绪成因。虽然模型在 Q2 训练时见过模态遮蔽,局部孤立遮蔽和只保留 10% 的输入仍可能偏离训练分布;相关数值按诊断结果报告,不称为解释准确率或因果效应。
此前 MoFE 复用检查点的第一轮解释和 Router 可视化保存在 `experiments/q3/legacy_mofe_first_round/`,用于保留历史记录。当前 `output/q3/` 只放本题采用的 ATI–HO 交付文件。
+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
+54
View File
@@ -0,0 +1,54 @@
from __future__ import annotations
import json
import torch
from .attribution import exact_shapley_audit
from .audit import structural_audit
from ...model.ati_ho import ATIHOModel, task_loss
from ...model.ati_ho_config import CONFIGS
def main() -> None:
torch.manual_seed(17)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
dims = (768, 74, 35)
xs = tuple(torch.randn(2, 50, dim, device=device) for dim in dims)
masks = torch.ones(2, 50, 3, dtype=torch.bool, device=device)
masks[0, 10:20, 1] = False
masks[1, 30:42, 2] = False
y_cls = torch.tensor([0, 2], device=device)
y_reg = torch.tensor([-1.5, 2.0], device=device)
reports = {}
for name in ("A0", "A1", "A2", "A3", "D0"):
model = ATIHOModel(dims, CONFIGS[name]).to(device)
result = model(xs, masks)
assert result["params"].shape == (2, 5)
assert torch.isfinite(result["params"]).all()
loss, _ = task_loss(
result,
y_cls,
y_reg,
lambda_interaction=CONFIGS[name].lambda_interaction,
lambda_mask=CONFIGS[name].lambda_mask,
mask_target=masks,
)
loss.backward()
assert torch.isfinite(loss)
report = structural_audit(model, xs, masks)
assert report["main_and_additivity_pass"]
assert report["nonfinite_value_count"] == 0
if CONFIGS[name].anchored:
assert report["pair_anchor_pass"]
reports[name] = report
model = ATIHOModel(dims, CONFIGS["A2"]).to(device).eval()
audit = exact_shapley_audit([model], xs, masks, batch_size=8)
assert audit["class_pass"].all(), audit["class_abs_error"]
reports["analytic_vs_exact_shapley_max_abs"] = float(audit["class_abs_error"].max())
print(json.dumps(reports, indent=2))
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()