diff --git a/final/README.md b/final/README.md index 5589310..6539934 100644 --- a/final/README.md +++ b/final/README.md @@ -11,12 +11,12 @@ | `q1/` | Q1 原生特征提取、物理时间特征包构建、五折对照和可视化 | | `q2/math/` | 数学方案 C0–C7、训练、缺失控制评估及附件 3 预测 | | `q2/deep_learning/q2/` | EarlyConcat + BiGRU、MoFE-7 + MLP Router 及训练协议 | -| `q3/` | Q3 验证、附件 4 全量预测和逐样本解释卡片 | +| `q3/` | ATI–HO Q3 训练、结构审计、归因与附件 4 推理 | | `output/q1/` | 已整理的 Q1 100 个特征文件、对齐审计和已有结果 | | `output/q2/` | 已整理的 Q2 未对齐数据模型对比结果表 | -| `output/q3/` | Q3 运行后生成的预测、解释和验证结果 | +| `output/q3/` | ATI–HO 附件 4 题目输出;验证结果和实验审计放在 `experiments/q3/` | | `experiments/q2/` | 本项目最新未对齐 Q2 对比运行、权重和审计结果 | -| `REPORTS.md` | Q1、Q2 实验结果和方法说明 | +| `REPORTS.md` | Q1、Q2、Q3 实验结果和方法说明 | | `data_paths.py` | 官方附件的默认位置与外部数据根目录设置 | ## 准备官方附件 @@ -137,6 +137,19 @@ python -m final.q2.deep_learning.q2.train_math_protocol \ 数学方案默认会用附件 3 未对齐样本生成 30 条 `attachment3_predictions.csv` 和 `attachment3_audit.csv`,预测文件包含极性、情感强度及类别概率。如只检查训练流程,可以增加 `--skip-attachment3`。深度学习运行会输出两模型的官方测试指标、验证缺失情景、AURC-MAE、Bootstrap 区间、权重和运行清单。为满足总附件大小限制,随项目提供的 Q2 归档保留逐情景指标和模型参数;逐行遮蔽/测试门控审计及官方测试单样本明细可由完整重训重新生成,未放入紧凑归档。 +若已存在数学分支检查点、只需重新生成附件 3 预测而不重训,可运行: + +```bash +export FINAL_DATA_DIR="/path/to/E题数据" +python -m final.q2.math.predict_attachment3 \ + --input-version unaligned_50 \ + --results-dir final/experiments/q2/unaligned_math_all_b128 \ + --output-dir final/output/q2 \ + --device auto +``` + +该入口读取保存的模型、校准与预处理参数,输出 `attachment3_predictions.csv`、`attachment3_audit.csv` 和 `attachment3_prediction_manifest.json`;不会覆盖模型检查点目录中的运行清单。 + 把两次运行的结果重新汇总到统一对比表: ```bash @@ -148,47 +161,36 @@ python final/compare_unaligned_q2.py \ 结果文件包括 `comparison_validation.csv`(全模型验证对比)、`comparison_test.csv`(预先选出的数学模型及两种深度模型测试结果)、`comparison_aurc.csv`(四种缺失模式的 AURC-MAE)。指标定义、已有对比数值和边界说明见 [REPORTS.md](REPORTS.md)。 -## Q3:分层反事实证据归因 +## Q3:ATI–HO 训练、评估与附件 4 预测 -Q3 复用 Q2 已训练的 EarlyConcat + BiGRU、MoFE-7 + MLP Router 检查点和同一套训练集 robust scaler,不重新拟合模型。第一轮含三个解释方案:E0 对 EarlyConcat 做精确三模态 Shapley 与局部遮蔽;E1 把 MoFE Router 当作待检验的内部信号;E2 对同一个 MoFE 检查点做精确 Shapley、交互和局部遮蔽。E1/E2 的预测完全相同,比较的是解释方式。 +Q3 的当前方案是 ATI–HO。模型定义集中在 `model/ati_ho.py` 和 `model/ati_ho_config.py`,训练、结构审计、Shapley/Owen 归因与附件 4 推理入口位于 `q3/ati_ho/`。附件 4 只用于最终预测和解释,不参与训练、选型或指标计算。 -在项目根目录设置附件位置后运行。项目默认查找 final/data/;也可以将 FINAL_DATA_DIR 指向包含附件 2、附件 4 文件夹的根目录: +训练与验证读取官方 `unaligned_50.pkl`,使用统一 Q1 adapter 投影到 50 个 Relative-Progress 槽,并复用只由 Q2 训练集拟合的 robust scaler。训练/验证/测试按来源视频组隔离;相对进度不是物理时间同步。 -~~~powershell -$env:FINAL_DATA_DIR = "D:\task-data" -python -m final.q3.run_experiments --output-dir final/output/q3/first_round --device auto -~~~ +在项目根目录设置官方附件位置后,运行两阶段训练: -~~~bash -export FINAL_DATA_DIR="/data/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 +``` -完整运行使用官方验证集生成预测指标、误差归因和分类 margin 的 Shapley。只想快速检查附件 4 可增加 --skip-validation;不抽取候选视频帧可增加 --no-frames。每次运行请给一个新的空输出目录。 +Stage I 以 seed 42 训练 A0/A1/A2/A3 和未锚定诊断 D0,执行结构与 8 联盟 Shapley 检查。Stage II 以 seeds 42、3407、2026 训练 EarlyConcat + BiGRU、MoFE-7 + MLP Router、初选 ATI 方案和两项消融。检查点和运行记录写入 `final/experiments/q3/ati_ho/`。重训已有模型时显式加 `--force`。 -主要产物: +训练完成后生成验证审计和附件 4 文件: -- attachment4_predictions.csv:20 个附件 4 样本在 E0/E1/E2 下的预测和类别概率。 -- attachment4_modal_shapley.csv、attachment4_pairwise_interactions.csv:分类 logit 与强度回归的模态贡献、绝对贡献比例、配对交互和 Shapley 完备性残差。 -- attachment4_local_evidence.csv、attachment4_router_local_evidence.csv:1/3/5 个相对进程 bin 的局部遮蔽响应与 MoFE Router 位置分数。 -- attachment4_evidence_segments.csv:每个方案的代表性文本、音频和视觉证据段;evidence_frames/ 中保存由进度比例估算位置抽取的候选帧。 -- faithfulness_by_sample.csv、q3_method_comparison.csv:删除/保留检验、0–70% 删除曲线、尺度稳定性和 Router–Shapley 一致性汇总。 -- validation_predictions.csv、validation_errors.csv、validation_error_attribution.csv:官方验证集预测、误差样本和分类 margin 归因。 -- explanation_cards/、typical_explanation_card.md、evidence_profiles/:逐样本解释卡、代表卡和 3×50 局部证据图。 +```bash +export FINAL_DATA_DIR="/path/to/E题数据" +python -m final.q3.ati_ho.evaluate --device auto +``` -想先查看文本、音频波形、视频帧中的遮蔽位置,再查看对应时间 bin 上七个 MoFE expert 的 Router 权重热力图,可运行: +评估按 A0/A1/A2 三 seed 固定验证情景损失选择最终 ATI 方案,计算来源视频组配对 Bootstrap、验证集 Shapley、结构审计、Owen 稳定性、删除/保留诊断和计算成本。题目输出放在 `final/output/q3/ati_ho/`: -~~~powershell -python -m final.q3.plot_router_heatmaps --output-dir final/output/q3 -~~~ +- `attachment4_predictions.csv`:20 个样本的预测类别、强度和类别概率。 +- `attachment4_explanations.csv`:五个参数的主效应、pairwise 参数项、分类 Shapley 和强度 Shapley。 +- `attachment4_local_evidence.csv`:分模态、分相对进度片段的 Owen 贡献与标准误。 +- `attachment4_prediction_manifest.json`:数据、adapter/scaler、模型权重哈希和无标签推理信息。 -它输出附件 4 样本 02、03 的完整输入和受控残缺对照图:蓝色斜线表示输入被遮蔽,下面按 T、A、V、TA、TV、AV、TAV 顺序显示 expert 的 α[t,e]。详细定义与限制见 [Q3 实验说明](q3/README.md#router-热力图样例),实验图也已收录在 [REPORTS.md](REPORTS.md)。 - -三种输入模态只有 8 个 coalition,脚本完整枚举而非近似 SHAP。分类解释固定使用完整输入的预测类别 logit,回归解释使用情感强度输出。局部证据是单模态窗口被遮蔽前后的 logit 差;正值表示该窗口支持当前预测,负值表示该窗口反对当前预测。Router 权重只描述路由机制,不能直接解释成预测贡献。 - -附件 4 特征行没有物理时间戳。文本可以回溯到 tokenizer token 和来源行;音频/视觉证据保留来源行与归一化进程。若视频可读,候选帧位置由相对进程乘视频时长估算,供人工回看,不表示特征已经物理时间对齐。删除/保留检验使用模型训练时见过的 mask 接口,但孤立稀疏遮蔽仍可能偏离训练分布;报告会把 10% 保留结果标为诊断,主要删除曲线范围限制在最多删除 70%。 - -细节和解释边界见 Q3 实验说明(q3/README.md)与 REPORTS.md。 +完整指标和实验审计位于 `final/experiments/q3/ati_ho/results/ati_ho/`,主要报告为 `ATI_HO_RESULTS.md`、`ATI_HO_PAPER.md` 和 `EXECUTIVE_SUMMARY.md`。方法细节、训练协议与文件边界见 [Q3 说明](q3/README.md);Q1/Q2/当前 Q3 的合并结果写在 [REPORTS.md](REPORTS.md)。 ## 模型与对齐约定 diff --git a/final/REPORTS.md b/final/REPORTS.md index 5e73ba9..0e1db99 100644 --- a/final/REPORTS.md +++ b/final/REPORTS.md @@ -1,4 +1,4 @@ -# Q1/Q2 实验结果与复现摘要 +# Q1/Q2/Q3 实验结果与复现摘要 本报告汇总 `final/` 内可复核的 Q1 物理时间对齐结果和 Q2 官方未对齐数据比较。模型代码、训练命令、数据目录和依赖说明见 [README.md](README.md)。数值结果链接均指向本项目内的 CSV、JSON 或实验审计文件。 @@ -96,45 +96,63 @@ Q1 adapter 使用带媒体时戳的原生特征进行物理时间聚合;Q2 未 完整精度及审计文件:[14 个版本验证指标](output/q2/comparison_validation.csv) · [3 个正式测试结果](output/q2/comparison_test.csv) · [全部版本 AURC-MAE](output/q2/comparison_aurc.csv) · [数学运行清单](experiments/q2/unaligned_math_all_b128/run_manifest.json) · [深度学习运行清单](experiments/q2/unaligned_deep_two_b128/run_manifest.json) · [数学 42 情景明细(gzip CSV)](experiments/q2/unaligned_math_all_b128/controlled_missingness.csv.gz) · [深度学习 42 情景明细](experiments/q2/unaligned_deep_two_b128/controlled_metrics_by_scenario.csv) · [深度学习测试配对 Bootstrap](experiments/q2/unaligned_deep_two_b128/official_test_paired_bootstrap.csv)。 -## Q3:分层反事实证据归因(第一轮) +## Q3:ATI–HO 锚定交互与分层 Owen 归因 -Q3 将 Q1 的统一输入/回溯接口、Q2 的缺失鲁棒预测器与附件 4 原始视频接起来。附件 4 的特征按统一 adapter 的 Relative-Progress 模式转换;预测器直接复用 Q2 seed 20260924 的 EarlyConcat + BiGRU、MoFE-7 + MLP Router 检查点和仅由训练集拟合的 robust scaler,本轮不重新训练。解释覆盖附件 4 的 20 个样本;官方验证集另有 728 条样本用于复核预测指标并输出误差归因。 +本轮按用户更正后的 ATI–HO 方案完成 Q3 训练与评估。输入使用官方 `unaligned_50.pkl`,经统一 Q1 adapter 转成 50 个 Relative-Progress 槽,复用仅由训练集拟合的 Q2 robust scaler。训练/验证/测试为 3,395/728/727 条,来源视频组为 1,528/239/381,组间无重叠。相对进度统一序列顺序,不代表物理时间同步。 -三模态只有 8 个输入组合,因此 E0/E2 对文本、音频、视觉 coalition 完整枚举,分别计算预测类别 logit 与强度输出的精确 Shapley 贡献及成对交互。局部证据由窗口宽度 1/3/5 的反事实遮蔽得到,并回溯到特征来源行、文本片段及候选视频帧。E1 展示 MoFE Router 权重,并用同一预测器做删除/保留检验;Router 权重作为内部路由信号分析,不直接当作预测贡献。 +### 训练与选型 -### 预测器验证表现 +Stage I 用 seed 42 训练 A0/A1/A2/A3 和未锚定诊断 D0;Stage II 对 EarlyConcat + BiGRU、MoFE-7 + MLP Router、Stage I 初选方案及两项关键消融使用 seeds 42、3407、2026。按锁定验证集四个固定遮蔽场景任务损失的三 seed 均值选择最终模型。Stage I 初选 A2;Stage II 三 seed 选出 **A0(仅锚定主效应)**: -下表为复用检查点在官方验证集自然缺失输入上的结果。E1 与 E2 使用同一个 MoFE 检查点,因此预测指标相同。 +| ATI 方案 | 固定场景验证损失(均值 ± 标准差) | +|---|---:| +| A0 | **0.86580 ± 0.00920** | +| A1(秩 4 pairwise) | 0.86728 ± 0.01282 | +| A2(锚定交叉注意力) | 0.86859 ± 0.01014 | -| 方案 / 预测器 | Accuracy ↑ | Macro-F1 ↑ | MAE ↓ | RMSE ↓ | Pearson ↑ | +结果不支持在最终模型中保留更复杂的 pairwise 分支;A1/A2 仍作为消融完整留档。附件 4 标签未参与训练、选型和结果计算。 + +### 官方验证集性能 + +表中数值为三 seed 均值 ± 标准差;分类保留 negative/neutral/positive 三类。 + +| 模型 | Seeds | Accuracy ↑ | Macro-F1 ↑ | 强度 MAE ↓ | Pearson ↑ | |---|---:|---:|---:|---:|---:| -| E0 EarlyConcat + BiGRU | **0.6236** | **0.5779** | 0.6614 | 0.8731 | **0.6092** | -| E1/E2 MoFE-7 + MLP Router | 0.6140 | 0.5661 | **0.6426** | **0.8499** | 0.6014 | +| EarlyConcat + BiGRU | 3 | 0.632 ± 0.010 | 0.594 ± 0.018 | **0.629 ± 0.008** | **0.621 ± 0.003** | +| MoFE-7 + MLP Router | 3 | 0.624 ± 0.010 | **0.598 ± 0.009** | 0.639 ± 0.004 | 0.614 ± 0.004 | +| ATI–HO A0 | 3 | **0.633 ± 0.002** | 0.593 ± 0.010 | 0.740 ± 0.015 | 0.565 ± 0.003 | -EarlyConcat 的 Accuracy/Macro-F1 点估计较高,MoFE 的 MAE/RMSE 较低。这里复核的是已训练检查点,不是 Q3 新模型训练结果;指标差别不代表显著差异。 +A0 的 Accuracy 点估计略高,但 Macro-F1 接近两条基线,强度 MAE 和 Pearson 明显较弱。因此“最终选中”只表示它在预先规定的组合验证损失上均值最低,不表示它在所有指标上优于基线。 -### 附件 4 解释行为检查 +按验证视频组进行 1,000 次配对 Bootstrap。正值表示 ATI–HO 更好;MAE 差值定义为“基线 MAE − ATI–HO MAE”。所有区间均为 95%: -下表数值是 20 个无人工解释标签样本的均值。Comprehensiveness 是删除高排名局部证据后预测类别 logit 的下降,越大说明当前预测对所选证据越敏感;Sufficiency 是只保留高排名证据后的完整输入 logit 绝对差,越小越接近完整输入。Deletion AUC 汇总删除 0%–70% 的 logit 下降。 +| 对比 | 指标 | 差值 | Bootstrap 区间 | +|---|---|---:|---:| +| A0 − EarlyConcat | Accuracy | +0.0027 | [-0.0157, 0.0218] | +| A0 − EarlyConcat | Macro-F1 | +0.0044 | [-0.0192, 0.0287] | +| EarlyConcat MAE − A0 MAE | MAE | -0.1211 | [-0.1553, -0.0892] | +| A0 − EarlyConcat | Pearson | -0.0619 | [-0.0922, -0.0338] | +| A0 − MoFE | Accuracy | +0.0000 | [-0.0177, 0.0182] | +| A0 − MoFE | Macro-F1 | -0.0118 | [-0.0329, 0.0079] | +| MoFE MAE − A0 MAE | MAE | -0.1158 | [-0.1488, -0.0835] | +| A0 − MoFE | Pearson | -0.0571 | [-0.0863, -0.0268] | -| 方案 | 删除 Top 10% ↑ | 删除 Top 30% ↑ | 仅保留 Top 10% 的绝对差 ↓ | 仅保留 Top 30% 的绝对差 ↓ | Deletion AUC ↑ | 尺度稳定性 Spearman ↑ | -|---|---:|---:|---:|---:|---:|---:| -| E0 EarlyConcat:Shapley + 局部遮蔽 | 0.4599 | **1.1218** | 0.5246 | 0.2542 | **0.9157** | 0.9047 | -| E1 MoFE:Router + 反事实检验 | 0.2100 | 0.8702 | 0.4954 | 0.2500 | 0.7925 | — | -| E2 MoFE:Shapley + 局部遮蔽 | **0.3515** | 0.9818 | **0.4116** | **0.1950** | 0.8371 | 0.8823 | +分类差异区间跨零;连续强度指标显示 A0 在当前验证组上弱于两条基线。Bootstrap 描述固定训练检查点在验证视频组上的抽样不确定性,不代表训练 seed 不确定性。 -在同一 MoFE 检查点上,E2 的删除分数和 deletion AUC 高于 E1,保留证据后的绝对差更低;这说明 E2 的局部 Shapley/遮蔽排序在本批样本上比 Router utility 更贴合预测器的遮蔽响应。两者的平均模态排序 Spearman 为 0.725,主导模态一致率为 17/20(85%),Router 与反事实贡献存在一定对应,但并不相同。E0/E2 的尺度稳定性是窗口宽度 1/3/5 得分排序相关的样本均值;E1 不是多尺度遮蔽解释,故不计算该值。 +### 结构与解释审计 -Shapley 完备性检查覆盖分类与强度两种输出,最大绝对残差为 `4.44×10⁻¹⁶`;E1/E2 的 20 个样本预测类别、强度和类别概率逐项一致。上述 faithfulness 数值是模型响应诊断,不是人工解释准确率;删除/保留输入也可能偏离训练时的掩码分布,尤其是仅保留 10% 的情景。附件 4 没有词、音频帧或视频帧的可信时间戳;输出的视频秒数由归一化进程乘视频时长估算,只作人工回看定位,不代表物理时间对齐。遮蔽贡献也不证明现实情绪成因。 +ATI 参数向量为三个居中类别 logit 与负/正条件强度参数;A0 的输出按锚定主效应加和重构,训练后最大重构残差为 0。A0–A3 的锚定结构检查通过;D0 未锚定诊断用于检查缺失模态基线泄漏。 -完整结果:[Q3 方案汇总](output/q3/first_round/q3_method_comparison.csv) · [运行清单与完备性审计](output/q3/first_round/run_manifest.json) · [附件 4 预测](output/q3/first_round/attachment4_predictions.csv) · [模态 Shapley 贡献](output/q3/first_round/attachment4_modal_shapley.csv) · [跨模态交互](output/q3/first_round/attachment4_pairwise_interactions.csv) · [局部证据](output/q3/first_round/attachment4_local_evidence.csv) · [逐样本 faithfulness](output/q3/first_round/faithfulness_by_sample.csv) · [验证集指标](output/q3/first_round/validation_metrics.json) · [验证集错误归因](output/q3/first_round/validation_error_attribution.csv) · [典型解释卡](output/q3/first_round/typical_explanation_card.md)。全部 20 个样本解释卡、证据图、候选帧与 CSV 保存在 [`output/q3/first_round/`](output/q3/first_round/)。 +最终 A0 在 728 条验证样本上的解析分类 Shapley 与 8 个模态联盟精确枚举通过率为 100%,最大绝对误差 6.81×10⁻⁷(容差 1e-6 绝对误差加 1e-5 相对误差)。Attachment 4 的 20 个无标签样本通过率也为 100%。类别选择和 sigmoid 解码后的强度输出为非线性量,因此强度贡献直接用 8 联盟精确枚举,没有套用参数线性分解。 -### MoFE Router 热力图示例 +Attachment 4 的 Hierarchical Owen 归因按三模态外层分组、每模态 10 个相对进度片段内层分组,从 8 次随机排列开始并依据稳定性增加至最多 64 次。20 个样本的最大守恒残差为 4.41×10⁻⁶。删除/保留诊断采用 10%/20%/30% 预算,并以同模态、同片段数随机选择作对照;这些数值是模型遮挡响应,不是人工解释准确率或因果效应。训练 seed 与 1% 输入扰动稳定性分别记录在 [稳定性审计表](experiments/q3/ati_ho/results/ati_ho/stability_results.csv)。 -下图选取附件 4 样本 02(完整输入时 MoFE 预测中性)和样本 03(预测负向)。每个样本各展示完整输入与一组受控遮蔽:样本 02 遮蔽文本 bins 6–20 和音频 bins 19–33,样本 03 遮蔽视觉 bins 19–33。完整样本三种模态各有 50/50 个可见 bin;残缺输入是人为构造的对照,不是附件 4 的原生缺失记录。上方文本、波形和视频帧只用蓝色斜线标记被遮蔽区;下方将 T、A、V、TA、TV、AV、TAV 七个 expert 横向并排,色块/数值表示该样本的平均 Router 权重,下方窄条显示 50 个相对进程 bin 上的 α[t,e]。每个 bin 的权重在可用 expert 中归一化;非可用 expert 权重为 0。色条显示原始 0–1 分数并用平方根归一化增强低权重可见度。顶部 T/A/V 数值是跨位置的 Router exposure share。 +本轮 14/20 个附件 4 样本在最多 64 次排列内满足 Owen 停止条件,平均 45.6 次。30% 删除时,ATI–HO 选中的局部片段使固定 logit margin 平均下降 0.433;同模态、同片段数随机片段下降 0.112。训练 seed 归因的平均相关为 0.252、top-5 Jaccard 为 0.380,说明细粒度位置排序存在 seed 波动;1% 输入扰动下的五个诊断样本保持了相同 top-5 证据。A0 每 seed 有 127,020 个可训练参数,三 seed 集成单样本推理约 4.842 ms;8 联盟 Shapley 约 0.044 秒/样本,Owen 约 0.131 秒/样本(RTX 5070 Ti)。这些时间只表示本轮运行环境。 -样本 02 的 T/A/V exposure share 从完整输入的 35.8%/28.9%/35.2% 变为遮蔽文本与音频后的 29.9%/25.9%/44.1%,分类预测从中性变为正向。样本 03 从 35.9%/29.2%/34.9% 变为遮蔽视觉后的 39.6%/34.2%/26.2%,分类仍为负向。这些是两个受控样本的 Router 路由变化示例,不表示缺失模态的真实情绪贡献;α[t,e] 也不是反事实预测贡献。 +### 题目输出和完整材料 -![附件 4 完整与受控残缺样本的 MoFE Router 热力图](output/q3/mofe_router_heatmap_examples.png) +题目输出目录 `output/q3/ati_ho/` 仅放附件 4 的预测与解释交付件:`attachment4_predictions.csv`(20 行类别、强度及概率)、`attachment4_explanations.csv`(参数分解及 Shapley)、`attachment4_local_evidence.csv`(Owen 片段贡献)、`attachment4_prediction_manifest.json`(输入/模型哈希和审计范围)。附件 4 没有标签,不报告其准确率或误差。 -[PNG 高清图](output/q3/mofe_router_heatmap_examples.png) · [PDF 矢量版](output/q3/mofe_router_heatmap_examples.pdf) · [逐位置 Router 分数](output/q3/mofe_router_heatmap_scores.csv) · [样本与条件摘要](output/q3/mofe_router_heatmap_summary.csv) · [热力图运行清单](output/q3/mofe_router_heatmap_manifest.json) · [可复现绘图脚本](q3/plot_router_heatmaps.py)。 +工作区 `experiments/q3/ati_ho/` 保留完整候选与基线权重,以及逐 seed 指标、验证集 Shapley、Bootstrap、结构、Owen、fidelity、稳定性和复杂度结果。精简提交包包含最终 A0 三 seed 权重和完整结果表。主要文档:[实验结果](experiments/q3/ati_ho/results/ati_ho/ATI_HO_RESULTS.md) · [论文式报告](experiments/q3/ati_ho/results/ati_ho/ATI_HO_PAPER.md) · [验收摘要](experiments/q3/ati_ho/results/ati_ho/EXECUTIVE_SUMMARY.md) · [运行清单](experiments/q3/ati_ho/results/ati_ho/run_manifest.json) · [三 seed 指标表](experiments/q3/ati_ho/results/ati_ho/main_results.csv) · [视频组 Bootstrap](experiments/q3/ati_ho/results/ati_ho/bootstrap_results.csv) · [Owen 与 fidelity 明细](experiments/q3/ati_ho/results/ati_ho/owen_audit.csv) · [稳定性审计](experiments/q3/ati_ho/results/ati_ho/stability_results.csv) · [附件 4 预测](output/q3/ati_ho/attachment4_predictions.csv) · [附件 4 局部证据](output/q3/ati_ho/attachment4_local_evidence.csv)。 + +模型定义在 `model/ati_ho.py` 与 `model/ati_ho_config.py`,训练和评估命令见 [Q3 运行说明](q3/ati_ho/README.md)。此前 MoFE 第一轮解释与 Router 热力图已从题目输出目录移入工作区历史记录;当前 Q3 结果以 ATI–HO 为准。 diff --git a/final/experiments/q3/ati_ho/final_selection.json b/final/experiments/q3/ati_ho/final_selection.json new file mode 100644 index 0000000..fc7aa74 --- /dev/null +++ b/final/experiments/q3/ati_ho/final_selection.json @@ -0,0 +1,37 @@ +{ + "selected_method": "A0", + "provisional_seed42_method": "A2", + "candidate_methods_with_three_seeds": [ + "A2", + "A0", + "A1" + ], + "selection_rule": "lowest mean fixed four-scenario validation task loss across seeds 42, 3407, 2026", + "candidate_summary": [ + { + "method": "A0", + "seed_losses": "[0.8643234267339601, 0.8756466648735842, 0.8574190991265433]", + "mean_validation_selection_loss": 0.8657963969113626, + "std_validation_selection_loss": 0.009202622948013908, + "seeds": 3, + "validation_only_selection": true + }, + { + "method": "A1", + "seed_losses": "[0.8649765662439577, 0.8810990981675766, 0.8557776766163963]", + "mean_validation_selection_loss": 0.8672844470093102, + "std_validation_selection_loss": 0.012817501026466038, + "seeds": 3, + "validation_only_selection": true + }, + { + "method": "A2", + "seed_losses": "[0.8634063961741689, 0.880271397449158, 0.862087192279952]", + "mean_validation_selection_loss": 0.8685883286344263, + "std_validation_selection_loss": 0.010139311979910028, + "seeds": 3, + "validation_only_selection": true + } + ], + "attachment4_labels_used": false +} diff --git a/final/experiments/q3/ati_ho/models/A0/seed_2026/model_best.pt b/final/experiments/q3/ati_ho/models/A0/seed_2026/model_best.pt new file mode 100644 index 0000000..a7953bb Binary files /dev/null and b/final/experiments/q3/ati_ho/models/A0/seed_2026/model_best.pt differ diff --git a/final/experiments/q3/ati_ho/models/A0/seed_2026/training_history.csv b/final/experiments/q3/ati_ho/models/A0/seed_2026/training_history.csv new file mode 100644 index 0000000..5b10a4e --- /dev/null +++ b/final/experiments/q3/ati_ho/models/A0/seed_2026/training_history.csv @@ -0,0 +1,8 @@ +method,seed,epoch,train_loss,valid_selection_loss,valid_clean_loss,lambda_interaction,lambda_mask +A0,2026,1,1.0065954625606537,0.9363003446833118,0.931940136375008,0.001,0.0 +A0,2026,2,0.8319815562831031,0.8692282438278198,0.8585401066056975,0.001,0.0 +A0,2026,3,0.7417113295307866,0.8601733071135951,0.8481150000959963,0.001,0.0 +A0,2026,4,0.6949168112542894,0.8574190991265433,0.8474381146850166,0.001,0.0 +A0,2026,5,0.6569236450725131,0.8673936297277827,0.8559725572774698,0.001,0.0 +A0,2026,6,0.6017359914603057,0.9111922525770062,0.9027906433566586,0.001,0.0 +A0,2026,7,0.5584396239784029,0.9282659668843826,0.9144766841615949,0.001,0.0 diff --git a/final/experiments/q3/ati_ho/models/A0/seed_2026/training_manifest.json b/final/experiments/q3/ati_ho/models/A0/seed_2026/training_manifest.json new file mode 100644 index 0000000..79a57e2 --- /dev/null +++ b/final/experiments/q3/ati_ho/models/A0/seed_2026/training_manifest.json @@ -0,0 +1,137 @@ +{ + "method": "A0", + "seed": 2026, + "best_epoch": 4, + "best_selection_loss": 0.8574190991265433, + "elapsed_seconds": 6.77071833203081, + "batch_size": 64, + "epoch_limit": 12, + "patience": 3, + "optimizer": "AdamW", + "learning_rate": 0.0003, + "weight_decay": 0.001, + "gradient_clip_norm": 1.0, + "training_mask_rates": [ + 0.0, + 0.1, + 0.3, + 0.5, + 0.7 + ], + "training_mask_patterns": [ + "single", + "sync", + "partial", + "async" + ], + "training_mask_seed_base": 20261227, + "same_orders_and_masks_across_methods_for_same_seed": true, + "config": { + "name": "A0_main_effects", + "low_rank": false, + "cross_attention": false, + "anchored": true, + "lambda_interaction": 0.001, + "lambda_mask": 0.0, + "rank": 4, + "hidden": 64, + "gru_hidden_per_direction": 32, + "attention_heads": 4, + "attention_ffn": 128, + "eta_init": 0.1 + }, + "history": [ + { + "method": "A0", + "seed": 2026, + "epoch": 1, + "train_loss": 1.0065954625606537, + "valid_selection_loss": 0.9363003446833118, + "valid_clean_loss": 0.931940136375008, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A0", + "seed": 2026, + "epoch": 2, + "train_loss": 0.8319815562831031, + "valid_selection_loss": 0.8692282438278198, + "valid_clean_loss": 0.8585401066056975, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A0", + "seed": 2026, + "epoch": 3, + "train_loss": 0.7417113295307866, + "valid_selection_loss": 0.8601733071135951, + "valid_clean_loss": 0.8481150000959963, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A0", + "seed": 2026, + "epoch": 4, + "train_loss": 0.6949168112542894, + "valid_selection_loss": 0.8574190991265433, + "valid_clean_loss": 0.8474381146850166, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A0", + "seed": 2026, + "epoch": 5, + "train_loss": 0.6569236450725131, + "valid_selection_loss": 0.8673936297277827, + "valid_clean_loss": 0.8559725572774698, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A0", + "seed": 2026, + "epoch": 6, + "train_loss": 0.6017359914603057, + "valid_selection_loss": 0.9111922525770062, + "valid_clean_loss": 0.9027906433566586, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A0", + "seed": 2026, + "epoch": 7, + "train_loss": 0.5584396239784029, + "valid_selection_loss": 0.9282659668843826, + "valid_clean_loss": 0.9144766841615949, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + } + ], + "mask_counts": { + "0.1/async": 1149, + "0.0/async": 1181, + "0.5/sync": 1205, + "0.0/sync": 1177, + "0.5/single": 1154, + "0.0/single": 1223, + "0.7/async": 1223, + "0.0/partial": 1254, + "0.7/sync": 1182, + "0.5/partial": 1182, + "0.3/partial": 1153, + "0.1/partial": 1216, + "0.3/single": 1178, + "0.3/async": 1180, + "0.5/async": 1182, + "0.1/single": 1208, + "0.7/single": 1129, + "0.3/sync": 1220, + "0.1/sync": 1187, + "0.7/partial": 1182 + } +} \ No newline at end of file diff --git a/final/experiments/q3/ati_ho/models/A0/seed_3407/model_best.pt b/final/experiments/q3/ati_ho/models/A0/seed_3407/model_best.pt new file mode 100644 index 0000000..eaaf75d Binary files /dev/null and b/final/experiments/q3/ati_ho/models/A0/seed_3407/model_best.pt differ diff --git a/final/experiments/q3/ati_ho/models/A0/seed_3407/training_history.csv b/final/experiments/q3/ati_ho/models/A0/seed_3407/training_history.csv new file mode 100644 index 0000000..e5e4bb4 --- /dev/null +++ b/final/experiments/q3/ati_ho/models/A0/seed_3407/training_history.csv @@ -0,0 +1,7 @@ +method,seed,epoch,train_loss,valid_selection_loss,valid_clean_loss,lambda_interaction,lambda_mask +A0,3407,1,1.0028652648131053,0.9436100423336029,0.9391928321712619,0.001,0.0 +A0,3407,2,0.8238278020311285,0.8820641582811272,0.8730134918139532,0.001,0.0 +A0,3407,3,0.74309202587163,0.8756466648735842,0.8646609815922413,0.001,0.0 +A0,3407,4,0.6922215135009201,0.8815795684253777,0.8710692741058685,0.001,0.0 +A0,3407,5,0.6518548451088093,0.9090294297579881,0.8968785254509895,0.001,0.0 +A0,3407,6,0.6016097575150154,0.9234253629878326,0.9113659098908141,0.001,0.0 diff --git a/final/experiments/q3/ati_ho/models/A0/seed_3407/training_manifest.json b/final/experiments/q3/ati_ho/models/A0/seed_3407/training_manifest.json new file mode 100644 index 0000000..718e319 --- /dev/null +++ b/final/experiments/q3/ati_ho/models/A0/seed_3407/training_manifest.json @@ -0,0 +1,127 @@ +{ + "method": "A0", + "seed": 3407, + "best_epoch": 3, + "best_selection_loss": 0.8756466648735842, + "elapsed_seconds": 5.952770586998668, + "batch_size": 64, + "epoch_limit": 12, + "patience": 3, + "optimizer": "AdamW", + "learning_rate": 0.0003, + "weight_decay": 0.001, + "gradient_clip_norm": 1.0, + "training_mask_rates": [ + 0.0, + 0.1, + 0.3, + 0.5, + 0.7 + ], + "training_mask_patterns": [ + "single", + "sync", + "partial", + "async" + ], + "training_mask_seed_base": 20261227, + "same_orders_and_masks_across_methods_for_same_seed": true, + "config": { + "name": "A0_main_effects", + "low_rank": false, + "cross_attention": false, + "anchored": true, + "lambda_interaction": 0.001, + "lambda_mask": 0.0, + "rank": 4, + "hidden": 64, + "gru_hidden_per_direction": 32, + "attention_heads": 4, + "attention_ffn": 128, + "eta_init": 0.1 + }, + "history": [ + { + "method": "A0", + "seed": 3407, + "epoch": 1, + "train_loss": 1.0028652648131053, + "valid_selection_loss": 0.9436100423336029, + "valid_clean_loss": 0.9391928321712619, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A0", + "seed": 3407, + "epoch": 2, + "train_loss": 0.8238278020311285, + "valid_selection_loss": 0.8820641582811272, + "valid_clean_loss": 0.8730134918139532, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A0", + "seed": 3407, + "epoch": 3, + "train_loss": 0.74309202587163, + "valid_selection_loss": 0.8756466648735842, + "valid_clean_loss": 0.8646609815922413, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A0", + "seed": 3407, + "epoch": 4, + "train_loss": 0.6922215135009201, + "valid_selection_loss": 0.8815795684253777, + "valid_clean_loss": 0.8710692741058685, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A0", + "seed": 3407, + "epoch": 5, + "train_loss": 0.6518548451088093, + "valid_selection_loss": 0.9090294297579881, + "valid_clean_loss": 0.8968785254509895, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A0", + "seed": 3407, + "epoch": 6, + "train_loss": 0.6016097575150154, + "valid_selection_loss": 0.9234253629878326, + "valid_clean_loss": 0.9113659098908141, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + } + ], + "mask_counts": { + "0.7/async": 1003, + "0.0/async": 1014, + "0.5/sync": 1033, + "0.1/sync": 1024, + "0.5/partial": 996, + "0.3/partial": 1019, + "0.5/async": 1053, + "0.3/single": 1003, + "0.0/sync": 1042, + "0.7/partial": 1079, + "0.7/sync": 986, + "0.0/single": 1029, + "0.1/async": 959, + "0.3/sync": 1055, + "0.0/partial": 1016, + "0.1/single": 1011, + "0.1/partial": 988, + "0.5/single": 1076, + "0.7/single": 997, + "0.3/async": 987 + } +} \ No newline at end of file diff --git a/final/experiments/q3/ati_ho/models/A0/seed_42/model_best.pt b/final/experiments/q3/ati_ho/models/A0/seed_42/model_best.pt new file mode 100644 index 0000000..cde9ddf Binary files /dev/null and b/final/experiments/q3/ati_ho/models/A0/seed_42/model_best.pt differ diff --git a/final/experiments/q3/ati_ho/models/A0/seed_42/training_history.csv b/final/experiments/q3/ati_ho/models/A0/seed_42/training_history.csv new file mode 100644 index 0000000..c1c81df --- /dev/null +++ b/final/experiments/q3/ati_ho/models/A0/seed_42/training_history.csv @@ -0,0 +1,8 @@ +method,seed,epoch,train_loss,valid_selection_loss,valid_clean_loss,lambda_interaction,lambda_mask +A0,42,1,0.9974622616061458,0.9360404843157465,0.9310548764008743,0.001,0.0 +A0,42,2,0.8197665159349088,0.8663458909307207,0.8535200449136587,0.001,0.0 +A0,42,3,0.7383816402267527,0.8669271827726575,0.8542619085573888,0.001,0.0 +A0,42,4,0.6841275570569215,0.8643234267339601,0.850496660222064,0.001,0.0 +A0,42,5,0.6515898009141287,0.8761576001460736,0.8627192581092918,0.001,0.0 +A0,42,6,0.6044590015102316,0.9111036181777388,0.8968614366028335,0.001,0.0 +A0,42,7,0.5584419562860772,0.9260857973124955,0.9113180768358838,0.001,0.0 diff --git a/final/experiments/q3/ati_ho/models/A0/seed_42/training_manifest.json b/final/experiments/q3/ati_ho/models/A0/seed_42/training_manifest.json new file mode 100644 index 0000000..52d0746 --- /dev/null +++ b/final/experiments/q3/ati_ho/models/A0/seed_42/training_manifest.json @@ -0,0 +1,137 @@ +{ + "method": "A0", + "seed": 42, + "best_epoch": 4, + "best_selection_loss": 0.8643234267339601, + "elapsed_seconds": 12.891749234986492, + "batch_size": 64, + "epoch_limit": 12, + "patience": 3, + "optimizer": "AdamW", + "learning_rate": 0.0003, + "weight_decay": 0.001, + "gradient_clip_norm": 1.0, + "training_mask_rates": [ + 0.0, + 0.1, + 0.3, + 0.5, + 0.7 + ], + "training_mask_patterns": [ + "single", + "sync", + "partial", + "async" + ], + "training_mask_seed_base": 20261227, + "same_orders_and_masks_across_methods_for_same_seed": true, + "config": { + "name": "A0_main_effects", + "low_rank": false, + "cross_attention": false, + "anchored": true, + "lambda_interaction": 0.001, + "lambda_mask": 0.0, + "rank": 4, + "hidden": 64, + "gru_hidden_per_direction": 32, + "attention_heads": 4, + "attention_ffn": 128, + "eta_init": 0.1 + }, + "history": [ + { + "method": "A0", + "seed": 42, + "epoch": 1, + "train_loss": 0.9974622616061458, + "valid_selection_loss": 0.9360404843157465, + "valid_clean_loss": 0.9310548764008743, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A0", + "seed": 42, + "epoch": 2, + "train_loss": 0.8197665159349088, + "valid_selection_loss": 0.8663458909307207, + "valid_clean_loss": 0.8535200449136587, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A0", + "seed": 42, + "epoch": 3, + "train_loss": 0.7383816402267527, + "valid_selection_loss": 0.8669271827726575, + "valid_clean_loss": 0.8542619085573888, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A0", + "seed": 42, + "epoch": 4, + "train_loss": 0.6841275570569215, + "valid_selection_loss": 0.8643234267339601, + "valid_clean_loss": 0.850496660222064, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A0", + "seed": 42, + "epoch": 5, + "train_loss": 0.6515898009141287, + "valid_selection_loss": 0.8761576001460736, + "valid_clean_loss": 0.8627192581092918, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A0", + "seed": 42, + "epoch": 6, + "train_loss": 0.6044590015102316, + "valid_selection_loss": 0.9111036181777388, + "valid_clean_loss": 0.8968614366028335, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A0", + "seed": 42, + "epoch": 7, + "train_loss": 0.5584419562860772, + "valid_selection_loss": 0.9260857973124955, + "valid_clean_loss": 0.9113180768358838, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + } + ], + "mask_counts": { + "0.0/single": 1169, + "0.3/single": 1187, + "0.5/single": 1158, + "0.7/sync": 1272, + "0.7/partial": 1207, + "0.0/async": 1145, + "0.3/async": 1200, + "0.3/partial": 1209, + "0.3/sync": 1170, + "0.1/single": 1205, + "0.5/async": 1185, + "0.7/single": 1232, + "0.5/sync": 1143, + "0.5/partial": 1236, + "0.0/partial": 1183, + "0.0/sync": 1159, + "0.7/async": 1179, + "0.1/sync": 1202, + "0.1/async": 1192, + "0.1/partial": 1132 + } +} \ No newline at end of file diff --git a/final/experiments/q3/ati_ho/models/A1/seed_2026/model_best.pt b/final/experiments/q3/ati_ho/models/A1/seed_2026/model_best.pt new file mode 100644 index 0000000..4b7eccd Binary files /dev/null and b/final/experiments/q3/ati_ho/models/A1/seed_2026/model_best.pt differ diff --git a/final/experiments/q3/ati_ho/models/A1/seed_2026/training_history.csv b/final/experiments/q3/ati_ho/models/A1/seed_2026/training_history.csv new file mode 100644 index 0000000..e88155d --- /dev/null +++ b/final/experiments/q3/ati_ho/models/A1/seed_2026/training_history.csv @@ -0,0 +1,8 @@ +method,seed,epoch,train_loss,valid_selection_loss,valid_clean_loss,lambda_interaction,lambda_mask +A1,2026,1,1.0153538673012346,0.9397767204177248,0.9357580549114353,0.001,0.0 +A1,2026,2,0.8344417214393616,0.8647092963968004,0.8555307774753361,0.001,0.0 +A1,2026,3,0.7413759132226309,0.8561831976358707,0.8470297724336058,0.001,0.0 +A1,2026,4,0.6923596881054066,0.8557776766163963,0.8481834858328432,0.001,0.0 +A1,2026,5,0.6530446432254933,0.8681209758742825,0.8606502813297313,0.001,0.0 +A1,2026,6,0.5979297994463532,0.9125843056283154,0.9086607209928743,0.001,0.0 +A1,2026,7,0.5576094913261908,0.9276157835355172,0.9214578185762677,0.001,0.0 diff --git a/final/experiments/q3/ati_ho/models/A1/seed_2026/training_manifest.json b/final/experiments/q3/ati_ho/models/A1/seed_2026/training_manifest.json new file mode 100644 index 0000000..926f190 --- /dev/null +++ b/final/experiments/q3/ati_ho/models/A1/seed_2026/training_manifest.json @@ -0,0 +1,137 @@ +{ + "method": "A1", + "seed": 2026, + "best_epoch": 4, + "best_selection_loss": 0.8557776766163963, + "elapsed_seconds": 7.712336958968081, + "batch_size": 64, + "epoch_limit": 12, + "patience": 3, + "optimizer": "AdamW", + "learning_rate": 0.0003, + "weight_decay": 0.001, + "gradient_clip_norm": 1.0, + "training_mask_rates": [ + 0.0, + 0.1, + 0.3, + 0.5, + 0.7 + ], + "training_mask_patterns": [ + "single", + "sync", + "partial", + "async" + ], + "training_mask_seed_base": 20261227, + "same_orders_and_masks_across_methods_for_same_seed": true, + "config": { + "name": "A1_low_rank_pairs", + "low_rank": true, + "cross_attention": false, + "anchored": true, + "lambda_interaction": 0.001, + "lambda_mask": 0.0, + "rank": 4, + "hidden": 64, + "gru_hidden_per_direction": 32, + "attention_heads": 4, + "attention_ffn": 128, + "eta_init": 0.1 + }, + "history": [ + { + "method": "A1", + "seed": 2026, + "epoch": 1, + "train_loss": 1.0153538673012346, + "valid_selection_loss": 0.9397767204177248, + "valid_clean_loss": 0.9357580549114353, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A1", + "seed": 2026, + "epoch": 2, + "train_loss": 0.8344417214393616, + "valid_selection_loss": 0.8647092963968004, + "valid_clean_loss": 0.8555307774753361, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A1", + "seed": 2026, + "epoch": 3, + "train_loss": 0.7413759132226309, + "valid_selection_loss": 0.8561831976358707, + "valid_clean_loss": 0.8470297724336058, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A1", + "seed": 2026, + "epoch": 4, + "train_loss": 0.6923596881054066, + "valid_selection_loss": 0.8557776766163963, + "valid_clean_loss": 0.8481834858328432, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A1", + "seed": 2026, + "epoch": 5, + "train_loss": 0.6530446432254933, + "valid_selection_loss": 0.8681209758742825, + "valid_clean_loss": 0.8606502813297313, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A1", + "seed": 2026, + "epoch": 6, + "train_loss": 0.5979297994463532, + "valid_selection_loss": 0.9125843056283154, + "valid_clean_loss": 0.9086607209928743, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A1", + "seed": 2026, + "epoch": 7, + "train_loss": 0.5576094913261908, + "valid_selection_loss": 0.9276157835355172, + "valid_clean_loss": 0.9214578185762677, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + } + ], + "mask_counts": { + "0.1/async": 1149, + "0.0/async": 1181, + "0.5/sync": 1205, + "0.0/sync": 1177, + "0.5/single": 1154, + "0.0/single": 1223, + "0.7/async": 1223, + "0.0/partial": 1254, + "0.7/sync": 1182, + "0.5/partial": 1182, + "0.3/partial": 1153, + "0.1/partial": 1216, + "0.3/single": 1178, + "0.3/async": 1180, + "0.5/async": 1182, + "0.1/single": 1208, + "0.7/single": 1129, + "0.3/sync": 1220, + "0.1/sync": 1187, + "0.7/partial": 1182 + } +} \ No newline at end of file diff --git a/final/experiments/q3/ati_ho/models/A1/seed_3407/model_best.pt b/final/experiments/q3/ati_ho/models/A1/seed_3407/model_best.pt new file mode 100644 index 0000000..a76d86e Binary files /dev/null and b/final/experiments/q3/ati_ho/models/A1/seed_3407/model_best.pt differ diff --git a/final/experiments/q3/ati_ho/models/A1/seed_3407/training_history.csv b/final/experiments/q3/ati_ho/models/A1/seed_3407/training_history.csv new file mode 100644 index 0000000..946c717 --- /dev/null +++ b/final/experiments/q3/ati_ho/models/A1/seed_3407/training_history.csv @@ -0,0 +1,7 @@ +method,seed,epoch,train_loss,valid_selection_loss,valid_clean_loss,lambda_interaction,lambda_mask +A1,3407,1,0.993991059285623,0.9350224992076119,0.9297198895569686,0.001,0.0 +A1,3407,2,0.8231704179887418,0.8853522272227885,0.8782615517521952,0.001,0.0 +A1,3407,3,0.7412442289016865,0.8810990981675766,0.872995504966149,0.001,0.0 +A1,3407,4,0.6908998235508248,0.8885695379186462,0.8816041920211289,0.001,0.0 +A1,3407,5,0.6549228540173283,0.9160289600655273,0.9081327443594461,0.001,0.0 +A1,3407,6,0.6015714597370889,0.9299172356233493,0.9222097593349415,0.001,0.0 diff --git a/final/experiments/q3/ati_ho/models/A1/seed_3407/training_manifest.json b/final/experiments/q3/ati_ho/models/A1/seed_3407/training_manifest.json new file mode 100644 index 0000000..6c46530 --- /dev/null +++ b/final/experiments/q3/ati_ho/models/A1/seed_3407/training_manifest.json @@ -0,0 +1,127 @@ +{ + "method": "A1", + "seed": 3407, + "best_epoch": 3, + "best_selection_loss": 0.8810990981675766, + "elapsed_seconds": 6.539086283999495, + "batch_size": 64, + "epoch_limit": 12, + "patience": 3, + "optimizer": "AdamW", + "learning_rate": 0.0003, + "weight_decay": 0.001, + "gradient_clip_norm": 1.0, + "training_mask_rates": [ + 0.0, + 0.1, + 0.3, + 0.5, + 0.7 + ], + "training_mask_patterns": [ + "single", + "sync", + "partial", + "async" + ], + "training_mask_seed_base": 20261227, + "same_orders_and_masks_across_methods_for_same_seed": true, + "config": { + "name": "A1_low_rank_pairs", + "low_rank": true, + "cross_attention": false, + "anchored": true, + "lambda_interaction": 0.001, + "lambda_mask": 0.0, + "rank": 4, + "hidden": 64, + "gru_hidden_per_direction": 32, + "attention_heads": 4, + "attention_ffn": 128, + "eta_init": 0.1 + }, + "history": [ + { + "method": "A1", + "seed": 3407, + "epoch": 1, + "train_loss": 0.993991059285623, + "valid_selection_loss": 0.9350224992076119, + "valid_clean_loss": 0.9297198895569686, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A1", + "seed": 3407, + "epoch": 2, + "train_loss": 0.8231704179887418, + "valid_selection_loss": 0.8853522272227885, + "valid_clean_loss": 0.8782615517521952, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A1", + "seed": 3407, + "epoch": 3, + "train_loss": 0.7412442289016865, + "valid_selection_loss": 0.8810990981675766, + "valid_clean_loss": 0.872995504966149, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A1", + "seed": 3407, + "epoch": 4, + "train_loss": 0.6908998235508248, + "valid_selection_loss": 0.8885695379186462, + "valid_clean_loss": 0.8816041920211289, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A1", + "seed": 3407, + "epoch": 5, + "train_loss": 0.6549228540173283, + "valid_selection_loss": 0.9160289600655273, + "valid_clean_loss": 0.9081327443594461, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A1", + "seed": 3407, + "epoch": 6, + "train_loss": 0.6015714597370889, + "valid_selection_loss": 0.9299172356233493, + "valid_clean_loss": 0.9222097593349415, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + } + ], + "mask_counts": { + "0.7/async": 1003, + "0.0/async": 1014, + "0.5/sync": 1033, + "0.1/sync": 1024, + "0.5/partial": 996, + "0.3/partial": 1019, + "0.5/async": 1053, + "0.3/single": 1003, + "0.0/sync": 1042, + "0.7/partial": 1079, + "0.7/sync": 986, + "0.0/single": 1029, + "0.1/async": 959, + "0.3/sync": 1055, + "0.0/partial": 1016, + "0.1/single": 1011, + "0.1/partial": 988, + "0.5/single": 1076, + "0.7/single": 997, + "0.3/async": 987 + } +} \ No newline at end of file diff --git a/final/experiments/q3/ati_ho/models/A1/seed_42/model_best.pt b/final/experiments/q3/ati_ho/models/A1/seed_42/model_best.pt new file mode 100644 index 0000000..f4c21cf Binary files /dev/null and b/final/experiments/q3/ati_ho/models/A1/seed_42/model_best.pt differ diff --git a/final/experiments/q3/ati_ho/models/A1/seed_42/training_history.csv b/final/experiments/q3/ati_ho/models/A1/seed_42/training_history.csv new file mode 100644 index 0000000..f3ecb11 --- /dev/null +++ b/final/experiments/q3/ati_ho/models/A1/seed_42/training_history.csv @@ -0,0 +1,8 @@ +method,seed,epoch,train_loss,valid_selection_loss,valid_clean_loss,lambda_interaction,lambda_mask +A1,42,1,0.9995549983448453,0.9426855438358182,0.9384826362788022,0.001,0.0 +A1,42,2,0.8219450292763887,0.8703031135456902,0.8586705464583176,0.001,0.0 +A1,42,3,0.7375812182823817,0.8694013199963413,0.8590540899025215,0.001,0.0 +A1,42,4,0.6825026941520197,0.8649765662439577,0.8543586613057734,0.001,0.0 +A1,42,5,0.648727031217681,0.877989024087623,0.8682481325589694,0.001,0.0 +A1,42,6,0.6008665031856961,0.9173483715935067,0.9082090225848523,0.001,0.0 +A1,42,7,0.5538152815015228,0.9287053037148255,0.9185833616571112,0.001,0.0 diff --git a/final/experiments/q3/ati_ho/models/A1/seed_42/training_manifest.json b/final/experiments/q3/ati_ho/models/A1/seed_42/training_manifest.json new file mode 100644 index 0000000..1cc9710 --- /dev/null +++ b/final/experiments/q3/ati_ho/models/A1/seed_42/training_manifest.json @@ -0,0 +1,137 @@ +{ + "method": "A1", + "seed": 42, + "best_epoch": 4, + "best_selection_loss": 0.8649765662439577, + "elapsed_seconds": 7.483738042006735, + "batch_size": 64, + "epoch_limit": 12, + "patience": 3, + "optimizer": "AdamW", + "learning_rate": 0.0003, + "weight_decay": 0.001, + "gradient_clip_norm": 1.0, + "training_mask_rates": [ + 0.0, + 0.1, + 0.3, + 0.5, + 0.7 + ], + "training_mask_patterns": [ + "single", + "sync", + "partial", + "async" + ], + "training_mask_seed_base": 20261227, + "same_orders_and_masks_across_methods_for_same_seed": true, + "config": { + "name": "A1_low_rank_pairs", + "low_rank": true, + "cross_attention": false, + "anchored": true, + "lambda_interaction": 0.001, + "lambda_mask": 0.0, + "rank": 4, + "hidden": 64, + "gru_hidden_per_direction": 32, + "attention_heads": 4, + "attention_ffn": 128, + "eta_init": 0.1 + }, + "history": [ + { + "method": "A1", + "seed": 42, + "epoch": 1, + "train_loss": 0.9995549983448453, + "valid_selection_loss": 0.9426855438358182, + "valid_clean_loss": 0.9384826362788022, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A1", + "seed": 42, + "epoch": 2, + "train_loss": 0.8219450292763887, + "valid_selection_loss": 0.8703031135456902, + "valid_clean_loss": 0.8586705464583176, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A1", + "seed": 42, + "epoch": 3, + "train_loss": 0.7375812182823817, + "valid_selection_loss": 0.8694013199963413, + "valid_clean_loss": 0.8590540899025215, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A1", + "seed": 42, + "epoch": 4, + "train_loss": 0.6825026941520197, + "valid_selection_loss": 0.8649765662439577, + "valid_clean_loss": 0.8543586613057734, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A1", + "seed": 42, + "epoch": 5, + "train_loss": 0.648727031217681, + "valid_selection_loss": 0.877989024087623, + "valid_clean_loss": 0.8682481325589694, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A1", + "seed": 42, + "epoch": 6, + "train_loss": 0.6008665031856961, + "valid_selection_loss": 0.9173483715935067, + "valid_clean_loss": 0.9082090225848523, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A1", + "seed": 42, + "epoch": 7, + "train_loss": 0.5538152815015228, + "valid_selection_loss": 0.9287053037148255, + "valid_clean_loss": 0.9185833616571112, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + } + ], + "mask_counts": { + "0.0/single": 1169, + "0.3/single": 1187, + "0.5/single": 1158, + "0.7/sync": 1272, + "0.7/partial": 1207, + "0.0/async": 1145, + "0.3/async": 1200, + "0.3/partial": 1209, + "0.3/sync": 1170, + "0.1/single": 1205, + "0.5/async": 1185, + "0.7/single": 1232, + "0.5/sync": 1143, + "0.5/partial": 1236, + "0.0/partial": 1183, + "0.0/sync": 1159, + "0.7/async": 1179, + "0.1/sync": 1202, + "0.1/async": 1192, + "0.1/partial": 1132 + } +} \ No newline at end of file diff --git a/final/experiments/q3/ati_ho/models/A2/seed_2026/model_best.pt b/final/experiments/q3/ati_ho/models/A2/seed_2026/model_best.pt new file mode 100644 index 0000000..0dcdcf5 Binary files /dev/null and b/final/experiments/q3/ati_ho/models/A2/seed_2026/model_best.pt differ diff --git a/final/experiments/q3/ati_ho/models/A2/seed_2026/training_history.csv b/final/experiments/q3/ati_ho/models/A2/seed_2026/training_history.csv new file mode 100644 index 0000000..94a6d35 --- /dev/null +++ b/final/experiments/q3/ati_ho/models/A2/seed_2026/training_history.csv @@ -0,0 +1,7 @@ +method,seed,epoch,train_loss,valid_selection_loss,valid_clean_loss,lambda_interaction,lambda_mask +A2,2026,1,0.9926471710205078,0.928315669953168,0.9239131356333639,0.001,0.0 +A2,2026,2,0.820087315859618,0.8673638532777409,0.8601592170013176,0.001,0.0 +A2,2026,3,0.7308831590193289,0.862087192279952,0.8561929890087673,0.001,0.0 +A2,2026,4,0.6851603929643277,0.8666458973190287,0.8619590087251349,0.001,0.0 +A2,2026,5,0.6444870498445299,0.88101957656525,0.8752215864894154,0.001,0.0 +A2,2026,6,0.5834257365376861,0.9354289185542327,0.9353187752294017,0.001,0.0 diff --git a/final/experiments/q3/ati_ho/models/A2/seed_2026/training_manifest.json b/final/experiments/q3/ati_ho/models/A2/seed_2026/training_manifest.json new file mode 100644 index 0000000..918993a --- /dev/null +++ b/final/experiments/q3/ati_ho/models/A2/seed_2026/training_manifest.json @@ -0,0 +1,127 @@ +{ + "method": "A2", + "seed": 2026, + "best_epoch": 3, + "best_selection_loss": 0.862087192279952, + "elapsed_seconds": 9.472242687013932, + "batch_size": 64, + "epoch_limit": 12, + "patience": 3, + "optimizer": "AdamW", + "learning_rate": 0.0003, + "weight_decay": 0.001, + "gradient_clip_norm": 1.0, + "training_mask_rates": [ + 0.0, + 0.1, + 0.3, + 0.5, + 0.7 + ], + "training_mask_patterns": [ + "single", + "sync", + "partial", + "async" + ], + "training_mask_seed_base": 20261227, + "same_orders_and_masks_across_methods_for_same_seed": true, + "config": { + "name": "A2_anchored_pairwise", + "low_rank": true, + "cross_attention": true, + "anchored": true, + "lambda_interaction": 0.001, + "lambda_mask": 0.0, + "rank": 4, + "hidden": 64, + "gru_hidden_per_direction": 32, + "attention_heads": 4, + "attention_ffn": 128, + "eta_init": 0.1 + }, + "history": [ + { + "method": "A2", + "seed": 2026, + "epoch": 1, + "train_loss": 0.9926471710205078, + "valid_selection_loss": 0.928315669953168, + "valid_clean_loss": 0.9239131356333639, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A2", + "seed": 2026, + "epoch": 2, + "train_loss": 0.820087315859618, + "valid_selection_loss": 0.8673638532777409, + "valid_clean_loss": 0.8601592170013176, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A2", + "seed": 2026, + "epoch": 3, + "train_loss": 0.7308831590193289, + "valid_selection_loss": 0.862087192279952, + "valid_clean_loss": 0.8561929890087673, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A2", + "seed": 2026, + "epoch": 4, + "train_loss": 0.6851603929643277, + "valid_selection_loss": 0.8666458973190287, + "valid_clean_loss": 0.8619590087251349, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A2", + "seed": 2026, + "epoch": 5, + "train_loss": 0.6444870498445299, + "valid_selection_loss": 0.88101957656525, + "valid_clean_loss": 0.8752215864894154, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A2", + "seed": 2026, + "epoch": 6, + "train_loss": 0.5834257365376861, + "valid_selection_loss": 0.9354289185542327, + "valid_clean_loss": 0.9353187752294017, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + } + ], + "mask_counts": { + "0.1/async": 970, + "0.0/async": 995, + "0.5/sync": 1040, + "0.0/sync": 1012, + "0.5/single": 999, + "0.0/single": 1046, + "0.7/async": 1048, + "0.0/partial": 1063, + "0.7/sync": 1020, + "0.5/partial": 1022, + "0.3/partial": 1002, + "0.1/partial": 1022, + "0.3/single": 1004, + "0.3/async": 1026, + "0.5/async": 1028, + "0.1/single": 1048, + "0.7/single": 959, + "0.3/sync": 1030, + "0.1/sync": 1007, + "0.7/partial": 1029 + } +} \ No newline at end of file diff --git a/final/experiments/q3/ati_ho/models/A2/seed_3407/model_best.pt b/final/experiments/q3/ati_ho/models/A2/seed_3407/model_best.pt new file mode 100644 index 0000000..02d474a Binary files /dev/null and b/final/experiments/q3/ati_ho/models/A2/seed_3407/model_best.pt differ diff --git a/final/experiments/q3/ati_ho/models/A2/seed_3407/training_history.csv b/final/experiments/q3/ati_ho/models/A2/seed_3407/training_history.csv new file mode 100644 index 0000000..c947bad --- /dev/null +++ b/final/experiments/q3/ati_ho/models/A2/seed_3407/training_history.csv @@ -0,0 +1,7 @@ +method,seed,epoch,train_loss,valid_selection_loss,valid_clean_loss,lambda_interaction,lambda_mask +A2,3407,1,0.9647035289693762,0.9118595260840197,0.9059003676686969,0.001,0.0 +A2,3407,2,0.8049548080673924,0.8821067369573719,0.8777574326965835,0.001,0.0 +A2,3407,3,0.7312485785396011,0.880271397449158,0.8748716823347322,0.001,0.0 +A2,3407,4,0.6844777796003554,0.88496050123985,0.881083797622513,0.001,0.0 +A2,3407,5,0.6395798369690224,0.9139724308317835,0.9094404204861148,0.001,0.0 +A2,3407,6,0.5926711849040456,0.9307978958873957,0.9263341426849365,0.001,0.0 diff --git a/final/experiments/q3/ati_ho/models/A2/seed_3407/training_manifest.json b/final/experiments/q3/ati_ho/models/A2/seed_3407/training_manifest.json new file mode 100644 index 0000000..0ece989 --- /dev/null +++ b/final/experiments/q3/ati_ho/models/A2/seed_3407/training_manifest.json @@ -0,0 +1,127 @@ +{ + "method": "A2", + "seed": 3407, + "best_epoch": 3, + "best_selection_loss": 0.880271397449158, + "elapsed_seconds": 9.558659436006565, + "batch_size": 64, + "epoch_limit": 12, + "patience": 3, + "optimizer": "AdamW", + "learning_rate": 0.0003, + "weight_decay": 0.001, + "gradient_clip_norm": 1.0, + "training_mask_rates": [ + 0.0, + 0.1, + 0.3, + 0.5, + 0.7 + ], + "training_mask_patterns": [ + "single", + "sync", + "partial", + "async" + ], + "training_mask_seed_base": 20261227, + "same_orders_and_masks_across_methods_for_same_seed": true, + "config": { + "name": "A2_anchored_pairwise", + "low_rank": true, + "cross_attention": true, + "anchored": true, + "lambda_interaction": 0.001, + "lambda_mask": 0.0, + "rank": 4, + "hidden": 64, + "gru_hidden_per_direction": 32, + "attention_heads": 4, + "attention_ffn": 128, + "eta_init": 0.1 + }, + "history": [ + { + "method": "A2", + "seed": 3407, + "epoch": 1, + "train_loss": 0.9647035289693762, + "valid_selection_loss": 0.9118595260840197, + "valid_clean_loss": 0.9059003676686969, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A2", + "seed": 3407, + "epoch": 2, + "train_loss": 0.8049548080673924, + "valid_selection_loss": 0.8821067369573719, + "valid_clean_loss": 0.8777574326965835, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A2", + "seed": 3407, + "epoch": 3, + "train_loss": 0.7312485785396011, + "valid_selection_loss": 0.880271397449158, + "valid_clean_loss": 0.8748716823347322, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A2", + "seed": 3407, + "epoch": 4, + "train_loss": 0.6844777796003554, + "valid_selection_loss": 0.88496050123985, + "valid_clean_loss": 0.881083797622513, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A2", + "seed": 3407, + "epoch": 5, + "train_loss": 0.6395798369690224, + "valid_selection_loss": 0.9139724308317835, + "valid_clean_loss": 0.9094404204861148, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A2", + "seed": 3407, + "epoch": 6, + "train_loss": 0.5926711849040456, + "valid_selection_loss": 0.9307978958873957, + "valid_clean_loss": 0.9263341426849365, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + } + ], + "mask_counts": { + "0.7/async": 1003, + "0.0/async": 1014, + "0.5/sync": 1033, + "0.1/sync": 1024, + "0.5/partial": 996, + "0.3/partial": 1019, + "0.5/async": 1053, + "0.3/single": 1003, + "0.0/sync": 1042, + "0.7/partial": 1079, + "0.7/sync": 986, + "0.0/single": 1029, + "0.1/async": 959, + "0.3/sync": 1055, + "0.0/partial": 1016, + "0.1/single": 1011, + "0.1/partial": 988, + "0.5/single": 1076, + "0.7/single": 997, + "0.3/async": 987 + } +} \ No newline at end of file diff --git a/final/experiments/q3/ati_ho/models/A2/seed_42/model_best.pt b/final/experiments/q3/ati_ho/models/A2/seed_42/model_best.pt new file mode 100644 index 0000000..017d5f1 Binary files /dev/null and b/final/experiments/q3/ati_ho/models/A2/seed_42/model_best.pt differ diff --git a/final/experiments/q3/ati_ho/models/A2/seed_42/training_history.csv b/final/experiments/q3/ati_ho/models/A2/seed_42/training_history.csv new file mode 100644 index 0000000..4e87fcd --- /dev/null +++ b/final/experiments/q3/ati_ho/models/A2/seed_42/training_history.csv @@ -0,0 +1,6 @@ +method,seed,epoch,train_loss,valid_selection_loss,valid_clean_loss,lambda_interaction,lambda_mask +A2,42,1,0.9741679154060505,0.9116741605512388,0.9052737184933254,0.001,0.0 +A2,42,2,0.7997706675970996,0.8634063961741689,0.8530163961452443,0.001,0.0 +A2,42,3,0.7264242089456983,0.8717996972602802,0.8635974290606739,0.001,0.0 +A2,42,4,0.6684221158976908,0.8713849653581996,0.8629672671412374,0.001,0.0 +A2,42,5,0.6371411096166681,0.8887755983805918,0.8809175399633554,0.001,0.0 diff --git a/final/experiments/q3/ati_ho/models/A2/seed_42/training_manifest.json b/final/experiments/q3/ati_ho/models/A2/seed_42/training_manifest.json new file mode 100644 index 0000000..9008d9d --- /dev/null +++ b/final/experiments/q3/ati_ho/models/A2/seed_42/training_manifest.json @@ -0,0 +1,117 @@ +{ + "method": "A2", + "seed": 42, + "best_epoch": 2, + "best_selection_loss": 0.8634063961741689, + "elapsed_seconds": 8.171100930019747, + "batch_size": 64, + "epoch_limit": 12, + "patience": 3, + "optimizer": "AdamW", + "learning_rate": 0.0003, + "weight_decay": 0.001, + "gradient_clip_norm": 1.0, + "training_mask_rates": [ + 0.0, + 0.1, + 0.3, + 0.5, + 0.7 + ], + "training_mask_patterns": [ + "single", + "sync", + "partial", + "async" + ], + "training_mask_seed_base": 20261227, + "same_orders_and_masks_across_methods_for_same_seed": true, + "config": { + "name": "A2_anchored_pairwise", + "low_rank": true, + "cross_attention": true, + "anchored": true, + "lambda_interaction": 0.001, + "lambda_mask": 0.0, + "rank": 4, + "hidden": 64, + "gru_hidden_per_direction": 32, + "attention_heads": 4, + "attention_ffn": 128, + "eta_init": 0.1 + }, + "history": [ + { + "method": "A2", + "seed": 42, + "epoch": 1, + "train_loss": 0.9741679154060505, + "valid_selection_loss": 0.9116741605512388, + "valid_clean_loss": 0.9052737184933254, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A2", + "seed": 42, + "epoch": 2, + "train_loss": 0.7997706675970996, + "valid_selection_loss": 0.8634063961741689, + "valid_clean_loss": 0.8530163961452443, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A2", + "seed": 42, + "epoch": 3, + "train_loss": 0.7264242089456983, + "valid_selection_loss": 0.8717996972602802, + "valid_clean_loss": 0.8635974290606739, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A2", + "seed": 42, + "epoch": 4, + "train_loss": 0.6684221158976908, + "valid_selection_loss": 0.8713849653581996, + "valid_clean_loss": 0.8629672671412374, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "A2", + "seed": 42, + "epoch": 5, + "train_loss": 0.6371411096166681, + "valid_selection_loss": 0.8887755983805918, + "valid_clean_loss": 0.8809175399633554, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + } + ], + "mask_counts": { + "0.0/single": 838, + "0.3/single": 841, + "0.5/single": 832, + "0.7/sync": 931, + "0.7/partial": 854, + "0.0/async": 818, + "0.3/async": 856, + "0.3/partial": 861, + "0.3/sync": 841, + "0.1/single": 857, + "0.5/async": 832, + "0.7/single": 872, + "0.5/sync": 848, + "0.5/partial": 854, + "0.0/partial": 870, + "0.0/sync": 827, + "0.7/async": 831, + "0.1/sync": 839, + "0.1/async": 847, + "0.1/partial": 826 + } +} \ No newline at end of file diff --git a/final/experiments/q3/ati_ho/models/A3/seed_42/model_best.pt b/final/experiments/q3/ati_ho/models/A3/seed_42/model_best.pt new file mode 100644 index 0000000..dd62696 Binary files /dev/null and b/final/experiments/q3/ati_ho/models/A3/seed_42/model_best.pt differ diff --git a/final/experiments/q3/ati_ho/models/A3/seed_42/training_history.csv b/final/experiments/q3/ati_ho/models/A3/seed_42/training_history.csv new file mode 100644 index 0000000..e122dea --- /dev/null +++ b/final/experiments/q3/ati_ho/models/A3/seed_42/training_history.csv @@ -0,0 +1,6 @@ +method,seed,epoch,train_loss,valid_selection_loss,valid_clean_loss,lambda_interaction,lambda_mask +A3,42,1,1.0065002750467371,0.9119112013460515,0.9055310710445865,0.001,0.05 +A3,42,2,0.8278417852189806,0.8637567355737581,0.8532818588581714,0.001,0.05 +A3,42,3,0.7517132185123585,0.8728943449127805,0.8647457944167839,0.001,0.05 +A3,42,4,0.6915535065862868,0.8723347707764133,0.8638496981872307,0.001,0.05 +A3,42,5,0.6589032621295364,0.8902343159521019,0.8822925837485345,0.001,0.05 diff --git a/final/experiments/q3/ati_ho/models/A3/seed_42/training_manifest.json b/final/experiments/q3/ati_ho/models/A3/seed_42/training_manifest.json new file mode 100644 index 0000000..11b2f8f --- /dev/null +++ b/final/experiments/q3/ati_ho/models/A3/seed_42/training_manifest.json @@ -0,0 +1,117 @@ +{ + "method": "A3", + "seed": 42, + "best_epoch": 2, + "best_selection_loss": 0.8637567355737581, + "elapsed_seconds": 8.470124277984723, + "batch_size": 64, + "epoch_limit": 12, + "patience": 3, + "optimizer": "AdamW", + "learning_rate": 0.0003, + "weight_decay": 0.001, + "gradient_clip_norm": 1.0, + "training_mask_rates": [ + 0.0, + 0.1, + 0.3, + 0.5, + 0.7 + ], + "training_mask_patterns": [ + "single", + "sync", + "partial", + "async" + ], + "training_mask_seed_base": 20261227, + "same_orders_and_masks_across_methods_for_same_seed": true, + "config": { + "name": "A3_pairwise_mask_aux", + "low_rank": true, + "cross_attention": true, + "anchored": true, + "lambda_interaction": 0.001, + "lambda_mask": 0.05, + "rank": 4, + "hidden": 64, + "gru_hidden_per_direction": 32, + "attention_heads": 4, + "attention_ffn": 128, + "eta_init": 0.1 + }, + "history": [ + { + "method": "A3", + "seed": 42, + "epoch": 1, + "train_loss": 1.0065002750467371, + "valid_selection_loss": 0.9119112013460515, + "valid_clean_loss": 0.9055310710445865, + "lambda_interaction": 0.001, + "lambda_mask": 0.05 + }, + { + "method": "A3", + "seed": 42, + "epoch": 2, + "train_loss": 0.8278417852189806, + "valid_selection_loss": 0.8637567355737581, + "valid_clean_loss": 0.8532818588581714, + "lambda_interaction": 0.001, + "lambda_mask": 0.05 + }, + { + "method": "A3", + "seed": 42, + "epoch": 3, + "train_loss": 0.7517132185123585, + "valid_selection_loss": 0.8728943449127805, + "valid_clean_loss": 0.8647457944167839, + "lambda_interaction": 0.001, + "lambda_mask": 0.05 + }, + { + "method": "A3", + "seed": 42, + "epoch": 4, + "train_loss": 0.6915535065862868, + "valid_selection_loss": 0.8723347707764133, + "valid_clean_loss": 0.8638496981872307, + "lambda_interaction": 0.001, + "lambda_mask": 0.05 + }, + { + "method": "A3", + "seed": 42, + "epoch": 5, + "train_loss": 0.6589032621295364, + "valid_selection_loss": 0.8902343159521019, + "valid_clean_loss": 0.8822925837485345, + "lambda_interaction": 0.001, + "lambda_mask": 0.05 + } + ], + "mask_counts": { + "0.0/single": 838, + "0.3/single": 841, + "0.5/single": 832, + "0.7/sync": 931, + "0.7/partial": 854, + "0.0/async": 818, + "0.3/async": 856, + "0.3/partial": 861, + "0.3/sync": 841, + "0.1/single": 857, + "0.5/async": 832, + "0.7/single": 872, + "0.5/sync": 848, + "0.5/partial": 854, + "0.0/partial": 870, + "0.0/sync": 827, + "0.7/async": 831, + "0.1/sync": 839, + "0.1/async": 847, + "0.1/partial": 826 + } +} \ No newline at end of file diff --git a/final/experiments/q3/ati_ho/models/B0_early_concat/seed_2026/model_best.pt b/final/experiments/q3/ati_ho/models/B0_early_concat/seed_2026/model_best.pt new file mode 100644 index 0000000..e4be73c Binary files /dev/null and b/final/experiments/q3/ati_ho/models/B0_early_concat/seed_2026/model_best.pt differ diff --git a/final/experiments/q3/ati_ho/models/B0_early_concat/seed_2026/training_history.csv b/final/experiments/q3/ati_ho/models/B0_early_concat/seed_2026/training_history.csv new file mode 100644 index 0000000..5813375 --- /dev/null +++ b/final/experiments/q3/ati_ho/models/B0_early_concat/seed_2026/training_history.csv @@ -0,0 +1,6 @@ +method,seed,epoch,train_loss,valid_selection_loss,valid_clean_loss,lambda_interaction,lambda_mask +B0_early_concat,2026,1,0.9894467581201483,0.9033395762626941,0.8919235042163304,0.0,0.0 +B0_early_concat,2026,2,0.8127173174310613,0.8722499026047006,0.8655281905289535,0.0,0.0 +B0_early_concat,2026,3,0.7274277905623118,0.9010504417039537,0.9002949733000535,0.0,0.0 +B0_early_concat,2026,4,0.6712073153919644,0.9205121268610378,0.9208705294263232,0.0,0.0 +B0_early_concat,2026,5,0.6352586569609465,0.9440761974879672,0.9462754883608975,0.0,0.0 diff --git a/final/experiments/q3/ati_ho/models/B0_early_concat/seed_2026/training_manifest.json b/final/experiments/q3/ati_ho/models/B0_early_concat/seed_2026/training_manifest.json new file mode 100644 index 0000000..bc53bb0 --- /dev/null +++ b/final/experiments/q3/ati_ho/models/B0_early_concat/seed_2026/training_manifest.json @@ -0,0 +1,118 @@ +{ + "method": "B0_early_concat", + "seed": 2026, + "best_epoch": 2, + "best_selection_loss": 0.8722499026047006, + "elapsed_seconds": 2.9758365789894015, + "batch_size": 64, + "epoch_limit": 12, + "patience": 3, + "optimizer": "AdamW", + "learning_rate": 0.0003, + "weight_decay": 0.001, + "gradient_clip_norm": 1.0, + "training_mask_rates": [ + 0.0, + 0.1, + 0.3, + 0.5, + 0.7 + ], + "training_mask_patterns": [ + "single", + "sync", + "partial", + "async" + ], + "training_mask_seed_base": 20261227, + "same_orders_and_masks_across_methods_for_same_seed": true, + "config": { + "model_config": { + "router": "mlp", + "expert_names": [ + "T", + "A", + "V", + "TA", + "TV", + "AV", + "TAV" + ], + "availability_mode": "hard" + } + }, + "history": [ + { + "method": "B0_early_concat", + "seed": 2026, + "epoch": 1, + "train_loss": 0.9894467581201483, + "valid_selection_loss": 0.9033395762626941, + "valid_clean_loss": 0.8919235042163304, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + }, + { + "method": "B0_early_concat", + "seed": 2026, + "epoch": 2, + "train_loss": 0.8127173174310613, + "valid_selection_loss": 0.8722499026047006, + "valid_clean_loss": 0.8655281905289535, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + }, + { + "method": "B0_early_concat", + "seed": 2026, + "epoch": 3, + "train_loss": 0.7274277905623118, + "valid_selection_loss": 0.9010504417039537, + "valid_clean_loss": 0.9002949733000535, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + }, + { + "method": "B0_early_concat", + "seed": 2026, + "epoch": 4, + "train_loss": 0.6712073153919644, + "valid_selection_loss": 0.9205121268610378, + "valid_clean_loss": 0.9208705294263232, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + }, + { + "method": "B0_early_concat", + "seed": 2026, + "epoch": 5, + "train_loss": 0.6352586569609465, + "valid_selection_loss": 0.9440761974879672, + "valid_clean_loss": 0.9462754883608975, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + } + ], + "mask_counts": { + "0.1/async": 795, + "0.0/async": 832, + "0.5/sync": 860, + "0.0/sync": 859, + "0.5/single": 824, + "0.0/single": 878, + "0.7/async": 880, + "0.0/partial": 860, + "0.7/sync": 842, + "0.5/partial": 867, + "0.3/partial": 815, + "0.1/partial": 874, + "0.3/single": 837, + "0.3/async": 873, + "0.5/async": 830, + "0.1/single": 886, + "0.7/single": 801, + "0.3/sync": 841, + "0.1/sync": 847, + "0.7/partial": 874 + } +} \ No newline at end of file diff --git a/final/experiments/q3/ati_ho/models/B0_early_concat/seed_3407/model_best.pt b/final/experiments/q3/ati_ho/models/B0_early_concat/seed_3407/model_best.pt new file mode 100644 index 0000000..1a93d6d Binary files /dev/null and b/final/experiments/q3/ati_ho/models/B0_early_concat/seed_3407/model_best.pt differ diff --git a/final/experiments/q3/ati_ho/models/B0_early_concat/seed_3407/training_history.csv b/final/experiments/q3/ati_ho/models/B0_early_concat/seed_3407/training_history.csv new file mode 100644 index 0000000..49e90f3 --- /dev/null +++ b/final/experiments/q3/ati_ho/models/B0_early_concat/seed_3407/training_history.csv @@ -0,0 +1,7 @@ +method,seed,epoch,train_loss,valid_selection_loss,valid_clean_loss,lambda_interaction,lambda_mask +B0_early_concat,3407,1,0.9901614431981687,0.9334849718507829,0.9280904976876228,0.0,0.0 +B0_early_concat,3407,2,0.8284304141998291,0.8979402235248587,0.8948127568423093,0.0,0.0 +B0_early_concat,3407,3,0.7437613738907708,0.8850730860626306,0.8822627912510882,0.0,0.0 +B0_early_concat,3407,4,0.682034013999833,0.8985609801915976,0.8979688529129867,0.0,0.0 +B0_early_concat,3407,5,0.6572728455066681,0.9590781547211028,0.9659944576221507,0.0,0.0 +B0_early_concat,3407,6,0.5803193593466723,0.9992094658888303,1.013474199798081,0.0,0.0 diff --git a/final/experiments/q3/ati_ho/models/B0_early_concat/seed_3407/training_manifest.json b/final/experiments/q3/ati_ho/models/B0_early_concat/seed_3407/training_manifest.json new file mode 100644 index 0000000..f56f287 --- /dev/null +++ b/final/experiments/q3/ati_ho/models/B0_early_concat/seed_3407/training_manifest.json @@ -0,0 +1,128 @@ +{ + "method": "B0_early_concat", + "seed": 3407, + "best_epoch": 3, + "best_selection_loss": 0.8850730860626306, + "elapsed_seconds": 3.5102919560158625, + "batch_size": 64, + "epoch_limit": 12, + "patience": 3, + "optimizer": "AdamW", + "learning_rate": 0.0003, + "weight_decay": 0.001, + "gradient_clip_norm": 1.0, + "training_mask_rates": [ + 0.0, + 0.1, + 0.3, + 0.5, + 0.7 + ], + "training_mask_patterns": [ + "single", + "sync", + "partial", + "async" + ], + "training_mask_seed_base": 20261227, + "same_orders_and_masks_across_methods_for_same_seed": true, + "config": { + "model_config": { + "router": "mlp", + "expert_names": [ + "T", + "A", + "V", + "TA", + "TV", + "AV", + "TAV" + ], + "availability_mode": "hard" + } + }, + "history": [ + { + "method": "B0_early_concat", + "seed": 3407, + "epoch": 1, + "train_loss": 0.9901614431981687, + "valid_selection_loss": 0.9334849718507829, + "valid_clean_loss": 0.9280904976876228, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + }, + { + "method": "B0_early_concat", + "seed": 3407, + "epoch": 2, + "train_loss": 0.8284304141998291, + "valid_selection_loss": 0.8979402235248587, + "valid_clean_loss": 0.8948127568423093, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + }, + { + "method": "B0_early_concat", + "seed": 3407, + "epoch": 3, + "train_loss": 0.7437613738907708, + "valid_selection_loss": 0.8850730860626306, + "valid_clean_loss": 0.8822627912510882, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + }, + { + "method": "B0_early_concat", + "seed": 3407, + "epoch": 4, + "train_loss": 0.682034013999833, + "valid_selection_loss": 0.8985609801915976, + "valid_clean_loss": 0.8979688529129867, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + }, + { + "method": "B0_early_concat", + "seed": 3407, + "epoch": 5, + "train_loss": 0.6572728455066681, + "valid_selection_loss": 0.9590781547211028, + "valid_clean_loss": 0.9659944576221507, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + }, + { + "method": "B0_early_concat", + "seed": 3407, + "epoch": 6, + "train_loss": 0.5803193593466723, + "valid_selection_loss": 0.9992094658888303, + "valid_clean_loss": 1.013474199798081, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + } + ], + "mask_counts": { + "0.7/async": 1003, + "0.0/async": 1014, + "0.5/sync": 1033, + "0.1/sync": 1024, + "0.5/partial": 996, + "0.3/partial": 1019, + "0.5/async": 1053, + "0.3/single": 1003, + "0.0/sync": 1042, + "0.7/partial": 1079, + "0.7/sync": 986, + "0.0/single": 1029, + "0.1/async": 959, + "0.3/sync": 1055, + "0.0/partial": 1016, + "0.1/single": 1011, + "0.1/partial": 988, + "0.5/single": 1076, + "0.7/single": 997, + "0.3/async": 987 + } +} \ No newline at end of file diff --git a/final/experiments/q3/ati_ho/models/B0_early_concat/seed_42/model_best.pt b/final/experiments/q3/ati_ho/models/B0_early_concat/seed_42/model_best.pt new file mode 100644 index 0000000..1f958ac Binary files /dev/null and b/final/experiments/q3/ati_ho/models/B0_early_concat/seed_42/model_best.pt differ diff --git a/final/experiments/q3/ati_ho/models/B0_early_concat/seed_42/training_history.csv b/final/experiments/q3/ati_ho/models/B0_early_concat/seed_42/training_history.csv new file mode 100644 index 0000000..6e94854 --- /dev/null +++ b/final/experiments/q3/ati_ho/models/B0_early_concat/seed_42/training_history.csv @@ -0,0 +1,6 @@ +method,seed,epoch,train_loss,valid_selection_loss,valid_clean_loss,lambda_interaction,lambda_mask +B0_early_concat,42,1,0.9883709682358636,0.9093860420551929,0.9040598345326853,0.0,0.0 +B0_early_concat,42,2,0.8108003448556971,0.8755410200619436,0.8639584504641019,0.0,0.0 +B0_early_concat,42,3,0.7378355805520658,0.8942322922604424,0.8897933003666637,0.0,0.0 +B0_early_concat,42,4,0.6694756401357828,0.9000292073239337,0.8993391833462558,0.0,0.0 +B0_early_concat,42,5,0.6367529040133512,0.9381853456680591,0.9431815828595843,0.0,0.0 diff --git a/final/experiments/q3/ati_ho/models/B0_early_concat/seed_42/training_manifest.json b/final/experiments/q3/ati_ho/models/B0_early_concat/seed_42/training_manifest.json new file mode 100644 index 0000000..9f445ae --- /dev/null +++ b/final/experiments/q3/ati_ho/models/B0_early_concat/seed_42/training_manifest.json @@ -0,0 +1,118 @@ +{ + "method": "B0_early_concat", + "seed": 42, + "best_epoch": 2, + "best_selection_loss": 0.8755410200619436, + "elapsed_seconds": 6.863459730986506, + "batch_size": 64, + "epoch_limit": 12, + "patience": 3, + "optimizer": "AdamW", + "learning_rate": 0.0003, + "weight_decay": 0.001, + "gradient_clip_norm": 1.0, + "training_mask_rates": [ + 0.0, + 0.1, + 0.3, + 0.5, + 0.7 + ], + "training_mask_patterns": [ + "single", + "sync", + "partial", + "async" + ], + "training_mask_seed_base": 20261227, + "same_orders_and_masks_across_methods_for_same_seed": true, + "config": { + "model_config": { + "router": "mlp", + "expert_names": [ + "T", + "A", + "V", + "TA", + "TV", + "AV", + "TAV" + ], + "availability_mode": "hard" + } + }, + "history": [ + { + "method": "B0_early_concat", + "seed": 42, + "epoch": 1, + "train_loss": 0.9883709682358636, + "valid_selection_loss": 0.9093860420551929, + "valid_clean_loss": 0.9040598345326853, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + }, + { + "method": "B0_early_concat", + "seed": 42, + "epoch": 2, + "train_loss": 0.8108003448556971, + "valid_selection_loss": 0.8755410200619436, + "valid_clean_loss": 0.8639584504641019, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + }, + { + "method": "B0_early_concat", + "seed": 42, + "epoch": 3, + "train_loss": 0.7378355805520658, + "valid_selection_loss": 0.8942322922604424, + "valid_clean_loss": 0.8897933003666637, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + }, + { + "method": "B0_early_concat", + "seed": 42, + "epoch": 4, + "train_loss": 0.6694756401357828, + "valid_selection_loss": 0.9000292073239337, + "valid_clean_loss": 0.8993391833462558, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + }, + { + "method": "B0_early_concat", + "seed": 42, + "epoch": 5, + "train_loss": 0.6367529040133512, + "valid_selection_loss": 0.9381853456680591, + "valid_clean_loss": 0.9431815828595843, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + } + ], + "mask_counts": { + "0.0/single": 838, + "0.3/single": 841, + "0.5/single": 832, + "0.7/sync": 931, + "0.7/partial": 854, + "0.0/async": 818, + "0.3/async": 856, + "0.3/partial": 861, + "0.3/sync": 841, + "0.1/single": 857, + "0.5/async": 832, + "0.7/single": 872, + "0.5/sync": 848, + "0.5/partial": 854, + "0.0/partial": 870, + "0.0/sync": 827, + "0.7/async": 831, + "0.1/sync": 839, + "0.1/async": 847, + "0.1/partial": 826 + } +} \ No newline at end of file diff --git a/final/experiments/q3/ati_ho/models/B5_mofe_mlp/seed_2026/model_best.pt b/final/experiments/q3/ati_ho/models/B5_mofe_mlp/seed_2026/model_best.pt new file mode 100644 index 0000000..a03d385 Binary files /dev/null and b/final/experiments/q3/ati_ho/models/B5_mofe_mlp/seed_2026/model_best.pt differ diff --git a/final/experiments/q3/ati_ho/models/B5_mofe_mlp/seed_2026/training_history.csv b/final/experiments/q3/ati_ho/models/B5_mofe_mlp/seed_2026/training_history.csv new file mode 100644 index 0000000..22be9d5 --- /dev/null +++ b/final/experiments/q3/ati_ho/models/B5_mofe_mlp/seed_2026/training_history.csv @@ -0,0 +1,7 @@ +method,seed,epoch,train_loss,valid_selection_loss,valid_clean_loss,lambda_interaction,lambda_mask +B5_mofe_mlp,2026,1,1.0356682870123122,0.9561890172106879,0.9502097872587351,0.0,0.0 +B5_mofe_mlp,2026,2,0.8621677259604136,0.8906986942658057,0.8774622966954996,0.0,0.0 +B5_mofe_mlp,2026,3,0.7639763178648772,0.8864746625934329,0.8795440236290732,0.0,0.0 +B5_mofe_mlp,2026,4,0.7062149555594833,0.8968080696496334,0.898658872960688,0.0,0.0 +B5_mofe_mlp,2026,5,0.6749658584594727,0.9011682678054977,0.9081905451449719,0.0,0.0 +B5_mofe_mlp,2026,6,0.6303896639082167,0.9725698265400562,0.9927981803705405,0.0,0.0 diff --git a/final/experiments/q3/ati_ho/models/B5_mofe_mlp/seed_2026/training_manifest.json b/final/experiments/q3/ati_ho/models/B5_mofe_mlp/seed_2026/training_manifest.json new file mode 100644 index 0000000..3e883b9 --- /dev/null +++ b/final/experiments/q3/ati_ho/models/B5_mofe_mlp/seed_2026/training_manifest.json @@ -0,0 +1,128 @@ +{ + "method": "B5_mofe_mlp", + "seed": 2026, + "best_epoch": 3, + "best_selection_loss": 0.8864746625934329, + "elapsed_seconds": 4.902324867958669, + "batch_size": 64, + "epoch_limit": 12, + "patience": 3, + "optimizer": "AdamW", + "learning_rate": 0.0003, + "weight_decay": 0.001, + "gradient_clip_norm": 1.0, + "training_mask_rates": [ + 0.0, + 0.1, + 0.3, + 0.5, + 0.7 + ], + "training_mask_patterns": [ + "single", + "sync", + "partial", + "async" + ], + "training_mask_seed_base": 20261227, + "same_orders_and_masks_across_methods_for_same_seed": true, + "config": { + "model_config": { + "router": "mlp", + "expert_names": [ + "T", + "A", + "V", + "TA", + "TV", + "AV", + "TAV" + ], + "availability_mode": "hard" + } + }, + "history": [ + { + "method": "B5_mofe_mlp", + "seed": 2026, + "epoch": 1, + "train_loss": 1.0356682870123122, + "valid_selection_loss": 0.9561890172106879, + "valid_clean_loss": 0.9502097872587351, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + }, + { + "method": "B5_mofe_mlp", + "seed": 2026, + "epoch": 2, + "train_loss": 0.8621677259604136, + "valid_selection_loss": 0.8906986942658057, + "valid_clean_loss": 0.8774622966954996, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + }, + { + "method": "B5_mofe_mlp", + "seed": 2026, + "epoch": 3, + "train_loss": 0.7639763178648772, + "valid_selection_loss": 0.8864746625934329, + "valid_clean_loss": 0.8795440236290732, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + }, + { + "method": "B5_mofe_mlp", + "seed": 2026, + "epoch": 4, + "train_loss": 0.7062149555594833, + "valid_selection_loss": 0.8968080696496334, + "valid_clean_loss": 0.898658872960688, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + }, + { + "method": "B5_mofe_mlp", + "seed": 2026, + "epoch": 5, + "train_loss": 0.6749658584594727, + "valid_selection_loss": 0.9011682678054977, + "valid_clean_loss": 0.9081905451449719, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + }, + { + "method": "B5_mofe_mlp", + "seed": 2026, + "epoch": 6, + "train_loss": 0.6303896639082167, + "valid_selection_loss": 0.9725698265400562, + "valid_clean_loss": 0.9927981803705405, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + } + ], + "mask_counts": { + "0.1/async": 970, + "0.0/async": 995, + "0.5/sync": 1040, + "0.0/sync": 1012, + "0.5/single": 999, + "0.0/single": 1046, + "0.7/async": 1048, + "0.0/partial": 1063, + "0.7/sync": 1020, + "0.5/partial": 1022, + "0.3/partial": 1002, + "0.1/partial": 1022, + "0.3/single": 1004, + "0.3/async": 1026, + "0.5/async": 1028, + "0.1/single": 1048, + "0.7/single": 959, + "0.3/sync": 1030, + "0.1/sync": 1007, + "0.7/partial": 1029 + } +} \ No newline at end of file diff --git a/final/experiments/q3/ati_ho/models/B5_mofe_mlp/seed_3407/model_best.pt b/final/experiments/q3/ati_ho/models/B5_mofe_mlp/seed_3407/model_best.pt new file mode 100644 index 0000000..bb6fd17 Binary files /dev/null and b/final/experiments/q3/ati_ho/models/B5_mofe_mlp/seed_3407/model_best.pt differ diff --git a/final/experiments/q3/ati_ho/models/B5_mofe_mlp/seed_3407/training_history.csv b/final/experiments/q3/ati_ho/models/B5_mofe_mlp/seed_3407/training_history.csv new file mode 100644 index 0000000..6f6d9c8 --- /dev/null +++ b/final/experiments/q3/ati_ho/models/B5_mofe_mlp/seed_3407/training_history.csv @@ -0,0 +1,8 @@ +method,seed,epoch,train_loss,valid_selection_loss,valid_clean_loss,lambda_interaction,lambda_mask +B5_mofe_mlp,3407,1,1.0173597953937672,0.9501426021803867,0.9452906283703479,0.0,0.0 +B5_mofe_mlp,3407,2,0.8566740729190685,0.922284932254435,0.9245965945851672,0.0,0.0 +B5_mofe_mlp,3407,3,0.7812010469260039,0.8953835820103739,0.901839215021867,0.0,0.0 +B5_mofe_mlp,3407,4,0.7314516603946686,0.8908331908367492,0.8956135635847574,0.0,0.0 +B5_mofe_mlp,3407,5,0.7109767960177528,0.941465525836735,0.9525768114970281,0.0,0.0 +B5_mofe_mlp,3407,6,0.6440567677771604,0.9769726883579086,1.0012760489851564,0.0,0.0 +B5_mofe_mlp,3407,7,0.609118523421111,1.0052462698339106,1.0282764369314843,0.0,0.0 diff --git a/final/experiments/q3/ati_ho/models/B5_mofe_mlp/seed_3407/training_manifest.json b/final/experiments/q3/ati_ho/models/B5_mofe_mlp/seed_3407/training_manifest.json new file mode 100644 index 0000000..99e9ab5 --- /dev/null +++ b/final/experiments/q3/ati_ho/models/B5_mofe_mlp/seed_3407/training_manifest.json @@ -0,0 +1,138 @@ +{ + "method": "B5_mofe_mlp", + "seed": 3407, + "best_epoch": 4, + "best_selection_loss": 0.8908331908367492, + "elapsed_seconds": 5.8394740910152905, + "batch_size": 64, + "epoch_limit": 12, + "patience": 3, + "optimizer": "AdamW", + "learning_rate": 0.0003, + "weight_decay": 0.001, + "gradient_clip_norm": 1.0, + "training_mask_rates": [ + 0.0, + 0.1, + 0.3, + 0.5, + 0.7 + ], + "training_mask_patterns": [ + "single", + "sync", + "partial", + "async" + ], + "training_mask_seed_base": 20261227, + "same_orders_and_masks_across_methods_for_same_seed": true, + "config": { + "model_config": { + "router": "mlp", + "expert_names": [ + "T", + "A", + "V", + "TA", + "TV", + "AV", + "TAV" + ], + "availability_mode": "hard" + } + }, + "history": [ + { + "method": "B5_mofe_mlp", + "seed": 3407, + "epoch": 1, + "train_loss": 1.0173597953937672, + "valid_selection_loss": 0.9501426021803867, + "valid_clean_loss": 0.9452906283703479, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + }, + { + "method": "B5_mofe_mlp", + "seed": 3407, + "epoch": 2, + "train_loss": 0.8566740729190685, + "valid_selection_loss": 0.922284932254435, + "valid_clean_loss": 0.9245965945851672, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + }, + { + "method": "B5_mofe_mlp", + "seed": 3407, + "epoch": 3, + "train_loss": 0.7812010469260039, + "valid_selection_loss": 0.8953835820103739, + "valid_clean_loss": 0.901839215021867, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + }, + { + "method": "B5_mofe_mlp", + "seed": 3407, + "epoch": 4, + "train_loss": 0.7314516603946686, + "valid_selection_loss": 0.8908331908367492, + "valid_clean_loss": 0.8956135635847574, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + }, + { + "method": "B5_mofe_mlp", + "seed": 3407, + "epoch": 5, + "train_loss": 0.7109767960177528, + "valid_selection_loss": 0.941465525836735, + "valid_clean_loss": 0.9525768114970281, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + }, + { + "method": "B5_mofe_mlp", + "seed": 3407, + "epoch": 6, + "train_loss": 0.6440567677771604, + "valid_selection_loss": 0.9769726883579086, + "valid_clean_loss": 1.0012760489851564, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + }, + { + "method": "B5_mofe_mlp", + "seed": 3407, + "epoch": 7, + "train_loss": 0.609118523421111, + "valid_selection_loss": 1.0052462698339106, + "valid_clean_loss": 1.0282764369314843, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + } + ], + "mask_counts": { + "0.7/async": 1162, + "0.0/async": 1173, + "0.5/sync": 1215, + "0.1/sync": 1215, + "0.5/partial": 1175, + "0.3/partial": 1195, + "0.5/async": 1228, + "0.3/single": 1162, + "0.0/sync": 1197, + "0.7/partial": 1263, + "0.7/sync": 1165, + "0.0/single": 1199, + "0.1/async": 1114, + "0.3/sync": 1218, + "0.0/partial": 1182, + "0.1/single": 1187, + "0.1/partial": 1158, + "0.5/single": 1238, + "0.7/single": 1161, + "0.3/async": 1158 + } +} \ No newline at end of file diff --git a/final/experiments/q3/ati_ho/models/B5_mofe_mlp/seed_42/model_best.pt b/final/experiments/q3/ati_ho/models/B5_mofe_mlp/seed_42/model_best.pt new file mode 100644 index 0000000..e16cbf5 Binary files /dev/null and b/final/experiments/q3/ati_ho/models/B5_mofe_mlp/seed_42/model_best.pt differ diff --git a/final/experiments/q3/ati_ho/models/B5_mofe_mlp/seed_42/training_history.csv b/final/experiments/q3/ati_ho/models/B5_mofe_mlp/seed_42/training_history.csv new file mode 100644 index 0000000..6b3bef2 --- /dev/null +++ b/final/experiments/q3/ati_ho/models/B5_mofe_mlp/seed_42/training_history.csv @@ -0,0 +1,8 @@ +method,seed,epoch,train_loss,valid_selection_loss,valid_clean_loss,lambda_interaction,lambda_mask +B5_mofe_mlp,42,1,1.0171989544674203,0.9322710561228322,0.9271787070966029,0.0,0.0 +B5_mofe_mlp,42,2,0.8356565082514728,0.9130383135525735,0.9016369266824408,0.0,0.0 +B5_mofe_mlp,42,3,0.7714545528093973,0.9085815730658207,0.9123272260466775,0.0,0.0 +B5_mofe_mlp,42,4,0.7048188387243836,0.9027628702121776,0.9053291953527011,0.0,0.0 +B5_mofe_mlp,42,5,0.6746232868344696,0.9353020928063236,0.9397049303893205,0.0,0.0 +B5_mofe_mlp,42,6,0.6332692697092339,1.025483563706115,1.0391453098464798,0.0,0.0 +B5_mofe_mlp,42,7,0.5814774891844502,1.0453854943369771,1.060925617322817,0.0,0.0 diff --git a/final/experiments/q3/ati_ho/models/B5_mofe_mlp/seed_42/training_manifest.json b/final/experiments/q3/ati_ho/models/B5_mofe_mlp/seed_42/training_manifest.json new file mode 100644 index 0000000..3da8ca3 --- /dev/null +++ b/final/experiments/q3/ati_ho/models/B5_mofe_mlp/seed_42/training_manifest.json @@ -0,0 +1,138 @@ +{ + "method": "B5_mofe_mlp", + "seed": 42, + "best_epoch": 4, + "best_selection_loss": 0.9027628702121776, + "elapsed_seconds": 6.507645265024621, + "batch_size": 64, + "epoch_limit": 12, + "patience": 3, + "optimizer": "AdamW", + "learning_rate": 0.0003, + "weight_decay": 0.001, + "gradient_clip_norm": 1.0, + "training_mask_rates": [ + 0.0, + 0.1, + 0.3, + 0.5, + 0.7 + ], + "training_mask_patterns": [ + "single", + "sync", + "partial", + "async" + ], + "training_mask_seed_base": 20261227, + "same_orders_and_masks_across_methods_for_same_seed": true, + "config": { + "model_config": { + "router": "mlp", + "expert_names": [ + "T", + "A", + "V", + "TA", + "TV", + "AV", + "TAV" + ], + "availability_mode": "hard" + } + }, + "history": [ + { + "method": "B5_mofe_mlp", + "seed": 42, + "epoch": 1, + "train_loss": 1.0171989544674203, + "valid_selection_loss": 0.9322710561228322, + "valid_clean_loss": 0.9271787070966029, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + }, + { + "method": "B5_mofe_mlp", + "seed": 42, + "epoch": 2, + "train_loss": 0.8356565082514728, + "valid_selection_loss": 0.9130383135525735, + "valid_clean_loss": 0.9016369266824408, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + }, + { + "method": "B5_mofe_mlp", + "seed": 42, + "epoch": 3, + "train_loss": 0.7714545528093973, + "valid_selection_loss": 0.9085815730658207, + "valid_clean_loss": 0.9123272260466775, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + }, + { + "method": "B5_mofe_mlp", + "seed": 42, + "epoch": 4, + "train_loss": 0.7048188387243836, + "valid_selection_loss": 0.9027628702121776, + "valid_clean_loss": 0.9053291953527011, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + }, + { + "method": "B5_mofe_mlp", + "seed": 42, + "epoch": 5, + "train_loss": 0.6746232868344696, + "valid_selection_loss": 0.9353020928063236, + "valid_clean_loss": 0.9397049303893205, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + }, + { + "method": "B5_mofe_mlp", + "seed": 42, + "epoch": 6, + "train_loss": 0.6332692697092339, + "valid_selection_loss": 1.025483563706115, + "valid_clean_loss": 1.0391453098464798, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + }, + { + "method": "B5_mofe_mlp", + "seed": 42, + "epoch": 7, + "train_loss": 0.5814774891844502, + "valid_selection_loss": 1.0453854943369771, + "valid_clean_loss": 1.060925617322817, + "lambda_interaction": 0.0, + "lambda_mask": 0.0 + } + ], + "mask_counts": { + "0.0/single": 1169, + "0.3/single": 1187, + "0.5/single": 1158, + "0.7/sync": 1272, + "0.7/partial": 1207, + "0.0/async": 1145, + "0.3/async": 1200, + "0.3/partial": 1209, + "0.3/sync": 1170, + "0.1/single": 1205, + "0.5/async": 1185, + "0.7/single": 1232, + "0.5/sync": 1143, + "0.5/partial": 1236, + "0.0/partial": 1183, + "0.0/sync": 1159, + "0.7/async": 1179, + "0.1/sync": 1202, + "0.1/async": 1192, + "0.1/partial": 1132 + } +} \ No newline at end of file diff --git a/final/experiments/q3/ati_ho/models/D0/seed_42/model_best.pt b/final/experiments/q3/ati_ho/models/D0/seed_42/model_best.pt new file mode 100644 index 0000000..187d50c Binary files /dev/null and b/final/experiments/q3/ati_ho/models/D0/seed_42/model_best.pt differ diff --git a/final/experiments/q3/ati_ho/models/D0/seed_42/training_history.csv b/final/experiments/q3/ati_ho/models/D0/seed_42/training_history.csv new file mode 100644 index 0000000..3dd02c1 --- /dev/null +++ b/final/experiments/q3/ati_ho/models/D0/seed_42/training_history.csv @@ -0,0 +1,6 @@ +method,seed,epoch,train_loss,valid_selection_loss,valid_clean_loss,lambda_interaction,lambda_mask +D0,42,1,0.9736692993729202,0.9111956561004722,0.9047908566810272,0.001,0.0 +D0,42,2,0.7995399965180291,0.8630276628575482,0.8526584676333836,0.001,0.0 +D0,42,3,0.7263524250851737,0.8714345292403147,0.863228269985744,0.001,0.0 +D0,42,4,0.6683531569110023,0.871083995961881,0.8627031235904484,0.001,0.0 +D0,42,5,0.6370693781861553,0.8886248764100966,0.8807592830815159,0.001,0.0 diff --git a/final/experiments/q3/ati_ho/models/D0/seed_42/training_manifest.json b/final/experiments/q3/ati_ho/models/D0/seed_42/training_manifest.json new file mode 100644 index 0000000..7b28191 --- /dev/null +++ b/final/experiments/q3/ati_ho/models/D0/seed_42/training_manifest.json @@ -0,0 +1,117 @@ +{ + "method": "D0", + "seed": 42, + "best_epoch": 2, + "best_selection_loss": 0.8630276628575482, + "elapsed_seconds": 8.120354533020873, + "batch_size": 64, + "epoch_limit": 12, + "patience": 3, + "optimizer": "AdamW", + "learning_rate": 0.0003, + "weight_decay": 0.001, + "gradient_clip_norm": 1.0, + "training_mask_rates": [ + 0.0, + 0.1, + 0.3, + 0.5, + 0.7 + ], + "training_mask_patterns": [ + "single", + "sync", + "partial", + "async" + ], + "training_mask_seed_base": 20261227, + "same_orders_and_masks_across_methods_for_same_seed": true, + "config": { + "name": "D0_unanchored_diagnostic", + "low_rank": true, + "cross_attention": true, + "anchored": false, + "lambda_interaction": 0.001, + "lambda_mask": 0.0, + "rank": 4, + "hidden": 64, + "gru_hidden_per_direction": 32, + "attention_heads": 4, + "attention_ffn": 128, + "eta_init": 0.1 + }, + "history": [ + { + "method": "D0", + "seed": 42, + "epoch": 1, + "train_loss": 0.9736692993729202, + "valid_selection_loss": 0.9111956561004722, + "valid_clean_loss": 0.9047908566810272, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "D0", + "seed": 42, + "epoch": 2, + "train_loss": 0.7995399965180291, + "valid_selection_loss": 0.8630276628575482, + "valid_clean_loss": 0.8526584676333836, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "D0", + "seed": 42, + "epoch": 3, + "train_loss": 0.7263524250851737, + "valid_selection_loss": 0.8714345292403147, + "valid_clean_loss": 0.863228269985744, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "D0", + "seed": 42, + "epoch": 4, + "train_loss": 0.6683531569110023, + "valid_selection_loss": 0.871083995961881, + "valid_clean_loss": 0.8627031235904484, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + }, + { + "method": "D0", + "seed": 42, + "epoch": 5, + "train_loss": 0.6370693781861553, + "valid_selection_loss": 0.8886248764100966, + "valid_clean_loss": 0.8807592830815159, + "lambda_interaction": 0.001, + "lambda_mask": 0.0 + } + ], + "mask_counts": { + "0.0/single": 838, + "0.3/single": 841, + "0.5/single": 832, + "0.7/sync": 931, + "0.7/partial": 854, + "0.0/async": 818, + "0.3/async": 856, + "0.3/partial": 861, + "0.3/sync": 841, + "0.1/single": 857, + "0.5/async": 832, + "0.7/single": 872, + "0.5/sync": 848, + "0.5/partial": 854, + "0.0/partial": 870, + "0.0/sync": 827, + "0.7/async": 831, + "0.1/sync": 839, + "0.1/async": 847, + "0.1/partial": 826 + } +} \ No newline at end of file diff --git a/final/experiments/q3/ati_ho/provisional_candidate.json b/final/experiments/q3/ati_ho/provisional_candidate.json new file mode 100644 index 0000000..878df14 --- /dev/null +++ b/final/experiments/q3/ati_ho/provisional_candidate.json @@ -0,0 +1,26 @@ +{ + "stage1_candidates": [ + { + "method": "A2", + "best_selection_loss": 0.8634063961741689, + "best_epoch": 2 + }, + { + "method": "A3", + "best_selection_loss": 0.8637567355737581, + "best_epoch": 2 + }, + { + "method": "A0", + "best_selection_loss": 0.8643234267339601, + "best_epoch": 4 + }, + { + "method": "A1", + "best_selection_loss": 0.8649765662439577, + "best_epoch": 4 + } + ], + "provisional_selected_candidate": "A2", + "selection_rule": "lowest fixed four-scenario ATI task loss on the locked official validation split; seed 42 only in Stage I" +} \ No newline at end of file diff --git a/final/experiments/q3/ati_ho/results/ati_ho/ATI_HO_PAPER.md b/final/experiments/q3/ati_ho/results/ati_ho/ATI_HO_PAPER.md new file mode 100644 index 0000000..84b3325 --- /dev/null +++ b/final/experiments/q3/ati_ho/results/ati_ho/ATI_HO_PAPER.md @@ -0,0 +1,61 @@ +# ATI–HO:基于锚定时间交互与分层 Owen 归因的多模态情感预测 + +## 摘要 + +本文在复杂场景多模态情感识别的第三问中实现 ATI–HO,并以官方未对齐输入和统一 Q1 adapter 为基础训练。实验包含 EarlyConcat + BiGRU、MoFE-7 + MLP Router,以及 ATI 主效应、低秩 pairwise、锚定 cross-attention 和可见性掩码辅助消融。ATI–HO 的最终方案由锁定验证集选择为 **A0**,三 seed 固定场景验证损失均值最小。最终模型在 728 条验证样本上的解析 Shapley 与 8 联盟精确枚举通过率为 100.000%。这里的结果支持“输出参数存在可核验的加和分解”,不构成对情绪因果机制的证明。 + +## 1. 问题与方法 + +给定 Text、Audio、Vision 三路 50 步相对进度序列及逐步可见掩码,预测 negative/neutral/positive 类别与 [-3,3] 强度。每个模态使用私有投影、双向 GRU(每方向 32 隐单元)和注意力池化。主效应以显式空输入前向相减锚定为零。ATI 参数向量为三个居中类别 logit 与负/正条件强度参数: + +`ξ = b + Σ_m G_m + Σ_{m torch.Tensor: + """Apply C to the three class logits while leaving magnitude parameters alone.""" + logits = value[..., :3] + logits = logits - logits.mean(dim=-1, keepdim=True) + return torch.cat((logits, value[..., 3:]), dim=-1) + + +class PrivateTemporalEncoder(nn.Module): + """One modality-private projection, BiGRU(32 each way), and attention pool.""" + + def __init__(self, input_dim: int, hidden: int, gru_hidden: int) -> None: + super().__init__() + self.projection = nn.Sequential( + nn.Linear(input_dim, hidden), nn.GELU(), nn.LayerNorm(hidden) + ) + self.temporal = nn.GRU( + input_size=hidden, + hidden_size=gru_hidden, + num_layers=1, + batch_first=True, + bidirectional=True, + ) + self.pool_score = nn.Linear(hidden, 1) + self.output_dim = 2 * gru_hidden + + def forward(self, x: torch.Tensor, observed: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + observed = observed.bool() + projected = self.projection(x) + projected = projected * observed.unsqueeze(-1).to(projected.dtype) + sequence, _ = self.temporal(projected) + sequence = sequence * observed.unsqueeze(-1).to(sequence.dtype) + scores = self.pool_score(torch.tanh(sequence)).squeeze(-1) + scores = scores.masked_fill(~observed, torch.finfo(scores.dtype).min) + has_any = observed.any(dim=1, keepdim=True) + weights = torch.softmax(scores, dim=1) + weights = torch.where(has_any, weights, torch.zeros_like(weights)) + pooled = torch.sum(sequence * weights.unsqueeze(-1), dim=1) + return sequence, pooled + + +class MainEffectHead(nn.Module): + def __init__(self, hidden: int) -> None: + super().__init__() + self.network = nn.Sequential(nn.Linear(hidden, hidden), nn.GELU(), nn.Linear(hidden, 5)) + + def forward(self, pooled: torch.Tensor) -> torch.Tensor: + return _center_class_parameters(self.network(pooled)) + + +class AnchoredPairBranch(nn.Module): + """A pair reads only two private streams; its four-term anchor is explicit.""" + + def __init__(self, hidden: int, config: ATIConfig) -> None: + super().__init__() + self.low_rank_enabled = config.low_rank + self.cross_attention_enabled = config.cross_attention + self.rank = config.rank + if self.low_rank_enabled: + self.left_factor = nn.Linear(hidden, config.rank) + self.right_factor = nn.Linear(hidden, config.rank) + self.low_rank_out = nn.Linear(config.rank, 5, bias=False) + else: + self.left_factor = None + self.right_factor = None + self.low_rank_out = None + + if self.cross_attention_enabled: + self.left_to_right = nn.MultiheadAttention( + hidden, config.attention_heads, batch_first=True + ) + self.right_to_left = nn.MultiheadAttention( + hidden, config.attention_heads, batch_first=True + ) + self.left_norm1 = nn.LayerNorm(hidden) + self.right_norm1 = nn.LayerNorm(hidden) + self.left_ffn = nn.Sequential( + nn.Linear(hidden, config.attention_ffn), + nn.GELU(), + nn.Linear(config.attention_ffn, hidden), + ) + self.right_ffn = nn.Sequential( + nn.Linear(hidden, config.attention_ffn), + nn.GELU(), + nn.Linear(config.attention_ffn, hidden), + ) + self.left_norm2 = nn.LayerNorm(hidden) + self.right_norm2 = nn.LayerNorm(hidden) + self.cross_out = nn.Linear(hidden * 2, 5, bias=False) + else: + self.left_to_right = None + self.right_to_left = None + self.left_norm1 = None + self.right_norm1 = None + self.left_ffn = None + self.right_ffn = None + self.left_norm2 = None + self.right_norm2 = None + self.cross_out = None + + init = min(max(config.eta_init, 1e-5), 1 - 1e-5) + self.eta_logit = nn.Parameter(torch.tensor(math.log(init / (1.0 - init)))) + # q(x0,y)=q(x,y0)=q(x0,y0)=offset. Four-term subtraction cancels it. + # D0 deliberately leaves this offset in the output as a leakage control. + self.anchor_offset = nn.Parameter(torch.zeros(5)) + + @staticmethod + def _masked_mean(sequence: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: + weights = mask.to(sequence.dtype).unsqueeze(-1) + return (sequence * weights).sum(dim=1) / weights.sum(dim=1).clamp_min(1.0) + + @staticmethod + def _safe_key_mask(mask: torch.Tensor) -> torch.Tensor: + safe = mask.clone() + empty = ~safe.any(dim=1) + if empty.any(): + safe[empty, 0] = True + return safe + + def _core( + self, + left: torch.Tensor, + right: torch.Tensor, + left_mask: torch.Tensor, + right_mask: torch.Tensor, + ) -> torch.Tensor: + joint = left_mask.bool() & right_mask.bool() + values: list[torch.Tensor] = [] + if self.low_rank_enabled: + assert self.left_factor is not None and self.right_factor is not None + assert self.low_rank_out is not None + product = torch.tanh(self.left_factor(left)) * torch.tanh(self.right_factor(right)) + values.append(self.low_rank_out(self._masked_mean(product, joint))) + if self.cross_attention_enabled: + assert self.left_to_right is not None and self.right_to_left is not None + assert self.left_norm1 is not None and self.right_norm1 is not None + assert self.left_ffn is not None and self.right_ffn is not None + assert self.left_norm2 is not None and self.right_norm2 is not None + assert self.cross_out is not None + safe_left = self._safe_key_mask(left_mask.bool()) + safe_right = self._safe_key_mask(right_mask.bool()) + left_msg, _ = self.left_to_right( + left, right, right, key_padding_mask=~safe_right, need_weights=False + ) + right_msg, _ = self.right_to_left( + right, left, left, key_padding_mask=~safe_left, need_weights=False + ) + left_context = self.left_norm1(left + left_msg) + right_context = self.right_norm1(right + right_msg) + left_context = self.left_norm2(left_context + self.left_ffn(left_context)) + right_context = self.right_norm2(right_context + self.right_ffn(right_context)) + left_context = left_context * left_mask.unsqueeze(-1).to(left_context.dtype) + right_context = right_context * right_mask.unsqueeze(-1).to(right_context.dtype) + pooled = torch.cat( + (self._masked_mean(left_context, joint), self._masked_mean(right_context, joint)), + dim=-1, + ) + cross = self.cross_out(pooled) + values.append(torch.sigmoid(self.eta_logit) * cross) + if not values: + return left.new_zeros((left.shape[0], 5)) + # Each branch has a bias-free output and a joint-observation gate. Thus + # core(x, y0)=core(x0, y)=core(x0, y0)=0 exactly. + return torch.stack(values, dim=0).sum(dim=0) + + def forward( + self, + left: torch.Tensor, + right: torch.Tensor, + left_mask: torch.Tensor, + right_mask: torch.Tensor, + *, + anchored: bool, + ) -> torch.Tensor: + raw_xy = self._core(left, right, left_mask, right_mask) + self.anchor_offset + if anchored: + # Four-term difference: + # q(x,y)-q(x,x0)-q(x0,y)+q(x0,y0) = core(x,y). + # The three absent-modality terms equal anchor_offset by the + # joint gate and bias-free core, so they cancel algebraically. + value = raw_xy - self.anchor_offset + else: + value = raw_xy + return _center_class_parameters(value) + + +class ATIHOModel(nn.Module): + """Five-parameter additive multimodal predictor with exact modality anchors.""" + + def __init__(self, dims: tuple[int, int, int], config: ATIConfig, steps: int = 50) -> None: + super().__init__() + self.dims = tuple(int(d) for d in dims) + self.steps = int(steps) + self.config = config + hidden = config.hidden + self.encoders = nn.ModuleList( + PrivateTemporalEncoder(dim, hidden, config.gru_hidden_per_direction) for dim in dims + ) + self.main_heads = nn.ModuleList(MainEffectHead(hidden) for _ in dims) + self.mask_heads = nn.ModuleList(nn.Linear(hidden, 1) for _ in dims) + self.pair_branches = nn.ModuleList( + AnchoredPairBranch(hidden, config) for _ in PAIR_INDICES + ) + self.baseline = nn.Parameter(torch.zeros(5)) + + def forward( + self, + xs: tuple[torch.Tensor, torch.Tensor, torch.Tensor], + masks: torch.Tensor, + *, + return_details: bool = True, + ) -> dict[str, Any]: + if len(xs) != 3: + raise ValueError("ATI–HO requires text, audio, and vision streams") + if masks.ndim != 3 or masks.shape[-1] != 3: + raise ValueError(f"masks must be B x T x 3, got {tuple(masks.shape)}") + if masks.shape[1] > self.steps: + raise ValueError(f"ATI–HO supports at most {self.steps} steps") + masks = masks.bool() + + sequences: list[torch.Tensor] = [] + pooled: list[torch.Tensor] = [] + mask_logits: list[torch.Tensor] = [] + main_effects: list[torch.Tensor] = [] + for modality, (encoder, head, mask_head, x) in enumerate( + zip(self.encoders, self.main_heads, self.mask_heads, xs) + ): + if x.shape[-1] != self.dims[modality]: + raise ValueError( + f"modality {modality} has {x.shape[-1]} features, expected {self.dims[modality]}" + ) + sequence, representation = encoder(x, masks[..., modality]) + # Missing-mask baseline has a zero pooled representation. Explicit + # subtraction makes every main effect zero at that baseline. + baseline_raw = head(representation.new_zeros(representation.shape)) + effect = _center_class_parameters(head(representation) - baseline_raw) + sequences.append(sequence) + pooled.append(representation) + mask_logits.append(mask_head(sequence).squeeze(-1)) + main_effects.append(effect) + + pair_effects: list[torch.Tensor] = [] + pair_penalties: list[torch.Tensor] = [] + for branch, (left_idx, right_idx) in zip(self.pair_branches, PAIR_INDICES): + pair = branch( + sequences[left_idx], + sequences[right_idx], + masks[..., left_idx], + masks[..., right_idx], + anchored=self.config.anchored, + ) + pair_effects.append(pair) + pair_penalties.append(pair.square().mean()) + + main_tensor = torch.stack(main_effects, dim=1) + pair_tensor = torch.stack(pair_effects, dim=1) + params = self.baseline.unsqueeze(0) + main_tensor.sum(dim=1) + pair_tensor.sum(dim=1) + 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) + hard_intensity = torch.where( + predicted_class == 0, + -nu_negative, + torch.where(predicted_class == 2, nu_positive, torch.zeros_like(nu_positive)), + ) + soft_intensity = probabilities[:, 2] * nu_positive - probabilities[:, 0] * nu_negative + result: dict[str, Any] = { + "logits": logits, + "probabilities": probabilities, + "predicted_class": predicted_class, + "intensity": hard_intensity, + "soft_intensity": soft_intensity, + "nu_negative": nu_negative, + "nu_positive": nu_positive, + "params": params, + "interaction_penalty": torch.stack(pair_penalties).mean(), + "mask_logits": torch.stack(mask_logits, dim=-1), + } + if return_details: + result.update( + { + "baseline": self.baseline.unsqueeze(0).expand(xs[0].shape[0], -1), + "main_effects": main_tensor, + "pair_effects": pair_tensor, + "main_sequences": torch.stack(sequences, dim=1), + } + ) + return result + + +def task_loss( + output: dict[str, Any], + y_cls: torch.Tensor, + y_reg: torch.Tensor, + *, + lambda_interaction: float, + lambda_mask: float, + mask_target: torch.Tensor | None = None, +) -> tuple[torch.Tensor, dict[str, torch.Tensor]]: + """CE + conditional polarity magnitude + low-weight continuous Huber.""" + class_loss = F.cross_entropy(output["logits"], y_cls) + negative = y_reg < 0 + positive = y_reg > 0 + target_mag = torch.abs(y_reg) / 3.0 + magnitude_parts: list[torch.Tensor] = [] + if negative.any(): + magnitude_parts.append( + F.smooth_l1_loss(output["nu_negative"][negative] / 3.0, target_mag[negative]) + ) + if positive.any(): + magnitude_parts.append( + F.smooth_l1_loss(output["nu_positive"][positive] / 3.0, target_mag[positive]) + ) + magnitude_loss = torch.stack(magnitude_parts).mean() if magnitude_parts else class_loss.new_zeros(()) + continuous_loss = F.huber_loss( + output["soft_intensity"] / 3.0, y_reg / 3.0, delta=0.25 + ) + interaction_loss = output["interaction_penalty"] + mask_loss = class_loss.new_zeros(()) + if lambda_mask > 0: + if mask_target is None: + raise ValueError("mask_target is required when the visibility-mask auxiliary loss is enabled") + mask_loss = F.binary_cross_entropy_with_logits( + output["mask_logits"], mask_target.to(output["mask_logits"].dtype) + ) + total = ( + class_loss + + magnitude_loss + + 0.2 * continuous_loss + + lambda_interaction * interaction_loss + + lambda_mask * mask_loss + ) + parts = { + "classification": class_loss, + "conditional_magnitude": magnitude_loss, + "continuous_huber": continuous_loss, + "interaction": interaction_loss, + "visibility_mask": mask_loss, + "total": total, + } + return total, parts diff --git a/final/model/ati_ho_config.py b/final/model/ati_ho_config.py new file mode 100644 index 0000000..de39a9e --- /dev/null +++ b/final/model/ati_ho_config.py @@ -0,0 +1,35 @@ +from __future__ import annotations + +from dataclasses import asdict, dataclass + + +@dataclass(frozen=True) +class ATIConfig: + name: str + low_rank: bool = False + cross_attention: bool = False + anchored: bool = True + lambda_interaction: float = 1e-3 + lambda_mask: float = 0.0 + rank: int = 4 + hidden: int = 64 + gru_hidden_per_direction: int = 32 + attention_heads: int = 4 + attention_ffn: int = 128 + eta_init: float = 0.1 + + def to_dict(self) -> dict[str, object]: + return asdict(self) + + +CONFIGS: dict[str, ATIConfig] = { + "A0": ATIConfig(name="A0_main_effects"), + "A1": ATIConfig(name="A1_low_rank_pairs", low_rank=True), + "A2": ATIConfig(name="A2_anchored_pairwise", low_rank=True, cross_attention=True), + "A3": ATIConfig( + name="A3_pairwise_mask_aux", low_rank=True, cross_attention=True, lambda_mask=0.05 + ), + "D0": ATIConfig( + name="D0_unanchored_diagnostic", low_rank=True, cross_attention=True, anchored=False + ), +} diff --git a/final/output/README.md b/final/output/README.md index 5082f5b..d55b08b 100644 --- a/final/output/README.md +++ b/final/output/README.md @@ -2,6 +2,6 @@ - [q1](q1/):附件一 100 条样本的三模态特征、全量汇总、典型样本对齐查询和完整性审计。 - [q2](q2/):附件二未对齐数据上全部保留模型的验证对比、选定模型测试结果及缺失曲线汇总。数学训练入口也会输出附件三的 30 条预测和审计。 -- [q3](q3/):Q3 的验证结果、附件四 20 条预测、输入审计、局部解释和典型解释卡生成位置;完整输出由训练入口写入。 +- [q3](q3/):当前 ATI–HO 方案的附件四预测、参数解释、局部 Owen 证据和推理清单;只保留题目输出文件。验证指标与完整实验审计保存在 `experiments/q3/ati_ho/`。 完整训练与复现说明见 [项目 README](../README.md),结果解释见 [REPORTS.md](../REPORTS.md)。附件二原始输入不包含在项目中;需按 README 的数据目录结构单独准备。提交赛题附件前,应按赛题规定检查最终附件体积。 diff --git a/final/output/q2/README.md b/final/output/q2/README.md index 64c8c91..2c97be4 100644 --- a/final/output/q2/README.md +++ b/final/output/q2/README.md @@ -7,5 +7,8 @@ | `comparison_validation.csv` | 全部 14 个数学/深度模型版本在官方验证集的指标 | | `comparison_test.csv` | 按预定选择规则保留的数学方案和两种深度模型测试指标 | | `comparison_aurc.csv` | 单模态、同步、部分重叠、异步缺失模式的归一化 AURC-MAE | +| `attachment3_predictions.csv` | 附件 3 未对齐版 30 个无标签样本的极性、情感强度、类别概率和预测区间 | +| `attachment3_audit.csv` | 输入来源哈希、可见模态位置、缺失率及推理审计 | +| `attachment3_prediction_manifest.json` | 输入版本、所用模型检查点、校准温度与输出清单 | 本目录是汇总结果。完整的权重、运行清单和情景审计位于 `final/experiments/q2/`。复现步骤、字段含义和限制见 [项目 README](../../README.md) 与 [REPORTS.md](../../REPORTS.md)。 diff --git a/final/output/q2/attachment3_audit.csv b/final/output/q2/attachment3_audit.csv new file mode 100644 index 0000000..dbf840e --- /dev/null +++ b/final/output/q2/attachment3_audit.csv @@ -0,0 +1,31 @@ +case_id,source_file,text_visible_steps,audio_visible_steps,vision_visible_steps,audio_missing_fraction,vision_missing_fraction,unknown_quality_flag,labels_available,input_note,source_coordinate_mode,source_sha256,visible_text_steps,visible_audio_steps,visible_vision_steps,low_information_prior_fallback,calibration_temperature,interval_90_lower,interval_90_upper,predictive_variance_mean_uncalibrated,predictive_variance_calibrated +附件3_未对齐版本_01,附件3_未对齐版本_01.pkl,50,49,48,0.020000000000000018,0.040000000000000036,True,False,unaligned features projected by the shared Q1 adapter; raw_text re-encoded with BERT; audio/vision lengths inferred from last nonzero row because attachment 3 omits trusted lengths,relative_progress,4a8cf3d22e80e29cbf042767fbd3b603b466e1f5e2adaaa2536e8a20df05778d,50,49,48,False,1.122980387832455,-2.7946219444274902,1.550757884979248,1.4611132144927979,1.6770378351211548 +附件3_未对齐版本_02,附件3_未对齐版本_02.pkl,50,50,50,0.0,0.0,True,False,unaligned features projected by the shared Q1 adapter; raw_text re-encoded with BERT; audio/vision lengths inferred from last nonzero row because attachment 3 omits trusted lengths,relative_progress,1c99664a44a42fffac167dfdaf21923d0acd402c210c3a8e61bea51bb3d0a9b6,50,50,50,False,1.122980387832455,-1.2931972742080688,1.0875630378723145,0.44498664140701294,0.4485733211040497 +附件3_未对齐版本_03,附件3_未对齐版本_03.pkl,50,50,30,0.0,0.4,True,False,unaligned features projected by the shared Q1 adapter; raw_text re-encoded with BERT; audio/vision lengths inferred from last nonzero row because attachment 3 omits trusted lengths,relative_progress,463185831b192816b256b4feef74054632e697939fc51b39d7d1a16903b613d4,50,50,30,False,1.122980387832455,-0.49461933970451355,1.5512335300445557,0.3907756507396698,0.3993552029132843 +附件3_未对齐版本_04,附件3_未对齐版本_04.pkl,50,50,50,0.0,0.0,True,False,unaligned features projected by the shared Q1 adapter; raw_text re-encoded with BERT; audio/vision lengths inferred from last nonzero row because attachment 3 omits trusted lengths,relative_progress,bcd6d1ddee46e9995806fa20132330f19c24472ae4150c561bb89ae1bcd2f235,50,50,50,False,1.122980387832455,-0.7178608775138855,1.2966943979263306,0.33880671858787537,0.3422609567642212 +附件3_未对齐版本_05,附件3_未对齐版本_05.pkl,50,50,50,0.0,0.0,True,False,unaligned features projected by the shared Q1 adapter; raw_text re-encoded with BERT; audio/vision lengths inferred from last nonzero row because attachment 3 omits trusted lengths,relative_progress,509990f6bf3ad0a14092824273c2de31ad46e48950b2c028dcaabdf965e56f9e,50,50,50,False,1.122980387832455,-2.460880994796753,0.9297983050346375,0.9528027772903442,1.0553690195083618 +附件3_未对齐版本_06,附件3_未对齐版本_06.pkl,50,50,50,0.0,0.0,True,False,unaligned features projected by the shared Q1 adapter; raw_text re-encoded with BERT; audio/vision lengths inferred from last nonzero row because attachment 3 omits trusted lengths,relative_progress,90df7acd34f00150233fa56621bfc0f1adb61fe2a8deea9fa3715d9499d6e61f,50,50,50,False,1.122980387832455,0.0,2.0461437702178955,0.43230101466178894,0.4604739844799042 +附件3_未对齐版本_07,附件3_未对齐版本_07.pkl,50,50,50,0.0,0.0,True,False,unaligned features projected by the shared Q1 adapter; raw_text re-encoded with BERT; audio/vision lengths inferred from last nonzero row because attachment 3 omits trusted lengths,relative_progress,775eac32a8fa8588e963640473f0d64d2c9753d034e0257adfa42959ffd23619,50,50,50,False,1.122980387832455,-0.1362186074256897,1.7385063171386719,0.3990931510925293,0.41373252868652344 +附件3_未对齐版本_08,附件3_未对齐版本_08.pkl,50,50,50,0.0,0.0,True,False,unaligned features projected by the shared Q1 adapter; raw_text re-encoded with BERT; audio/vision lengths inferred from last nonzero row because attachment 3 omits trusted lengths,relative_progress,4d50aa2a374872732e5e53cc8d6c0ca2b16c95f142b78dc291ce4782036b0b25,50,50,50,False,1.122980387832455,-1.6569616794586182,0.9865659475326538,0.5980508923530579,0.6036122441291809 +附件3_未对齐版本_09,附件3_未对齐版本_09.pkl,50,50,50,0.0,0.0,True,False,unaligned features projected by the shared Q1 adapter; raw_text re-encoded with BERT; audio/vision lengths inferred from last nonzero row because attachment 3 omits trusted lengths,relative_progress,868c890476dbb622e5bd8df2469f46ff5e6fb1c546a77d40f8acd45dab8315e8,50,50,50,False,1.122980387832455,-0.03433350846171379,1.4888426065444946,0.2824844419956207,0.2886645495891571 +附件3_未对齐版本_10,附件3_未对齐版本_10.pkl,50,49,49,0.020000000000000018,0.020000000000000018,True,False,unaligned features projected by the shared Q1 adapter; raw_text re-encoded with BERT; audio/vision lengths inferred from last nonzero row because attachment 3 omits trusted lengths,relative_progress,585dcfdb4c3856538732dea7acc25600198308433e6ab4f26e97a2c939d081e6,50,49,49,False,1.122980387832455,-0.5964999794960022,1.0180858373641968,0.2146325260400772,0.21938827633857727 +附件3_未对齐版本_11,附件3_未对齐版本_11.pkl,50,49,46,0.020000000000000018,0.07999999999999996,True,False,unaligned features projected by the shared Q1 adapter; raw_text re-encoded with BERT; audio/vision lengths inferred from last nonzero row because attachment 3 omits trusted lengths,relative_progress,ea43a1cb3d93841c2cfc063ffca804888fcd2341b59386ee26e4e05fceca665f,50,49,46,False,1.122980387832455,-0.5481101274490356,1.093954086303711,0.23209431767463684,0.23480401933193207 +附件3_未对齐版本_12,附件3_未对齐版本_12.pkl,50,50,50,0.0,0.0,True,False,unaligned features projected by the shared Q1 adapter; raw_text re-encoded with BERT; audio/vision lengths inferred from last nonzero row because attachment 3 omits trusted lengths,relative_progress,54db89460b17078acd4aa4683aa06982213deef1b94dfee3effe57f76980c1a9,50,50,50,False,1.122980387832455,-0.7832371592521667,1.5990359783172607,0.48600462079048157,0.49588483572006226 +附件3_未对齐版本_13,附件3_未对齐版本_13.pkl,50,50,50,0.0,0.0,True,False,unaligned features projected by the shared Q1 adapter; raw_text re-encoded with BERT; audio/vision lengths inferred from last nonzero row because attachment 3 omits trusted lengths,relative_progress,f680fff053e4949dbef6e81739da54f961fa6b17642957cf8df536060c58e444,50,50,50,False,1.122980387832455,-1.2298040390014648,1.722484827041626,0.7403188347816467,0.7457730174064636 +附件3_未对齐版本_14,附件3_未对齐版本_14.pkl,50,49,50,0.020000000000000018,0.0,True,False,unaligned features projected by the shared Q1 adapter; raw_text re-encoded with BERT; audio/vision lengths inferred from last nonzero row because attachment 3 omits trusted lengths,relative_progress,36c25a1251f1afce266098671a620a036c61148aa89e466c113f70ded1355f65,50,49,50,False,1.122980387832455,-1.5283849239349365,1.0474128723144531,0.5521396398544312,0.5541964173316956 +附件3_未对齐版本_15,附件3_未对齐版本_15.pkl,50,50,50,0.0,0.0,True,False,unaligned features projected by the shared Q1 adapter; raw_text re-encoded with BERT; audio/vision lengths inferred from last nonzero row because attachment 3 omits trusted lengths,relative_progress,97cd1ac0320af870b6057c6ba0aaa1c4cae377f108bfed3b3b6764f38e0793d9,50,50,50,False,1.122980387832455,0.0,2.2798967361450195,0.47987255454063416,0.5250018835067749 +附件3_未对齐版本_16,附件3_未对齐版本_16.pkl,50,50,50,0.0,0.0,True,False,unaligned features projected by the shared Q1 adapter; raw_text re-encoded with BERT; audio/vision lengths inferred from last nonzero row because attachment 3 omits trusted lengths,relative_progress,4861ab00b303cb7a89b7a50f1ba17a6e0125837c2d64ff8b7d0a366d31926c54,50,50,50,False,1.122980387832455,0.0,2.5877037048339844,0.47664394974708557,0.5373051762580872 +附件3_未对齐版本_17,附件3_未对齐版本_17.pkl,50,50,49,0.0,0.020000000000000018,True,False,unaligned features projected by the shared Q1 adapter; raw_text re-encoded with BERT; audio/vision lengths inferred from last nonzero row because attachment 3 omits trusted lengths,relative_progress,0285c121a125b95434ea319f1938286f92c66ee2ffe4bc9d46c2509b6fad1eee,50,50,49,False,1.122980387832455,-0.42875346541404724,1.4985958337783813,0.3550581634044647,0.3624156415462494 +附件3_未对齐版本_18,附件3_未对齐版本_18.pkl,50,50,50,0.0,0.0,True,False,unaligned features projected by the shared Q1 adapter; raw_text re-encoded with BERT; audio/vision lengths inferred from last nonzero row because attachment 3 omits trusted lengths,relative_progress,f25009fd8965236a5c51f864a6c60a22b92a57c6e41712422e2d6f4a37af2543,50,50,50,False,1.122980387832455,-0.2905426323413849,1.4048357009887695,0.29468634724617004,0.29980164766311646 +附件3_未对齐版本_19,附件3_未对齐版本_19.pkl,50,50,50,0.0,0.0,True,False,unaligned features projected by the shared Q1 adapter; raw_text re-encoded with BERT; audio/vision lengths inferred from last nonzero row because attachment 3 omits trusted lengths,relative_progress,a1ec22a528d865998b9c15817486efa5adb62a3162dbe20814aa1d3bb02b0452,50,50,50,False,1.122980387832455,-1.5002697706222534,0.8932285904884338,0.47930681705474854,0.4835101366043091 +附件3_未对齐版本_20,附件3_未对齐版本_20.pkl,50,50,50,0.0,0.0,True,False,unaligned features projected by the shared Q1 adapter; raw_text re-encoded with BERT; audio/vision lengths inferred from last nonzero row because attachment 3 omits trusted lengths,relative_progress,470ba389947ff86e5e9c05072b9345858690a57cf498dd02dbc1db0f688ee8c0,50,50,50,False,1.122980387832455,-0.7409619092941284,1.3107733726501465,0.3477743864059448,0.3519274592399597 +附件3_未对齐版本_21,附件3_未对齐版本_21.pkl,50,44,45,0.12,0.09999999999999998,True,False,unaligned features projected by the shared Q1 adapter; raw_text re-encoded with BERT; audio/vision lengths inferred from last nonzero row because attachment 3 omits trusted lengths,relative_progress,ce761fc3f77857b265b4500ef98a5439dab3f45df265ee985b3db218b4730755,50,44,45,False,1.122980387832455,0.0,1.706007719039917,0.32180947065353394,0.33276495337486267 +附件3_未对齐版本_22,附件3_未对齐版本_22.pkl,50,50,30,0.0,0.4,True,False,unaligned features projected by the shared Q1 adapter; raw_text re-encoded with BERT; audio/vision lengths inferred from last nonzero row because attachment 3 omits trusted lengths,relative_progress,261cacc583d096a853cc9181768528f2e29a4b701f9f89653517e83dce08bc45,50,50,30,False,1.122980387832455,-1.5561679601669312,1.679451823234558,0.9327147603034973,0.9247239232063293 +附件3_未对齐版本_23,附件3_未对齐版本_23.pkl,50,47,47,0.06000000000000005,0.06000000000000005,True,False,unaligned features projected by the shared Q1 adapter; raw_text re-encoded with BERT; audio/vision lengths inferred from last nonzero row because attachment 3 omits trusted lengths,relative_progress,c1386ab902160eb55d1b8c0aa79e8690693daba7258a4f16b9b1963fffa7b1f2,50,47,47,False,1.122980387832455,0.0,2.1167213916778564,0.40735018253326416,0.4355928599834442 +附件3_未对齐版本_24,附件3_未对齐版本_24.pkl,50,49,47,0.020000000000000018,0.06000000000000005,True,False,unaligned features projected by the shared Q1 adapter; raw_text re-encoded with BERT; audio/vision lengths inferred from last nonzero row because attachment 3 omits trusted lengths,relative_progress,a7f73ddd97e406440abf16336789fec83622884bf5bb1439c3da5c75a1ae9420,50,49,47,False,1.122980387832455,-0.49775230884552,1.5678426027297974,0.39866694808006287,0.40745168924331665 +附件3_未对齐版本_25,附件3_未对齐版本_25.pkl,50,44,40,0.12,0.19999999999999996,True,False,unaligned features projected by the shared Q1 adapter; raw_text re-encoded with BERT; audio/vision lengths inferred from last nonzero row because attachment 3 omits trusted lengths,relative_progress,62d5b24ef181b5d46207b4c48b54452e0fd1e4f7ce6f64945c5ac93a813ba001,50,44,40,False,1.122980387832455,-0.4213324785232544,1.4042198657989502,0.3172762989997864,0.3226009011268616 +附件3_未对齐版本_26,附件3_未对齐版本_26.pkl,50,49,48,0.020000000000000018,0.040000000000000036,True,False,unaligned features projected by the shared Q1 adapter; raw_text re-encoded with BERT; audio/vision lengths inferred from last nonzero row because attachment 3 omits trusted lengths,relative_progress,327c66c63b793e5cac0f652a77f63d6e144d0447867464108b861b441357be50,50,49,48,False,1.122980387832455,-0.012923063710331917,1.4388282299041748,0.2632901668548584,0.26744163036346436 +附件3_未对齐版本_27,附件3_未对齐版本_27.pkl,50,50,50,0.0,0.0,True,False,unaligned features projected by the shared Q1 adapter; raw_text re-encoded with BERT; audio/vision lengths inferred from last nonzero row because attachment 3 omits trusted lengths,relative_progress,cba6f3e088380f2378bc35b989df054ac1aabbaca5c831037ce542f7e5f7a47b,50,50,50,False,1.122980387832455,0.0,2.379847764968872,0.48261168599128723,0.530409574508667 +附件3_未对齐版本_28,附件3_未对齐版本_28.pkl,50,50,49,0.0,0.020000000000000018,True,False,unaligned features projected by the shared Q1 adapter; raw_text re-encoded with BERT; audio/vision lengths inferred from last nonzero row because attachment 3 omits trusted lengths,relative_progress,a4697405668cb31a8493fb591553170d9c0338971b122f8ee5ee0268a4cab91c,50,50,49,False,1.122980387832455,0.0,1.9912506341934204,0.39680054783821106,0.42043906450271606 +附件3_未对齐版本_29,附件3_未对齐版本_29.pkl,50,46,46,0.07999999999999996,0.07999999999999996,True,False,unaligned features projected by the shared Q1 adapter; raw_text re-encoded with BERT; audio/vision lengths inferred from last nonzero row because attachment 3 omits trusted lengths,relative_progress,86105acdca8927a1665c5480e9282306bcf8433fb4be069c7aa1ec5260f47574,50,46,46,False,1.122980387832455,0.0,2.5532665252685547,0.48900747299194336,0.5494855046272278 +附件3_未对齐版本_30,附件3_未对齐版本_30.pkl,50,50,45,0.0,0.09999999999999998,True,False,unaligned features projected by the shared Q1 adapter; raw_text re-encoded with BERT; audio/vision lengths inferred from last nonzero row because attachment 3 omits trusted lengths,relative_progress,3a04ca9fffd401dc476968d4cfd6f1a2cbf86ae52887d339d25265acdf0d23d8,50,50,45,False,1.122980387832455,0.0,1.9468833208084106,0.43973156809806824,0.4657270610332489 diff --git a/final/output/q2/attachment3_prediction_manifest.json b/final/output/q2/attachment3_prediction_manifest.json new file mode 100644 index 0000000..e59ef63 --- /dev/null +++ b/final/output/q2/attachment3_prediction_manifest.json @@ -0,0 +1,12 @@ +{ + "task": "unlabeled Attachment 3 inference", + "input_version": "unaligned_50", + "selected_model": "C6", + "checkpoint_run": "experiments/q2/unaligned_math_all_b128", + "prediction_count": 30, + "temperature": 1.122980387832455, + "labels_available": false, + "prediction_file": "attachment3_predictions.csv", + "audit_file": "attachment3_audit.csv", + "completed_utc": "2026-09-26T06:06:53Z" +} \ No newline at end of file diff --git a/final/output/q2/attachment3_predictions.csv b/final/output/q2/attachment3_predictions.csv new file mode 100644 index 0000000..7a4d126 --- /dev/null +++ b/final/output/q2/attachment3_predictions.csv @@ -0,0 +1,31 @@ +case_id,predicted_class,predicted_class_name,predicted_sentiment,p_negative,p_neutral,p_positive,interval_90_lower,interval_90_upper,predictive_variance_mean_uncalibrated,within_trajectory_variance,between_trajectory_variance,predictive_mean_calibrated,predictive_variance_calibrated,beta_negative_alpha,beta_negative_beta,beta_positive_alpha,beta_positive_beta,low_information_prior_fallback,output_note +附件3_未对齐版本_01,0,negative,-2.161813497543335,0.7967485189437866,0.1071634516119957,0.0960879996418953,-2.7946219444274902,1.550757884979248,1.4611132144927979,1.4611130952835083,4.174087564479123e-08,-1.5156621932983398,1.6770378351211548,4.154862403869629,1.8028476238250732,3.133524179458618,2.8363540172576904,False,unlabeled attachment-3 case; no accuracy/F1 is defined +附件3_未对齐版本_02,0,negative,-0.6189733147621155,0.3779905438423157,0.36179712414741516,0.26021236181259155,-1.2931972742080688,1.0875630378723145,0.44498664140701294,0.44498664140701294,0.0,-0.09244058281183243,0.4485733211040497,1.4110205173492432,4.546689510345459,1.3469334840774536,4.6229448318481445,False,unlabeled attachment-3 case; no accuracy/F1 is defined +附件3_未对齐版本_03,2,positive,0.695109486579895,0.11934426426887512,0.2899485230445862,0.5907072424888611,-0.49461933970451355,1.5512335300445557,0.3907756507396698,0.39077427983283997,1.3719078424401232e-06,0.398899108171463,0.3993552029132843,1.0179466009140015,4.93976354598999,1.549193024635315,4.420685291290283,False,unlabeled attachment-3 case; no accuracy/F1 is defined +附件3_未对齐版本_04,2,positive,0.5424688458442688,0.18944765627384186,0.3276962339878082,0.48285606503486633,-0.7178608775138855,1.2966943979263306,0.33880671858787537,0.33880671858787537,0.0,0.21264974772930145,0.3422609567642212,1.0159424543380737,4.941767692565918,1.2757951021194458,4.694083213806152,False,unlabeled attachment-3 case; no accuracy/F1 is defined +附件3_未对齐版本_05,0,negative,-1.6132756471633911,0.7777170538902283,0.12479734420776367,0.09748563170433044,-2.460880994796753,0.9297983050346375,0.9528027772903442,0.9528027772903442,0.0,-1.1471410989761353,1.0553690195083618,3.1798951625823975,2.777815103530884,2.0038387775421143,3.9660396575927734,False,unlabeled attachment-3 case; no accuracy/F1 is defined +附件3_未对齐版本_06,2,positive,1.075721025466919,0.029009997844696045,0.1393413245677948,0.8316486477851868,0.0,2.0461437702178955,0.43230101466178894,0.43230101466178894,0.0,0.9193391799926758,0.4604739844799042,0.8610076904296875,5.096702575683594,2.229879856109619,3.7399985790252686,False,unlabeled attachment-3 case; no accuracy/F1 is defined +附件3_未对齐版本_07,2,positive,0.8112409710884094,0.06607470661401749,0.22776396572589874,0.7061613202095032,-0.1362186074256897,1.7385063171386719,0.3990931510925293,0.3990931510925293,0.0,0.5933741927146912,0.41373252868652344,0.9166243076324463,5.041085720062256,1.7580684423446655,4.2118096351623535,False,unlabeled attachment-3 case; no accuracy/F1 is defined +附件3_未对齐版本_08,0,negative,-0.8177421689033508,0.5444804430007935,0.24635988473892212,0.20915962755680084,-1.6569616794586182,0.9865659475326538,0.5980508923530579,0.5980508923530579,0.0,-0.34253618121147156,0.6036122441291809,1.7663525342941284,4.191357612609863,1.348613977432251,4.621264457702637,False,unlabeled attachment-3 case; no accuracy/F1 is defined +附件3_未对齐版本_09,2,positive,0.584557056427002,0.060164276510477066,0.22881989181041718,0.7110158205032349,-0.03433350846171379,1.4888426065444946,0.2824844419956207,0.2824844419956207,0.0,0.46303731203079224,0.2886645495891571,0.657180666923523,5.300529479980469,1.351650595664978,4.618227481842041,False,unlabeled attachment-3 case; no accuracy/F1 is defined +附件3_未对齐版本_10,1,neutral,0.0,0.19817937910556793,0.4307863712310791,0.371034175157547,-0.5964999794960022,1.0180858373641968,0.2146325260400772,0.21463251113891602,1.7206826186111357e-08,0.11066635698080063,0.21938827633857727,0.8156828284263611,5.142027378082275,1.030116081237793,4.939762115478516,False,unlabeled attachment-3 case; no accuracy/F1 is defined +附件3_未对齐版本_11,2,positive,0.4067777991294861,0.18321716785430908,0.36699020862579346,0.4497925639152527,-0.5481101274490356,1.093954086303711,0.23209431767463684,0.2320941984653473,1.2958071238244884e-07,0.1593105047941208,0.23480401933193207,0.7960365414619446,5.161673545837402,1.0298019647598267,4.9400763511657715,False,unlabeled attachment-3 case; no accuracy/F1 is defined +附件3_未对齐版本_12,2,positive,0.7418312430381775,0.16085933148860931,0.25514480471611023,0.5839958190917969,-0.7832371592521667,1.5990359783172607,0.48600462079048157,0.48600462079048157,0.0,0.3795293867588043,0.49588483572006226,1.2342653274536133,4.723444938659668,1.633910059928894,4.335968017578125,False,unlabeled attachment-3 case; no accuracy/F1 is defined +附件3_未对齐版本_13,2,positive,0.9053856730461121,0.22559532523155212,0.2667962610721588,0.5076084136962891,-1.2298040390014648,1.722484827041626,0.7403188347816467,0.7403188347816467,0.0,0.3018410801887512,0.7457730174064636,1.6682190895080566,4.289491176605225,1.9262142181396484,4.04366397857666,False,unlabeled attachment-3 case; no accuracy/F1 is defined +附件3_未对齐版本_14,0,negative,-0.7449327111244202,0.4758836627006531,0.29052573442459106,0.23359058797359467,-1.5283849239349365,1.0474128723144531,0.5521396398544312,0.5521395802497864,2.380643060462262e-08,-0.2323387712240219,0.5541964173316956,1.6363940238952637,4.321316242218018,1.361243486404419,4.6086344718933105,False,unlabeled attachment-3 case; no accuracy/F1 is defined +附件3_未对齐版本_15,2,positive,1.3370250463485718,0.023975681513547897,0.10470940917730331,0.8713149428367615,0.0,2.2798967361450195,0.47987255454063416,0.47987255454063416,0.0,1.167418360710144,0.5250018835067749,1.044642686843872,4.913067817687988,2.695021390914917,3.2748568058013916,False,unlabeled attachment-3 case; no accuracy/F1 is defined +附件3_未对齐版本_16,2,positive,1.7538543939590454,0.012006768025457859,0.06832415610551834,0.9196690917015076,0.0,2.5877037048339844,0.47664394974708557,0.47664394974708557,0.0,1.5816236734390259,0.5373051762580872,1.0920591354370117,4.8656511306762695,3.43656849861145,2.5333096981048584,False,unlabeled attachment-3 case; no accuracy/F1 is defined +附件3_未对齐版本_17,2,positive,0.6437320113182068,0.11531270295381546,0.28419703245162964,0.6004902720451355,-0.42875346541404724,1.4985958337783813,0.3550581634044647,0.35505813360214233,1.3726113579082266e-08,0.38549983501434326,0.3624156415462494,0.9383652806282043,5.019345283508301,1.4581032991409302,4.511775016784668,False,unlabeled attachment-3 case; no accuracy/F1 is defined +附件3_未对齐版本_18,2,positive,0.5576586127281189,0.10332871973514557,0.2823880910873413,0.6142831444740295,-0.2905426323413849,1.4048357009887695,0.29468634724617004,0.29468634724617004,0.0,0.3613300621509552,0.29980164766311646,0.7871065735816956,5.1706037521362305,1.303192138671875,4.666686058044434,False,unlabeled attachment-3 case; no accuracy/F1 is defined +附件3_未对齐版本_19,0,negative,-0.6961055397987366,0.5107713937759399,0.2870360016822815,0.20219261944293976,-1.5002697706222534,0.8932285904884338,0.47930681705474854,0.47930681705474854,0.0,-0.2732691168785095,0.4835101366043091,1.5491523742675781,4.408557891845703,1.2319159507751465,4.737962245941162,False,unlabeled attachment-3 case; no accuracy/F1 is defined +附件3_未对齐版本_20,2,positive,0.5754238963127136,0.19199968874454498,0.35681214928627014,0.45118817687034607,-0.7409619092941284,1.3107733726501465,0.3477743864059448,0.3477743864059448,0.0,0.202140673995018,0.3519274592399597,1.040464162826538,4.917246341705322,1.3352046012878418,4.634673595428467,False,unlabeled attachment-3 case; no accuracy/F1 is defined +附件3_未对齐版本_21,2,positive,0.7424408197402954,0.031176764518022537,0.17010530829429626,0.7987179756164551,0.0,1.706007719039917,0.32180947065353394,0.321806401014328,3.0546812013199087e-06,0.6466689705848694,0.33276495337486267,0.6083263158798218,5.34938383102417,1.6308528184890747,4.339025497436523,False,unlabeled attachment-3 case; no accuracy/F1 is defined +附件3_未对齐版本_22,2,positive,0.9128319621086121,0.3136769235134125,0.24424688518047333,0.4420761466026306,-1.5561679601669312,1.679451823234558,0.9327147603034973,0.9327132105827332,1.59684725531406e-06,0.1216338723897934,0.9247239232063293,1.9581161737442017,3.999594211578369,1.9394835233688354,4.030395030975342,False,unlabeled attachment-3 case; no accuracy/F1 is defined +附件3_未对齐版本_23,2,positive,1.1373980045318604,0.017837069928646088,0.10193044692277908,0.8802324533462524,0.0,2.1167213916778564,0.40735018253326416,0.4073500633239746,1.3645041008203407e-07,1.0282129049301147,0.4355928599834442,0.7500343322753906,5.207675933837891,2.3408405780792236,3.629037857055664,False,unlabeled attachment-3 case; no accuracy/F1 is defined +附件3_未对齐版本_24,2,positive,0.7078329920768738,0.12042908370494843,0.2849310338497162,0.5946398973464966,-0.49775230884552,1.5678426027297974,0.39866694808006287,0.39866626262664795,6.773137215532188e-07,0.40837085247039795,0.40745168924331665,1.016898274421692,4.940811634063721,1.5724689960479736,4.397408962249756,False,unlabeled attachment-3 case; no accuracy/F1 is defined +附件3_未对齐版本_25,2,positive,0.5815068483352661,0.12130436301231384,0.3121408522129059,0.566554844379425,-0.4213324785232544,1.4042198657989502,0.3172762989997864,0.3172754645347595,8.146940331243968e-07,0.32919952273368835,0.3226009011268616,0.8835678696632385,5.0741424560546875,1.3444410562515259,4.625437259674072,False,unlabeled attachment-3 case; no accuracy/F1 is defined +附件3_未对齐版本_26,2,positive,0.5501649975776672,0.056601475924253464,0.24990615248680115,0.6934923529624939,-0.012923063710331917,1.4388282299041748,0.2632901668548584,0.2632899284362793,2.120857800491649e-07,0.4325553774833679,0.26744163036346436,0.5925376415252686,5.365172863006592,1.2903460264205933,4.679532051086426,False,unlabeled attachment-3 case; no accuracy/F1 is defined +附件3_未对齐版本_27,2,positive,1.4617162942886353,0.016829069703817368,0.09628183394670486,0.8868891000747681,0.0,2.379847764968872,0.48261168599128723,0.48261168599128723,0.0,1.291656255722046,0.530409574508667,0.9833105206489563,4.974400043487549,2.916853666305542,3.0530247688293457,False,unlabeled attachment-3 case; no accuracy/F1 is defined +附件3_未对齐版本_28,2,positive,1.010158896446228,0.024440867826342583,0.12888681888580322,0.8466722965240479,0.0,1.9912506341934204,0.39680054783821106,0.3968004584312439,7.125536427565748e-08,0.8896118998527527,0.42043906450271606,0.76673823595047,5.190971851348877,2.113187074661255,3.8566908836364746,False,unlabeled attachment-3 case; no accuracy/F1 is defined +附件3_未对齐版本_29,2,positive,1.7028013467788696,0.013574469834566116,0.07568830251693726,0.9107372760772705,0.0,2.5532665252685547,0.48900747299194336,0.48900729417800903,1.7003151242533932e-07,1.5236588716506958,0.5494855046272278,1.1064831018447876,4.851227283477783,3.3458335399627686,2.62404465675354,False,unlabeled attachment-3 case; no accuracy/F1 is defined +附件3_未对齐版本_30,2,positive,0.9856576919555664,0.045877669006586075,0.16616477072238922,0.7879575490951538,0.0,1.9468833208084106,0.43973156809806824,0.4397314190864563,1.4706266426856018e-07,0.7972231507301331,0.4657270610332489,0.9605550169944763,4.99715518951416,2.0706770420074463,3.8992011547088623,False,unlabeled attachment-3 case; no accuracy/F1 is defined diff --git a/final/output/q3/README.md b/final/output/q3/README.md new file mode 100644 index 0000000..0b800ab --- /dev/null +++ b/final/output/q3/README.md @@ -0,0 +1,10 @@ +# Q3 题目输出 + +`ati_ho/` 保存 ATI–HO 对附件 4 的最终预测与解释交付文件。附件 4 无真实标签,因此这些文件不含准确率或误差指标。 + +- `ati_ho/attachment4_predictions.csv`:类别、情感强度和类别概率。 +- `ati_ho/attachment4_explanations.csv`:ATI 参数分解与分类/强度 Shapley 结果。 +- `ati_ho/attachment4_local_evidence.csv`:相对进度片段的局部 Owen 贡献和标准误。 +- `ati_ho/attachment4_prediction_manifest.json`:模型、输入、scaler 哈希与推理范围。 + +训练权重、验证指标和完整实验审计保存在 `final/experiments/q3/ati_ho/`,不放入题目输出目录。历史 MoFE 第一轮材料保存在 `final/experiments/q3/legacy_mofe_first_round/`。 diff --git a/final/output/q3/ati_ho/README.md b/final/output/q3/ati_ho/README.md new file mode 100644 index 0000000..fab689f --- /dev/null +++ b/final/output/q3/ati_ho/README.md @@ -0,0 +1,10 @@ +# Q3 ATI–HO 提交输出 + +| 文件 | 内容 | +|---|---| +| `attachment4_predictions.csv` | 官方附件4的 20 条预测类别、强度与类别概率 | +| `attachment4_explanations.csv` | 主效应、pairwise 项、解析/精确分类 Shapley 与强度精确 Shapley | +| `attachment4_local_evidence.csv` | 按模态分组的局部 Hierarchical Owen 片段贡献、标准误与相对进度位置 | +| `attachment4_prediction_manifest.json` | adapter、模型权重哈希、文件计数和无标签推理审计 | + +附件4没有标签,本目录不提供准确率或误差指标。所有位置均为归一化进度槽,不是秒数。 diff --git a/final/output/q3/ati_ho/attachment4_explanations.csv b/final/output/q3/ati_ho/attachment4_explanations.csv new file mode 100644 index 0000000..d129e93 --- /dev/null +++ b/final/output/q3/ati_ho/attachment4_explanations.csv @@ -0,0 +1,21 @@ +case_id,fixed_target_class,fixed_runner_up_class,full_logit_margin,baseline_r_negative,baseline_r_positive,analytic_vs_exact_shapley_max_abs,analytic_vs_exact_shapley_all_pass,exact_intensity_shapley_sum,intensity_full_minus_empty_coalition,intensity_shapley_efficiency_residual,coordinate_mode,physical_time_alignment,G_T_logit_negative,G_T_logit_neutral,G_T_logit_positive,G_T_r_negative,G_T_r_positive,analytic_class_shapley_T,exact_class_shapley_T,exact_intensity_shapley_T,G_A_logit_negative,G_A_logit_neutral,G_A_logit_positive,G_A_r_negative,G_A_r_positive,analytic_class_shapley_A,exact_class_shapley_A,exact_intensity_shapley_A,G_V_logit_negative,G_V_logit_neutral,G_V_logit_positive,G_V_r_negative,G_V_r_positive,analytic_class_shapley_V,exact_class_shapley_V,exact_intensity_shapley_V,G_TA_logit_negative,G_TA_logit_neutral,G_TA_logit_positive,G_TA_r_negative,G_TA_r_positive,G_TV_logit_negative,G_TV_logit_neutral,G_TV_logit_positive,G_TV_r_negative,G_TV_r_positive,G_AV_logit_negative,G_AV_logit_neutral,G_AV_logit_positive,G_AV_r_negative,G_AV_r_positive,baseline_logit_negative,full_parameter_logit_negative,baseline_logit_neutral,full_parameter_logit_neutral,baseline_logit_positive,full_parameter_logit_positive,full_parameter_r_negative,full_parameter_r_positive +01,neutral,positive,0.17453938722610474,-0.03336911275982857,-0.02806602418422699,1.5522042984272844e-08,True,-1.478951811790466,-1.4789518117904663,2.220446049250313e-16,relative_progress,False,-0.904476523399353,0.7555410265922546,0.1489354968070984,-0.3240058124065399,-0.23826012015342712,0.6066055297851562,0.6066055142631133,-1.2055534323056538,-0.07994242757558823,-0.13918815553188324,0.21913057565689087,-0.32910090684890747,-0.25575292110443115,-0.3583187460899353,-0.3583187318096558,-0.09216936429341632,-0.06264422088861465,-0.0008340037311427295,0.06347822397947311,-0.3798588216304779,-0.5035299062728882,-0.06431222707033157,-0.06431221279005209,-0.18122901519139606,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,-0.004946508444845676,-1.0520095825195312,-0.003916915971785784,0.6116019487380981,0.005518266931176186,0.4370625615119934,-1.0663347244262695,-1.025609016418457 +02,positive,neutral,0.6727461628615856,-0.03336911275982857,-0.02806602418422699,5.5258472686503524e-08,True,-0.40346074104309076,-0.4034607410430908,5.551115123125783e-17,relative_progress,False,-0.6345188021659851,-0.13300542533397675,0.7675241827964783,-0.2942814826965332,-0.10120987892150879,0.9005296230316162,0.9005296782900889,0.5349497596422831,-0.060333430767059326,-0.038159824907779694,0.09849324822425842,-0.2349826991558075,-0.2913568913936615,0.13665306568145752,0.13665306878586608,-0.1772233446439107,0.12454855442047119,0.12466159462928772,-0.2492101490497589,-0.07290247082710266,-0.16126064956188202,-0.37387174367904663,-0.3738717480252186,-0.7611871560414631,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,-0.004946508444845676,-0.5752502083778381,-0.003916915971785784,-0.05042056366801262,0.005518266931176186,0.622325599193573,-0.6355358362197876,-0.5818934440612793 +03,negative,neutral,0.2300729900598526,-0.03336911275982857,-0.02806602418422699,5.6965897499150486e-08,True,-2.6271680593490596,-2.62716805934906,4.440892098500626e-16,relative_progress,False,0.6437126398086548,0.233465313911438,-0.877177894115448,-0.0156564861536026,-0.19630727171897888,0.4102473258972168,0.4102472689313193,-2.6401662031809487,-0.044556912034749985,-0.10961116850376129,0.15416809916496277,-0.21152648329734802,-0.20100197196006775,0.06505425274372101,0.0650542665583392,0.0028151472409566197,-0.24782779812812805,-0.003628835082054138,0.251456618309021,-0.21738770604133606,-0.18694338202476501,-0.2441989630460739,-0.24419895295674598,0.010182996590932206,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,-0.004946508444845676,0.34638139605522156,-0.003916915971785784,0.11630840599536896,0.005518266931176186,-0.4660349488258362,-0.47793978452682495,-0.6123186349868774 +04,negative,positive,0.7093599438667297,-0.03336911275982857,-0.02806602418422699,1.0337680578231812e-07,True,-2.610314011573791,-2.6103140115737915,4.440892098500626e-16,relative_progress,False,1.1683094501495361,-0.7724652290344238,-0.3958442211151123,0.14337058365345,0.0068162246607244015,1.5641536712646484,1.5641535678878427,-2.699520587921142,-0.11619937419891357,-0.15902423858642578,0.27522361278533936,-0.3240225911140442,-0.2022765576839447,-0.39142298698425293,-0.39142295625060797,0.043889820575714104,-0.13935381174087524,-0.17419832944869995,0.3135521411895752,-0.2877661883831024,-0.16212013363838196,-0.45290595293045044,-0.4529058923944831,0.045316755771636956,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,-0.004946508444845676,0.9078097343444824,-0.003916915971785784,-1.109604835510254,0.005518266931176186,0.19844979047775269,-0.5017873048782349,-0.3856464624404907 +05,positive,neutral,1.5049121379852295,-0.03336911275982857,-0.02806602418422699,8.133550488675922e-08,True,-0.4806082248687743,-0.4806082248687744,1.1102230246251565e-16,relative_progress,False,-1.709882140159607,0.5530490875244141,1.1568331718444824,-0.35266321897506714,-0.29165583848953247,0.6037840843200684,0.6037841221938529,-0.2102790276209513,-0.05014742910861969,-0.11470775306224823,0.16485518217086792,-0.22012154757976532,-0.2014065831899643,0.27956295013427734,0.27956292840341723,-0.14488289753595987,-0.21467027068138123,-0.19872985780239105,0.4134001135826111,-0.24117706716060638,-0.1745043247938156,0.6121299862861633,0.6121299049506584,-0.1254462997118632,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,-0.004946508444845676,-1.9796464443206787,-0.003916915971785784,0.2356945276260376,0.005518266931176186,1.740606665611267,-0.8473309278488159,-0.695632815361023 +06,positive,neutral,1.9159240424633026,-0.03336911275982857,-0.02806602418422699,7.388492426207982e-08,True,-0.2154364585876466,-0.21543645858764648,-1.1102230246251565e-16,relative_progress,False,-1.2203913927078247,-0.4495927691459656,1.669984221458435,-0.37521523237228394,0.019710635766386986,2.119576930999756,2.1195768682907024,0.659255842367808,-0.12370359897613525,0.02123902551829815,0.10246457159519196,-0.13953426480293274,-0.16336387395858765,0.08122554421424866,0.0812254703293244,-0.1013184388478597,0.015853652730584145,0.1392299085855484,-0.155083566904068,-0.013930173590779305,-0.14624570310115814,-0.2943134903907776,-0.29431344879170257,-0.7733738621075948,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,-0.004946508444845676,-1.3331878185272217,-0.003916915971785784,-0.2930407226085663,0.005518266931176186,1.6228833198547363,-0.5620487928390503,-0.31796497106552124 +07,positive,neutral,1.3011534810066223,-0.03336911275982857,-0.02806602418422699,5.65002362673539e-08,True,-0.723442018032074,-0.723442018032074,0.0,relative_progress,False,-1.6689116954803467,0.3794674873352051,1.289444088935852,-0.42971712350845337,-0.2734326720237732,0.909976601600647,0.9099765451004107,-0.1858192980289459,-0.0658918172121048,-0.08562538027763367,0.15151719748973846,-0.22639398276805878,-0.3063961863517761,0.23714257776737213,0.2371425957729419,-0.20845401287078855,-0.12855705618858337,-0.008021022193133831,0.13657806813716888,-0.4119300842285156,-0.48094600439071655,0.1445990949869156,0.14459909809132415,-0.3291687071323395,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,-0.004946508444845676,-1.8683069944381714,-0.003916915971785784,0.2819041609764099,0.005518266931176186,1.5830576419830322,-1.1014103889465332,-1.0888409614562988 +08,positive,neutral,1.6082637459039688,-0.03336911275982857,-0.02806602418422699,1.0492901014735878e-07,True,-0.5725139379501343,-0.5725139379501343,0.0,relative_progress,False,-1.888288974761963,0.20506727695465088,1.6832218170166016,-0.439189612865448,-0.17151769995689392,1.4781545400619507,1.4781544351329405,0.06195046504338582,-0.041887879371643066,-0.052495844662189484,0.09438371658325195,-0.14891745150089264,-0.11223925650119781,0.14687955379486084,0.1468795457233985,0.1037632425626119,-0.04141402989625931,0.03380971401929855,0.007604313548654318,-0.35752567648887634,-0.525276780128479,-0.02620540000498295,-0.0262054322908322,-0.738227645556132,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,-0.004946508444845676,-1.9765374660491943,-0.003916915971785784,0.18246422708034515,0.005518266931176186,1.790727972984314,-0.9790017604827881,-0.8370997309684753 +09,negative,positive,4.036897420883179,-0.03336911275982857,-0.02806602418422699,3.5235037376679657e-07,True,-2.9891992807388306,-2.9891992807388306,0.0,relative_progress,False,2.800884246826172,-1.3056433200836182,-1.4952411651611328,0.5610975623130798,0.024273857474327087,4.296125411987305,4.2961257643376785,-2.801382025082906,-0.038079407066106796,-0.1371668577194214,0.17524628341197968,-0.26829323172569275,-0.14377743005752563,-0.21332569420337677,-0.213325595172743,0.2530961434046427,-0.04993594437837601,0.0644339919090271,-0.01449805311858654,-0.24577169120311737,-0.3168761730194092,-0.03543788939714432,-0.03543773448715607,-0.4409133990605672,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,-0.004946508444845676,2.7079226970672607,-0.003916915971785784,-1.3822932243347168,0.005518266931176186,-1.328974723815918,0.013663536868989468,-0.4644457697868347 +10,negative,positive,1.9884863495826721,-0.03336911275982857,-0.02806602418422699,9.53053431729245e-08,True,-2.8972601890563965,-2.8972601890563965,0.0,relative_progress,False,1.9594717025756836,-1.0742477178573608,-0.8852239847183228,0.3760708272457123,0.06034252047538757,2.844695568084717,2.84469566339006,-2.984575629234314,-0.1273432970046997,-0.04922483488917351,0.17656812071800232,-0.16357895731925964,-0.07018119096755981,-0.303911417722702,-0.3039113522196809,0.034755587577819824,-0.13882485032081604,-0.26418352127075195,0.4030084013938904,-0.2881527543067932,-0.1471811681985855,-0.5418332815170288,-0.5418332458163301,0.052559852600097656,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,-0.004946508444845676,1.6883571147918701,-0.003916915971785784,-1.391573190689087,0.005518266931176186,-0.300129234790802,-0.10903000831604004,-0.18508586287498474 +11,negative,positive,1.2935872972011566,-0.03336911275982857,-0.02806602418422699,8.195638656616211e-08,True,-2.848456025123596,-2.848456025123596,0.0,relative_progress,False,1.1382856369018555,-0.7049045562744141,-0.43338102102279663,0.16562658548355103,-0.005292731337249279,1.5716667175292969,1.5716666355729103,-2.8356826305389404,0.015153793618083,-0.04410065710544586,0.02894686535000801,-0.22009000182151794,-0.24944591522216797,-0.01379307173192501,-0.013793108053505419,-0.010202229022979736,-0.13108593225479126,0.00835040770471096,0.12273551523685455,-0.08660280704498291,-0.09405757486820221,-0.2538214325904846,-0.253821425139904,-0.0025711655616760254,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,-0.004946508444845676,1.017406940460205,-0.003916915971785784,-0.7445716857910156,0.005518266931176186,-0.27618035674095154,-0.1744353473186493,-0.3768622875213623 +12,negative,neutral,1.9524167478084564,-0.03336911275982857,-0.02806602418422699,1.8114224076271057e-07,True,-2.605541944503784,-2.605541944503784,0.0,relative_progress,False,1.6268107891082764,-0.4549051523208618,-1.171905755996704,0.12826871871948242,-0.14643552899360657,2.0817160606384277,2.081715879496187,-2.024578471978505,-0.01550658605992794,0.08049377799034119,-0.0649871975183487,-0.22043946385383606,-0.1405203640460968,-0.09600036591291428,-0.09600043902173638,-0.5963235100110372,-0.03571124002337456,-0.003442181274294853,0.03915341943502426,-0.3830249011516571,-0.5144790410995483,-0.032269060611724854,-0.03226907039061189,0.015360037485758468,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,-0.004946508444845676,1.5706462860107422,-0.003916915971785784,-0.38177046179771423,0.005518266931176186,-1.1922211647033691,-0.5085647702217102,-0.8295010328292847 +13,positive,neutral,0.03907114267349243,-0.03336911275982857,-0.02806602418422699,3.290673100675434e-08,True,-0.6098982095718384,-0.6098982095718384,0.0,relative_progress,False,-1.2801198959350586,0.8114702105522156,0.468649685382843,-0.2939906418323517,-0.38156017661094666,-0.34282052516937256,-0.3428205431749423,-1.0097482800483704,0.023065369576215744,-0.03675944358110428,0.013694070279598236,-0.1403023600578308,-0.20620891451835632,0.050453513860702515,0.05045352938274542,0.21370625495910645,-0.15847381949424744,-0.0817645788192749,0.24023841321468353,-0.1458529531955719,-0.2810816764831543,0.32200300693511963,0.3220029740283886,0.18614381551742554,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,-0.004946508444845676,-1.4204747676849365,-0.003916915971785784,0.689029335975647,0.005518266931176186,0.7281004786491394,-0.6135150790214539,-0.8969167470932007 +14,positive,neutral,0.17552608251571655,-0.03336911275982857,-0.02806602418422699,3.787378477504433e-08,True,-0.5350744724273682,-0.5350744724273682,0.0,relative_progress,False,-1.025192379951477,0.941837728023529,0.083354651927948,-0.2702452540397644,-0.3983325660228729,-0.858483076095581,-0.8584830382217963,-1.0336409012476602,-0.14746510982513428,-0.10266011953353882,0.2501252293586731,-0.3350585103034973,-0.25513797998428345,0.3527853488922119,0.3527853271613518,0.2199622591336568,-0.29058393836021423,-0.19060233235359192,0.48118627071380615,-0.23630091547966003,-0.09704461693763733,0.6717885732650757,0.6717886111388603,0.2786041696866353,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,-0.004946508444845676,-1.4681880474090576,-0.003916915971785784,0.6446583271026611,0.005518266931176186,0.8201844096183777,-0.8749737739562988,-0.7785811424255371 +15,positive,neutral,3.1091278791427612,-0.03336911275982857,-0.02806602418422699,1.6577541828155518e-07,True,-0.3441582918167114,-0.3441582918167114,0.0,relative_progress,False,-1.7186782360076904,-0.38246893882751465,2.101147174835205,-0.5896317958831787,-0.02962455153465271,2.4836161136627197,2.4836159478873014,-0.021669268608093258,-0.14169877767562866,-0.11777668446302414,0.2594754695892334,-0.3810732960700989,-0.16558730602264404,0.37725216150283813,0.3772521745413542,-0.12142491340637207,-0.25492432713508606,0.008049803785979748,0.24687449634075165,-0.34001070261001587,-0.2736431062221527,0.23882469534873962,0.23882469348609447,-0.20106410980224607,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,-0.004946508444845676,-2.1202478408813477,-0.003916915971785784,-0.4961127042770386,0.005518266931176186,2.6130151748657227,-1.3440849781036377,-0.49692100286483765 +16,negative,neutral,4.228886842727661,-0.03336911275982857,-0.02806602418422699,1.5444432695937982e-07,True,-3.0126639604568477,-3.012663960456848,4.440892098500626e-16,relative_progress,False,2.9560515880584717,-1.2988104820251465,-1.657240867614746,0.566375195980072,0.0027166109066456556,4.254861831665039,4.254861882111678,-2.9760817686716714,-0.02732587233185768,-0.04761290177702904,0.07493877410888672,-0.11062884330749512,-0.1420450508594513,0.020287029445171356,0.020286875000844397,-0.011243383089701336,-0.03721527382731438,0.008017238229513168,0.02919803373515606,-0.3774202764034271,-0.44900593161582947,-0.045232512056827545,-0.04523256033038099,-0.02533880869547525,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,-0.004946508444845676,2.886563777923584,-0.003916915971785784,-1.3423230648040771,0.005518266931176186,-1.5475858449935913,0.0449569895863533,-0.6164003610610962 +17,positive,neutral,3.5790738463401794,-0.03336911275982857,-0.02806602418422699,1.2231369783677337e-07,True,-0.19588220119476318,-0.19588220119476318,0.0,relative_progress,False,-1.7191108465194702,-0.765514075756073,2.4846248626708984,-0.5405866503715515,0.05220368504524231,3.250138998031616,3.250139120345314,0.26819344361623126,-0.06238323450088501,0.11898847669363022,-0.05660523474216461,-0.19090893864631653,-0.1367032825946808,-0.17559370398521423,-0.1755937902877728,-0.5605612794558207,-0.27668532729148865,-0.10920405387878418,0.3858893811702728,-0.134759321808815,-0.17871713638305664,0.495093435049057,0.49509339344998193,0.09648563464482625,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,-0.004946508444845676,-2.0631260871887207,-0.003916915971785784,-0.7596465945243835,0.005518266931176186,2.819427251815796,-0.8996240496635437,-0.2912827730178833 +18,neutral,negative,0.24999400973320007,-0.03336911275982857,-0.02806602418422699,1.2262413889851942e-08,True,-1.4789518117904663,-1.4789518117904663,0.0,relative_progress,False,0.28653019666671753,0.403386652469635,-0.6899168491363525,-0.019597142934799194,-0.21449951827526093,0.11685645580291748,0.11685644571358958,-0.7159723242123921,-0.0920357033610344,-0.11092132329940796,0.20295703411102295,-0.2633468806743622,-0.1888737976551056,-0.018885619938373566,-0.018885626302411158,-0.04700716336568196,0.007763783447444439,0.15875737369060516,-0.16652116179466248,-0.2628101110458374,-0.28852805495262146,0.15099358558654785,0.15099359784896174,-0.7159723242123921,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,-0.004946508444845676,0.19731178879737854,-0.003916915971785784,0.4473057985305786,0.005518266931176186,-0.6479626893997192,-0.5791232585906982,-0.7199673652648926 +19,positive,negative,0.12766507267951965,-0.03336911275982857,-0.02806602418422699,2.7939677238464355e-09,True,-0.08535563945770264,-0.08535563945770264,0.0,relative_progress,False,0.7118504047393799,-0.72557532787323,0.013724908232688904,-0.0308525450527668,-0.001771043986082077,-0.6981254816055298,-0.6981254825368524,-1.2033302386601765,-0.041151344776153564,0.06523413956165314,-0.02408279851078987,-0.1305142492055893,-0.05948138236999512,0.017068546265363693,0.017068549059331417,-0.4989778598149618,-0.2795937955379486,-0.23906968533992767,0.5186634659767151,-0.14898139238357544,-0.05279204994440079,0.7982572317123413,0.7982572307810187,1.6169524590174356,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,-0.004946508444845676,0.3861587941646576,-0.003916915971785784,-0.9033277034759521,0.005518266931176186,0.5138238668441772,-0.3437173068523407,-0.14211049675941467 +20,positive,neutral,1.7363310009241104,-0.03336911275982857,-0.02806602418422699,4.4393042741841526e-08,True,-0.40033042430877674,-0.40033042430877686,1.1102230246251565e-16,relative_progress,False,-1.9999498128890991,0.14454472064971924,1.8554052114486694,-0.46275147795677185,-0.2674782872200012,1.7108604907989502,1.7108604659636812,1.316787560780843,0.05088898167014122,-0.08756403625011444,0.03667505830526352,-0.10067225992679596,-0.0886862576007843,0.12423909455537796,0.12423905016233522,-1.1897882024447122,-0.020659327507019043,0.06443151831626892,-0.04377218708395958,-0.28119325637817383,-0.19312873482704163,-0.1082037091255188,-0.10820371254036823,-0.5273297826449076,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,0.0,-0.004946508444845676,-1.9746668338775635,-0.003916915971785784,0.11749528348445892,0.005518266931176186,1.8538262844085693,-0.8779861330986023,-0.5773593187332153 diff --git a/final/output/q3/ati_ho/attachment4_local_evidence.csv b/final/output/q3/ati_ho/attachment4_local_evidence.csv new file mode 100644 index 0000000..f764140 --- /dev/null +++ b/final/output/q3/ati_ho/attachment4_local_evidence.csv @@ -0,0 +1,601 @@ +case_id,modality,relative_bin,relative_position_start,relative_position_end,local_owen_margin_contribution,owen_standard_error,permutations,stopping_status,physical_time_alignment +01,text,0,0.0,0.1,0.15501678062719293,0.02824737907762575,64,max_permutations,False +01,text,1,0.1,0.2,0.3393877267371863,0.03891857387116652,64,max_permutations,False +01,text,2,0.2,0.3,0.13160342909395695,0.021407469220121024,64,max_permutations,False +01,text,3,0.3,0.4,0.1541730676253792,0.023352077620926736,64,max_permutations,False +01,text,4,0.4,0.5,0.1187975586799439,0.02708052178375021,64,max_permutations,False +01,text,5,0.5,0.6,-0.32819975397433154,0.021070444987702094,64,max_permutations,False +01,text,6,0.6,0.7,-0.18518312580999918,0.01611749256140915,64,max_permutations,False +01,text,7,0.7,0.8,-0.09827927622245625,0.015110065326582888,64,max_permutations,False +01,text,8,0.8,0.9,0.041710176767082885,0.01712487503973442,64,max_permutations,False +01,text,9,0.9,1.0,0.27757893287343904,0.04676788039390862,64,max_permutations,False +01,audio,0,0.0,0.1,-0.011243712127907202,0.009233671019475894,64,max_permutations,False +01,audio,1,0.1,0.2,-0.032589772396022454,0.008617888532936106,64,max_permutations,False +01,audio,2,0.2,0.3,-0.06452517947764136,0.012695545986241358,64,max_permutations,False +01,audio,3,0.3,0.4,-0.058146502036834136,0.012725619090246807,64,max_permutations,False +01,audio,4,0.4,0.5,-0.04902930595562793,0.01048247600797557,64,max_permutations,False +01,audio,5,0.5,0.6,-0.039003983518341556,0.009180882634066611,64,max_permutations,False +01,audio,6,0.6,0.7,-0.009039346128702164,0.005976381690355603,64,max_permutations,False +01,audio,7,0.7,0.8,-0.05669806583318859,0.012973192997537856,64,max_permutations,False +01,audio,8,0.8,0.9,-0.05080851863021962,0.011617200421893223,64,max_permutations,False +01,audio,9,0.9,1.0,0.012765651888912544,0.004659085312381929,64,max_permutations,False +01,vision,0,0.0,0.1,-0.006502474949229509,0.0015052195559562167,64,max_permutations,False +01,vision,1,0.1,0.2,-0.006554739462444559,0.0011094309519753375,64,max_permutations,False +01,vision,2,0.2,0.3,-0.011689859733451158,0.0013052881660502993,64,max_permutations,False +01,vision,3,0.3,0.4,-0.026898140262346715,0.0023536208796073924,64,max_permutations,False +01,vision,4,0.4,0.5,-0.008135400246828794,0.0009906486108777516,64,max_permutations,False +01,vision,5,0.5,0.6,0.005125108757056296,0.0010578993554463802,64,max_permutations,False +01,vision,6,0.6,0.7,-0.004117644391953945,0.0013342637103614845,64,max_permutations,False +01,vision,7,0.7,0.8,0.014664946414995939,0.0015623694916137292,64,max_permutations,False +01,vision,8,0.8,0.9,-0.006133872899226844,0.0012644093533939029,64,max_permutations,False +01,vision,9,0.9,1.0,-0.014070135744987056,0.001974886787411523,64,max_permutations,False +02,text,0,0.0,0.1,-0.2435244449879974,0.03784503512777638,16,stable,False +02,text,1,0.1,0.2,0.08069909503683448,0.02251316758096149,16,stable,False +02,text,2,0.2,0.3,0.13866681058425456,0.06606381171215224,16,stable,False +02,text,3,0.3,0.4,0.3282391676912084,0.109628547483824,16,stable,False +02,text,4,0.4,0.5,0.4153550040209666,0.12304339544019145,16,stable,False +02,text,5,0.5,0.6,0.3692910491954535,0.11042240368777048,16,stable,False +02,text,6,0.6,0.7,0.0652436160016805,0.05736249030197223,16,stable,False +02,text,7,0.7,0.8,0.025608156574890018,0.03900034766287585,16,stable,False +02,text,8,0.8,0.9,-0.19390656682662666,0.06900286960562475,16,stable,False +02,text,9,0.9,1.0,-0.08514219988137484,0.034105645498513665,16,stable,False +02,audio,0,0.0,0.1,-0.025586919859051704,0.006379168684057744,16,stable,False +02,audio,1,0.1,0.2,0.016373574268072844,0.008831531476611865,16,stable,False +02,audio,2,0.2,0.3,0.05392544995993376,0.017164153137310234,16,stable,False +02,audio,3,0.3,0.4,0.01686276914551854,0.009908323474383376,16,stable,False +02,audio,4,0.4,0.5,0.014328758406918496,0.011732993410349142,16,stable,False +02,audio,5,0.5,0.6,0.04691262322012335,0.018712270387235944,16,stable,False +02,audio,6,0.6,0.7,0.003989448130596429,0.008562717075729453,16,stable,False +02,audio,7,0.7,0.8,0.016533869318664074,0.01057255051161289,16,stable,False +02,audio,8,0.8,0.9,0.01014680229127407,0.008664929199207605,16,stable,False +02,audio,9,0.9,1.0,-0.016833303146995604,0.0029446053018819273,16,stable,False +02,vision,0,0.0,0.1,-0.0151264633750543,0.01664132360117765,16,stable,False +02,vision,1,0.1,0.2,-0.06713183771353215,0.027702583825160908,16,stable,False +02,vision,2,0.2,0.3,-0.0959150152048096,0.028765221253865606,16,stable,False +02,vision,3,0.3,0.4,-0.05922563676722348,0.025965331517317503,16,stable,False +02,vision,4,0.4,0.5,-0.04631277953740209,0.023417126665055382,16,stable,False +02,vision,5,0.5,0.6,-0.015756300650537014,0.004390764367608424,16,stable,False +02,vision,6,0.6,0.7,-0.03570611774921417,0.01888455822860072,16,stable,False +02,vision,7,0.7,0.8,-0.002132096327841282,0.0026840496248077004,16,stable,False +02,vision,8,0.8,0.9,-0.019219718873500824,0.01745792859172882,16,stable,False +02,vision,9,0.9,1.0,-0.01734579389449209,0.017423208839136055,16,stable,False +03,text,0,0.0,0.1,-0.1468745181336999,0.0319421234683677,16,stable,False +03,text,1,0.1,0.2,-0.20071300026029348,0.028826606134011213,16,stable,False +03,text,2,0.2,0.3,-0.10554380202665925,0.052538191507580274,16,stable,False +03,text,3,0.3,0.4,-0.11763328965753317,0.022682032979912947,16,stable,False +03,text,4,0.4,0.5,-0.25409176759421825,0.04328636854479889,16,stable,False +03,text,5,0.5,0.6,0.025146916043013334,0.018089850784241857,16,stable,False +03,text,6,0.6,0.7,0.37964809220284224,0.0796977127780905,16,stable,False +03,text,7,0.7,0.8,0.4260731328104157,0.08012881849553437,16,stable,False +03,text,8,0.8,0.9,0.3716314242046792,0.07532427432576366,16,stable,False +03,text,9,0.9,1.0,0.03260408714413643,0.00958417025002297,16,stable,False +03,audio,0,0.0,0.1,0.007766876835376024,0.003271961759439887,16,stable,False +03,audio,1,0.1,0.2,0.006488126469776034,0.003324157638883152,16,stable,False +03,audio,2,0.2,0.3,-0.016845647449372336,0.005399341791008501,16,stable,False +03,audio,3,0.3,0.4,0.02497231960296631,0.0029256090190421784,16,stable,False +03,audio,4,0.4,0.5,0.012833932181820273,0.0021935914267862686,16,stable,False +03,audio,5,0.5,0.6,0.004549576668068767,0.0034734011702652755,16,stable,False +03,audio,6,0.6,0.7,-0.006960852537304163,0.002996739408371313,16,stable,False +03,audio,7,0.7,0.8,0.01306554430630058,0.0035671004404954194,16,stable,False +03,audio,8,0.8,0.9,-0.005483502696733922,0.002068398928076061,16,stable,False +03,audio,9,0.9,1.0,0.024667887715622783,0.007094296879298716,16,stable,False +03,vision,0,0.0,0.1,-0.006535648368299007,0.004425495990915771,16,stable,False +03,vision,1,0.1,0.2,-0.031847413833020255,0.0160869496446553,16,stable,False +03,vision,2,0.2,0.3,-0.07149078717338853,0.02230828022286466,16,stable,False +03,vision,3,0.3,0.4,-0.07968386419815943,0.024541203945209027,16,stable,False +03,vision,4,0.4,0.5,-0.016761592589318752,0.014055884186525796,16,stable,False +03,vision,5,0.5,0.6,-0.03342015130328946,0.013764353612724005,16,stable,False +03,vision,6,0.6,0.7,-0.06931015150621533,0.023818303790187446,16,stable,False +03,vision,7,0.7,0.8,-0.014363787340698764,0.014199624226887066,16,stable,False +03,vision,8,0.8,0.9,0.05309187390957959,0.007874248770874274,16,stable,False +03,vision,9,0.9,1.0,0.026122569106519222,0.0071736234178903235,16,stable,False +04,text,0,0.0,0.1,0.9141831923043355,0.24524596536621782,16,stable,False +04,text,1,0.1,0.2,0.5140680961194448,0.1650650755456328,16,stable,False +04,text,2,0.2,0.3,-0.5401804260909557,0.09564544199803597,16,stable,False +04,text,3,0.3,0.4,0.112938666716218,0.09031893851514408,16,stable,False +04,text,4,0.4,0.5,0.6277451105415821,0.18271232691496506,16,stable,False +04,text,5,0.5,0.6,-0.514852661639452,0.11074220090820447,16,stable,False +04,text,6,0.6,0.7,-0.5971456747502089,0.09624872271555468,16,stable,False +04,text,7,0.7,0.8,0.16591994650661945,0.11185014873586584,16,stable,False +04,text,8,0.8,0.9,0.2307781662675552,0.1340995377302692,16,stable,False +04,text,9,0.9,1.0,0.6506991538917646,0.17125446269256822,16,stable,False +04,audio,0,0.0,0.1,0.02359095774590969,0.016732583123800063,16,stable,False +04,audio,1,0.1,0.2,-0.07095338852377608,0.02680729641708691,16,stable,False +04,audio,2,0.2,0.3,-0.032523399975616485,0.026842560963308136,16,stable,False +04,audio,3,0.3,0.4,0.020568767562508583,0.004943455944633252,16,stable,False +04,audio,4,0.4,0.5,-0.062041693192441016,0.02984964496093275,16,stable,False +04,audio,5,0.5,0.6,-0.038894188415724784,0.021505186204123474,16,stable,False +04,audio,6,0.6,0.7,-0.0678289420902729,0.03067686039614309,16,stable,False +04,audio,7,0.7,0.8,-0.041840177786070853,0.023952643573328304,16,stable,False +04,audio,8,0.8,0.9,-0.09816189447883517,0.03659130925672187,16,stable,False +04,audio,9,0.9,1.0,-0.023339003324508667,0.00383657670443348,16,stable,False +04,vision,0,0.0,0.1,-0.025206708174664527,0.026378259950959587,16,stable,False +04,vision,1,0.1,0.2,-0.013624444603919983,0.019716376773975103,16,stable,False +04,vision,2,0.2,0.3,-0.024314943701028824,0.005766432950559478,16,stable,False +04,vision,3,0.3,0.4,-0.05704651493579149,0.026471305143213303,16,stable,False +04,vision,4,0.4,0.5,-0.04453487694263458,0.03174308285765449,16,stable,False +04,vision,5,0.5,0.6,0.08410330553306267,0.020474813190971948,16,stable,False +04,vision,6,0.6,0.7,0.05325228441506624,0.01108765528050272,16,stable,False +04,vision,7,0.7,0.8,-0.1704186499118805,0.04873071861783258,16,stable,False +04,vision,8,0.8,0.9,-0.12445163010852411,0.04287187047836418,16,stable,False +04,vision,9,0.9,1.0,-0.13066370971500874,0.040374675786863254,16,stable,False +05,text,0,0.0,0.1,0.23293456807732582,0.044398432047631275,32,stable,False +05,text,1,0.1,0.2,0.6020670213038102,0.08588462093473578,32,stable,False +05,text,2,0.2,0.3,0.01162069255951792,0.0264620228429983,32,stable,False +05,text,3,0.3,0.4,0.041673462837934494,0.022701154674003157,32,stable,False +05,text,4,0.4,0.5,-0.24565158819314092,0.02699554567613218,32,stable,False +05,text,5,0.5,0.6,-0.03989897854626179,0.031097782791919607,32,stable,False +05,text,6,0.6,0.7,-0.2312568612396717,0.03439740072569168,32,stable,False +05,text,7,0.7,0.8,0.23074009083211422,0.050355089571634086,32,stable,False +05,text,8,0.8,0.9,-0.06878424267051741,0.02565297276253453,32,stable,False +05,text,9,0.9,1.0,0.0703399513149634,0.0362178359824638,32,stable,False +05,audio,0,0.0,0.1,-0.09578148537548259,0.0098309928222143,32,stable,False +05,audio,1,0.1,0.2,-0.00981973874149844,0.007592593964265901,32,stable,False +05,audio,2,0.2,0.3,0.05399813351687044,0.014812493209196262,32,stable,False +05,audio,3,0.3,0.4,0.08642724622040987,0.019177949831390068,32,stable,False +05,audio,4,0.4,0.5,0.026131076039746404,0.009260606751003778,32,stable,False +05,audio,5,0.5,0.6,0.07434694247785956,0.018903685651552662,32,stable,False +05,audio,6,0.6,0.7,0.0071559567004442215,0.0062492049444848475,32,stable,False +05,audio,7,0.7,0.8,0.06161860638530925,0.012505995089940969,32,stable,False +05,audio,8,0.8,0.9,0.07158913864986971,0.01793481855764846,32,stable,False +05,audio,9,0.9,1.0,0.0038970530149526894,0.009892918874126777,32,stable,False +05,vision,0,0.0,0.1,-0.014343914575874805,0.015287948382323027,32,stable,False +05,vision,1,0.1,0.2,0.1029385207220912,0.02239155389065744,32,stable,False +05,vision,2,0.2,0.3,0.13193559815408662,0.03449586787907593,32,stable,False +05,vision,3,0.3,0.4,0.08305018657119945,0.028398797419757114,32,stable,False +05,vision,4,0.4,0.5,0.1675075776875019,0.035316019190087186,32,stable,False +05,vision,5,0.5,0.6,0.06359046476427466,0.0246648711208595,32,stable,False +05,vision,6,0.6,0.7,0.020825696003157645,0.0223136347433165,32,stable,False +05,vision,7,0.7,0.8,-0.05489160324214026,0.013125629329992788,32,stable,False +05,vision,8,0.8,0.9,0.059608321462292224,0.029116810394589545,32,stable,False +05,vision,9,0.9,1.0,0.05190906283678487,0.026304850943417998,32,stable,False +06,text,0,0.0,0.1,0.18846334269619547,0.06730884201190947,64,stable,False +06,text,1,0.1,0.2,0.3957692921103444,0.07577594795535471,64,stable,False +06,text,2,0.2,0.3,0.21731206859112717,0.06105360304872178,64,stable,False +06,text,3,0.3,0.4,0.4803523705340922,0.08313842672325379,64,stable,False +06,text,4,0.4,0.5,0.2823158281680662,0.06717924423897143,64,stable,False +06,text,5,0.5,0.6,0.24262981908395886,0.035196081320967264,64,stable,False +06,text,6,0.6,0.7,0.5434295730374288,0.09424949768326749,64,stable,False +06,text,7,0.7,0.8,-0.0884781887580175,0.045925553845011666,64,stable,False +06,text,8,0.8,0.9,-0.07236569913220592,0.042965322659466294,64,stable,False +06,text,9,0.9,1.0,-0.06985152451670729,0.03293487728170707,64,stable,False +06,audio,0,0.0,0.1,-0.05214787973091006,0.0036863432908589276,64,stable,False +06,audio,1,0.1,0.2,-0.007022940117167309,0.0029470715911910565,64,stable,False +06,audio,2,0.2,0.3,0.030832908902084455,0.0060431408682659745,64,stable,False +06,audio,3,0.3,0.4,-0.0023749730025883764,0.0035272925520899194,64,stable,False +06,audio,4,0.4,0.5,0.016031424805987626,0.004621540220354766,64,stable,False +06,audio,5,0.5,0.6,0.08691148844081908,0.009988336362468354,64,stable,False +06,audio,6,0.6,0.7,-0.0011681779578793794,0.0029408348640113913,64,stable,False +06,audio,7,0.7,0.8,0.04691894384450279,0.007146872018158698,64,stable,False +06,audio,8,0.8,0.9,-0.06488956298562698,0.006090834440525607,64,stable,False +06,audio,9,0.9,1.0,0.028134224092354998,0.005504207740849057,64,stable,False +06,vision,0,0.0,0.1,0.00015291321324184537,0.006396887328841082,64,stable,False +06,vision,1,0.1,0.2,-0.00969747846829705,0.00890091467134087,64,stable,False +06,vision,2,0.2,0.3,-0.027639445528620854,0.008574950208810841,64,stable,False +06,vision,3,0.3,0.4,-0.031400189152918756,0.010176875602015844,64,stable,False +06,vision,4,0.4,0.5,-0.023442256380803883,0.00794074657376766,64,stable,False +06,vision,5,0.5,0.6,-0.007102909963577986,0.006265396068488513,64,stable,False +06,vision,6,0.6,0.7,-0.0197655062074773,0.007618730929945213,64,stable,False +06,vision,7,0.7,0.8,-0.03513202635804191,0.00942967169768761,64,stable,False +06,vision,8,0.8,0.9,-0.11203585282783024,0.016250395019693643,64,stable,False +06,vision,9,0.9,1.0,-0.02825069660320878,0.009925887578503517,64,stable,False +07,text,0,0.0,0.1,-0.2807314048986882,0.07199795779186322,16,stable,False +07,text,1,0.1,0.2,-0.09571871161460876,0.03664843045032331,16,stable,False +07,text,2,0.2,0.3,-0.13468044064939022,0.020678854689247187,16,stable,False +07,text,3,0.3,0.4,0.22829848306719214,0.07451636374234474,16,stable,False +07,text,4,0.4,0.5,0.6596076075220481,0.17048749390447948,16,stable,False +07,text,5,0.5,0.6,0.3983154036104679,0.08918721333228016,16,stable,False +07,text,6,0.6,0.7,0.10073965415358543,0.04440199355530042,16,stable,False +07,text,7,0.7,0.8,-0.12250468949787319,0.047021136776418,16,stable,False +07,text,8,0.8,0.9,0.03219079377595335,0.058013483781366754,16,stable,False +07,text,9,0.9,1.0,0.12445983663201332,0.03178747045053555,16,stable,False +07,audio,0,0.0,0.1,0.026075719855725765,0.013401243452956714,16,stable,False +07,audio,1,0.1,0.2,0.027768587926402688,0.015812423563727133,16,stable,False +07,audio,2,0.2,0.3,0.030620926059782505,0.004426927026903294,16,stable,False +07,audio,3,0.3,0.4,0.03348503413144499,0.013051232010705604,16,stable,False +07,audio,4,0.4,0.5,0.07670619210693985,0.023307670673857236,16,stable,False +07,audio,5,0.5,0.6,0.022687664488330483,0.014040747590711115,16,stable,False +07,audio,6,0.6,0.7,-0.00988765712827444,0.004206218124756885,16,stable,False +07,audio,7,0.7,0.8,0.0007308972999453545,0.006877662469259831,16,stable,False +07,audio,8,0.8,0.9,0.03251865284983069,0.01768278860743525,16,stable,False +07,audio,9,0.9,1.0,-0.003563420264981687,0.008202134202086153,16,stable,False +07,vision,0,0.0,0.1,-0.0024851257912814617,0.000789153633491801,16,stable,False +07,vision,1,0.1,0.2,0.0007033385336399078,0.007718101656332319,16,stable,False +07,vision,2,0.2,0.3,0.04117334447801113,0.010553154090053005,16,stable,False +07,vision,3,0.3,0.4,0.027776760049164295,0.007851572454693596,16,stable,False +07,vision,4,0.4,0.5,0.020356823690235615,0.008475799972083118,16,stable,False +07,vision,5,0.5,0.6,0.02346079656854272,0.008896903677679282,16,stable,False +07,vision,6,0.6,0.7,-0.005256508942693472,0.0029639681909905766,16,stable,False +07,vision,7,0.7,0.8,0.025500426650978625,0.009935565536515522,16,stable,False +07,vision,8,0.8,0.9,-0.01366146607324481,0.005163228255806352,16,stable,False +07,vision,9,0.9,1.0,0.027030720375478268,0.010367075626397076,16,stable,False +08,text,0,0.0,0.1,0.17770834435941651,0.05086179014937832,32,stable,False +08,text,1,0.1,0.2,0.2499188704532571,0.04833664988674012,32,stable,False +08,text,2,0.2,0.3,0.026259816775564104,0.07045408004852691,32,stable,False +08,text,3,0.3,0.4,0.029634585371240973,0.051879886859682006,32,stable,False +08,text,4,0.4,0.5,-0.0343631996656768,0.04545107010431922,32,stable,False +08,text,5,0.5,0.6,-0.16044770478038117,0.023143879327351045,32,stable,False +08,text,6,0.6,0.7,0.32741222821641713,0.08759215548279477,32,stable,False +08,text,7,0.7,0.8,0.3970816290238872,0.09063187172760538,32,stable,False +08,text,8,0.8,0.9,0.4221285214298405,0.09525691508464883,32,stable,False +08,text,9,0.9,1.0,0.04282134445384145,0.03360955767376876,32,stable,False +08,audio,0,0.0,0.1,0.05316140741342679,0.010873419555925998,32,stable,False +08,audio,1,0.1,0.2,0.0871599295642227,0.01438313462474187,32,stable,False +08,audio,2,0.2,0.3,0.03795783658279106,0.009266816219941269,32,stable,False +08,audio,3,0.3,0.4,-0.014795155206229538,0.0037579314828385443,32,stable,False +08,audio,4,0.4,0.5,-0.05113702290691435,0.007596154718520131,32,stable,False +08,audio,5,0.5,0.6,0.020376423664856702,0.006369232285698349,32,stable,False +08,audio,6,0.6,0.7,0.0728136120014824,0.014421761209324037,32,stable,False +08,audio,7,0.7,0.8,0.02373598760459572,0.008822048706576926,32,stable,False +08,audio,8,0.8,0.9,-0.06333806028123945,0.0073744831773825455,32,stable,False +08,audio,9,0.9,1.0,-0.01905539585277438,0.0038848974439001736,32,stable,False +08,vision,0,0.0,0.1,-0.009772224439075217,0.002849140495600707,32,stable,False +08,vision,1,0.1,0.2,-0.008020571316592395,0.0025751063934248753,32,stable,False +08,vision,2,0.2,0.3,0.0039413339691236615,0.001381266016106183,32,stable,False +08,vision,3,0.3,0.4,-0.0007593688205815852,0.0022579792590337626,32,stable,False +08,vision,4,0.4,0.5,-0.01496548531576991,0.0035957861125950168,32,stable,False +08,vision,5,0.5,0.6,0.008766540500801057,0.0015528403463582147,32,stable,False +08,vision,6,0.6,0.7,0.005108536744955927,0.0014318259607011613,32,stable,False +08,vision,7,0.7,0.8,5.5977609008550644e-05,0.0019346942286685511,32,stable,False +08,vision,8,0.8,0.9,-0.007509635470341891,0.0029836460446540864,32,stable,False +08,vision,9,0.9,1.0,-0.0030505531176459044,0.00142592014577895,32,stable,False +09,text,0,0.0,0.1,0.10768596082925797,0.07193388358513367,64,max_permutations,False +09,text,1,0.1,0.2,0.15403015854826663,0.08336394126897574,64,max_permutations,False +09,text,2,0.2,0.3,0.18160415906459093,0.08997383499850248,64,max_permutations,False +09,text,3,0.3,0.4,0.9826138447388075,0.16587439742202253,64,max_permutations,False +09,text,4,0.4,0.5,0.6115690304577583,0.13047027452942908,64,max_permutations,False +09,text,5,0.5,0.6,0.49902549386024475,0.12256492458760214,64,max_permutations,False +09,text,6,0.6,0.7,0.41781722227460705,0.12378921411803138,64,max_permutations,False +09,text,7,0.7,0.8,0.399600631557405,0.10196689811575514,64,max_permutations,False +09,text,8,0.8,0.9,0.417607262119418,0.10928116768965355,64,max_permutations,False +09,text,9,0.9,1.0,0.524572035545134,0.14065626478549254,64,max_permutations,False +09,audio,0,0.0,0.1,0.023669981528655626,0.0051520133308288595,64,max_permutations,False +09,audio,1,0.1,0.2,-0.00939754476712551,0.005473113660761007,64,max_permutations,False +09,audio,2,0.2,0.3,0.040614094366901554,0.0041078750316379965,64,max_permutations,False +09,audio,3,0.3,0.4,-0.016439422048279084,0.0072136628696657665,64,max_permutations,False +09,audio,4,0.4,0.5,-0.05226275355380494,0.0064904620363677385,64,max_permutations,False +09,audio,5,0.5,0.6,-0.10833994395215996,0.013296084753314452,64,max_permutations,False +09,audio,6,0.6,0.7,-0.03989495833229739,0.007320245875938111,64,max_permutations,False +09,audio,7,0.7,0.8,-0.013566266730776988,0.005035313936052559,64,max_permutations,False +09,audio,8,0.8,0.9,-0.038148356121382676,0.00886424122318373,64,max_permutations,False +09,audio,9,0.9,1.0,0.00043955372530035675,0.006197565930567541,64,max_permutations,False +09,vision,0,0.0,0.1,-0.08023086354660336,0.008839803929316286,64,max_permutations,False +09,vision,1,0.1,0.2,-0.07771580701228231,0.0070556649751515365,64,max_permutations,False +09,vision,2,0.2,0.3,-0.08755727780226152,0.009575022261140967,64,max_permutations,False +09,vision,3,0.3,0.4,0.0292091522278497,0.003217859298158789,64,max_permutations,False +09,vision,4,0.4,0.5,0.03450135386083275,0.005715388020666905,64,max_permutations,False +09,vision,5,0.5,0.6,0.01119528801064007,0.0025980599041537816,64,max_permutations,False +09,vision,6,0.6,0.7,0.028356000926578417,0.0036435989460922116,64,max_permutations,False +09,vision,7,0.7,0.8,0.015950414250255562,0.002734138120982871,64,max_permutations,False +09,vision,8,0.8,0.9,0.017621606966713443,0.0036707584059330954,64,max_permutations,False +09,vision,9,0.9,1.0,0.07323238368553575,0.006916269456994557,64,max_permutations,False +10,text,0,0.0,0.1,0.4140143576951232,0.1003773425374924,64,stable,False +10,text,1,0.1,0.2,0.01121065801999066,0.05246440018063112,64,stable,False +10,text,2,0.2,0.3,-0.2394247827178333,0.04743588983292265,64,stable,False +10,text,3,0.3,0.4,0.17991569339937996,0.06999216901021174,64,stable,False +10,text,4,0.4,0.5,0.1494577877165284,0.06218177702920901,64,stable,False +10,text,5,0.5,0.6,0.19738677851273678,0.0656933121703012,64,stable,False +10,text,6,0.6,0.7,-0.033821260469267145,0.042924653056262786,64,stable,False +10,text,7,0.7,0.8,0.255947959041805,0.0672016446508023,64,stable,False +10,text,8,0.8,0.9,0.9977771317498991,0.1185689206389685,64,stable,False +10,text,9,0.9,1.0,0.9122313311236212,0.12249904316796838,64,stable,False +10,audio,0,0.0,0.1,-0.02243388789065648,0.00799590230531354,64,stable,False +10,audio,1,0.1,0.2,-0.02562665325240232,0.009263428262468369,64,stable,False +10,audio,2,0.2,0.3,-0.06718353304313496,0.012705761263230067,64,stable,False +10,audio,3,0.3,0.4,-0.044055100399418734,0.010404476609097359,64,stable,False +10,audio,4,0.4,0.5,-0.055598239006940275,0.010983050320342466,64,stable,False +10,audio,5,0.5,0.6,-0.026588419001200236,0.009006937078574168,64,stable,False +10,audio,6,0.6,0.7,-0.048232322980766185,0.0092315476739546,64,stable,False +10,audio,7,0.7,0.8,-0.032605959902866744,0.009424441098727502,64,stable,False +10,audio,8,0.8,0.9,0.0025217785441782326,0.00500671756094839,64,stable,False +10,audio,9,0.9,1.0,0.015890972776105627,0.006464776977388028,64,stable,False +10,vision,0,0.0,0.1,-0.045404995515127666,0.015533622691539765,64,stable,False +10,vision,1,0.1,0.2,-0.12280285042652395,0.02296312934995374,64,stable,False +10,vision,2,0.2,0.3,-0.10574701335281134,0.01811528331885136,64,stable,False +10,vision,3,0.3,0.4,-0.07697499348432757,0.015453804387098997,64,stable,False +10,vision,4,0.4,0.5,-0.020605886122211814,0.010291220091879721,64,stable,False +10,vision,5,0.5,0.6,0.024123090362991206,0.010118095100470485,64,stable,False +10,vision,6,0.6,0.7,-0.036647429878939874,0.01316007433152269,64,stable,False +10,vision,7,0.7,0.8,-0.002815269588609226,0.011069457121929277,64,stable,False +10,vision,8,0.8,0.9,-0.07251107609772589,0.019379081696173132,64,stable,False +10,vision,9,0.9,1.0,-0.08244680045754649,0.01965900097928553,64,stable,False +11,text,0,0.0,0.1,0.08317508539767005,0.04954351005429095,32,stable,False +11,text,1,0.1,0.2,1.0367089239880443,0.12217696010638715,32,stable,False +11,text,2,0.2,0.3,1.2976987939327955,0.11042302703400987,32,stable,False +11,text,3,0.3,0.4,0.5735141683544498,0.09506295243377091,32,stable,False +11,text,4,0.4,0.5,0.15150513211847283,0.0625122025037525,32,stable,False +11,text,5,0.5,0.6,-0.38421271872357465,0.050680664789031195,32,stable,False +11,text,6,0.6,0.7,-0.3501955643296242,0.04217346507123575,32,stable,False +11,text,7,0.7,0.8,-0.47991615513456054,0.052959887051836574,32,stable,False +11,text,8,0.8,0.9,-0.5079016497475095,0.058471655584300994,32,stable,False +11,text,9,0.9,1.0,0.15129060903564095,0.05744181582104048,32,stable,False +11,audio,0,0.0,0.1,0.0018712930323090404,0.0016871066184651411,32,stable,False +11,audio,1,0.1,0.2,-0.011189746350282803,0.0028909078455731205,32,stable,False +11,audio,2,0.2,0.3,0.00937529353541322,0.001603648872778038,32,stable,False +11,audio,3,0.3,0.4,0.0382854119525291,0.004117772505279889,32,stable,False +11,audio,4,0.4,0.5,-0.0017727971717249602,0.0013478594362271007,32,stable,False +11,audio,5,0.5,0.6,-0.014202558100805618,0.0027746389203791243,32,stable,False +11,audio,6,0.6,0.7,0.014589822400012054,0.0017130217499706257,32,stable,False +11,audio,7,0.7,0.8,-0.015179806243395433,0.0032442644218416243,32,stable,False +11,audio,8,0.8,0.9,-0.012747056491207331,0.0019680061212582643,32,stable,False +11,audio,9,0.9,1.0,-0.022822956962045282,0.002532858065194674,32,stable,False +11,vision,0,0.0,0.1,-0.01954834582284093,0.012090989810776621,32,stable,False +11,vision,1,0.1,0.2,0.05016999944928102,0.004326543193533188,32,stable,False +11,vision,2,0.2,0.3,0.04453502534306608,0.007363136352267293,32,stable,False +11,vision,3,0.3,0.4,0.01773451207554899,0.011142940924327957,32,stable,False +11,vision,4,0.4,0.5,0.017307151400018483,0.00985831921864745,32,stable,False +11,vision,5,0.5,0.6,-0.02267471485538408,0.010323882270265166,32,stable,False +11,vision,6,0.6,0.7,-0.10345343902008608,0.01803664928996661,32,stable,False +11,vision,7,0.7,0.8,-0.14296824307530187,0.02796628354308197,32,stable,False +11,vision,8,0.8,0.9,-0.026814361510332674,0.011924244492282523,32,stable,False +11,vision,9,0.9,1.0,-0.0681090060970746,0.0207660319901182,32,stable,False +12,text,0,0.0,0.1,-0.22706957020272966,0.028270720163018168,64,max_permutations,False +12,text,1,0.1,0.2,-0.05126071081031114,0.045658680137213685,64,max_permutations,False +12,text,2,0.2,0.3,-0.0007502666339860298,0.028853364095135444,64,max_permutations,False +12,text,3,0.3,0.4,0.21081679909548257,0.05735916403069648,64,max_permutations,False +12,text,4,0.4,0.5,0.12954307746258564,0.05510274057041151,64,max_permutations,False +12,text,5,0.5,0.6,0.4112846369534964,0.07828389201504268,64,max_permutations,False +12,text,6,0.6,0.7,0.21491040698310826,0.0586021692647383,64,max_permutations,False +12,text,7,0.7,0.8,0.5635053327641799,0.08192870595866984,64,max_permutations,False +12,text,8,0.8,0.9,0.27470678662939463,0.054513203557270434,64,max_permutations,False +12,text,9,0.9,1.0,0.5560293968446786,0.08999529164226666,64,max_permutations,False +12,audio,0,0.0,0.1,-0.006170861408463679,0.004798780931243528,64,max_permutations,False +12,audio,1,0.1,0.2,-0.015362115103926044,0.00258995441478962,64,max_permutations,False +12,audio,2,0.2,0.3,-0.0036864284047624096,0.0025797287784515488,64,max_permutations,False +12,audio,3,0.3,0.4,-0.015235978120472282,0.003646172476064035,64,max_permutations,False +12,audio,4,0.4,0.5,-0.02139794803224504,0.0035675354515246264,64,max_permutations,False +12,audio,5,0.5,0.6,-0.02214156695845304,0.0035430130764409293,64,max_permutations,False +12,audio,6,0.6,0.7,-0.0007115534390322864,0.002588959114002684,64,max_permutations,False +12,audio,7,0.7,0.8,0.0035326022407389246,0.0027588227666979635,64,max_permutations,False +12,audio,8,0.8,0.9,-0.011062974917877,0.00326390416772475,64,max_permutations,False +12,audio,9,0.9,1.0,-0.0037636279957951047,0.00324061542754634,64,max_permutations,False +12,vision,0,0.0,0.1,-0.006337987069855444,0.002301637977363389,64,max_permutations,False +12,vision,1,0.1,0.2,-0.003579808588256128,0.002670007175739108,64,max_permutations,False +12,vision,2,0.2,0.3,-0.0029591818747576326,0.0024554396975754325,64,max_permutations,False +12,vision,3,0.3,0.4,-0.004654540549381636,0.0025672597662441243,64,max_permutations,False +12,vision,4,0.4,0.5,-0.007126071686798241,0.0032008051787007353,64,max_permutations,False +12,vision,5,0.5,0.6,0.0013174179548514076,0.0021634175654936597,64,max_permutations,False +12,vision,6,0.6,0.7,-0.0006302003675955348,0.0021550750661950445,64,max_permutations,False +12,vision,7,0.7,0.8,0.0035120519023621455,0.0019089263253645567,64,max_permutations,False +12,vision,8,0.8,0.9,-0.0024549047666369006,0.0023895079936194896,64,max_permutations,False +12,vision,9,0.9,1.0,-0.00935584181570448,0.0025952001050152935,64,max_permutations,False +13,text,0,0.0,0.1,-0.17346966860350221,0.04551824983440317,16,stable,False +13,text,1,0.1,0.2,-0.1855074695777148,0.048209407087956735,16,stable,False +13,text,2,0.2,0.3,-0.01550071535166353,0.02172503173230116,16,stable,False +13,text,3,0.3,0.4,0.03290311177261174,0.023499213226645957,16,stable,False +13,text,4,0.4,0.5,0.1900026180082932,0.03912008599103367,16,stable,False +13,text,5,0.5,0.6,-0.13953795097768307,0.03286516961115438,16,stable,False +13,text,6,0.6,0.7,-0.12420026247855276,0.034847916907905535,16,stable,False +13,text,7,0.7,0.8,0.06725766509771347,0.013366956065681956,16,stable,False +13,text,8,0.8,0.9,0.019064208026975393,0.009466481989791187,16,stable,False +13,text,9,0.9,1.0,-0.013832075521349907,0.016580771920967817,16,stable,False +13,audio,0,0.0,0.1,-0.003313009685371071,0.0026748742871660547,16,stable,False +13,audio,1,0.1,0.2,0.009257545345462859,0.00395976823501722,16,stable,False +13,audio,2,0.2,0.3,0.019524238829035312,0.009409215219406412,16,stable,False +13,audio,3,0.3,0.4,0.04462461557704955,0.009484220323037713,16,stable,False +13,audio,4,0.4,0.5,0.014484378276392817,0.006951426960012512,16,stable,False +13,audio,5,0.5,0.6,0.012509375344961882,0.0042440774833653536,16,stable,False +13,audio,6,0.6,0.7,-0.005791692950879224,0.0033683296649581777,16,stable,False +13,audio,7,0.7,0.8,-0.06496133211476263,0.010975644658470165,16,stable,False +13,audio,8,0.8,0.9,0.024504587054252625,0.008614278400996093,16,stable,False +13,audio,9,0.9,1.0,-0.0003851763322018087,0.002261146242311167,16,stable,False +13,vision,0,0.0,0.1,0.027977536898106337,0.01829329870391982,16,stable,False +13,vision,1,0.1,0.2,0.027950285119004548,0.01830403216449929,16,stable,False +13,vision,2,0.2,0.3,0.04380461911205202,0.01990456380465491,16,stable,False +13,vision,3,0.3,0.4,0.04016937082633376,0.019893152391324876,16,stable,False +13,vision,4,0.4,0.5,0.02692084072623402,0.014545380772931507,16,stable,False +13,vision,5,0.5,0.6,0.018954191356897354,0.0034149259512102565,16,stable,False +13,vision,6,0.6,0.7,0.06265713414177299,0.022916999451708442,16,stable,False +13,vision,7,0.7,0.8,0.035298386588692665,0.01428720167143491,16,stable,False +13,vision,8,0.8,0.9,0.016107458621263504,0.0028639631718960505,16,stable,False +13,vision,9,0.9,1.0,0.0221631471067667,0.015163890241149372,16,stable,False +14,text,0,0.0,0.1,-0.24021617890684865,0.033546336635416424,64,stable,False +14,text,1,0.1,0.2,-0.16829417634289712,0.031083983220011557,64,stable,False +14,text,2,0.2,0.3,-0.026778001891216263,0.016457377836681397,64,stable,False +14,text,3,0.3,0.4,-0.23496718815295026,0.03471429644640515,64,stable,False +14,text,4,0.4,0.5,-0.07602742733433843,0.018601321194474392,64,stable,False +14,text,5,0.5,0.6,-0.1848787011404056,0.027805338574264888,64,stable,False +14,text,6,0.6,0.7,-0.11567743131308816,0.026433178143646617,64,stable,False +14,text,7,0.7,0.8,0.010257720685331151,0.01339085884044731,64,stable,False +14,text,8,0.8,0.9,-0.02071622767834924,0.016884469675709564,64,stable,False +14,text,9,0.9,1.0,0.19881457716110162,0.016145402060817323,64,stable,False +14,audio,0,0.0,0.1,0.04010411523631774,0.01146093680602742,64,stable,False +14,audio,1,0.1,0.2,0.04478503065183759,0.00739354475285859,64,stable,False +14,audio,2,0.2,0.3,0.0689044322934933,0.014591065681536624,64,stable,False +14,audio,3,0.3,0.4,0.01389613194623962,0.008440384415967835,64,stable,False +14,audio,4,0.4,0.5,0.050327645847573876,0.008836046028886341,64,stable,False +14,audio,5,0.5,0.6,0.06914294773014262,0.012219951146377578,64,stable,False +14,audio,6,0.6,0.7,-0.0026600561977829784,0.006211330361771348,64,stable,False +14,audio,7,0.7,0.8,0.031035248859552667,0.01052433811992653,64,stable,False +14,audio,8,0.8,0.9,0.04291828433633782,0.010814406550777271,64,stable,False +14,audio,9,0.9,1.0,-0.005668461497407407,0.00619452354852307,64,stable,False +14,vision,0,0.0,0.1,0.15063073497731239,0.030452212115541423,64,stable,False +14,vision,1,0.1,0.2,0.12773362622829154,0.025200581382778576,64,stable,False +14,vision,2,0.2,0.3,0.12687427486525849,0.025112801421750597,64,stable,False +14,vision,3,0.3,0.4,-0.005829234694829211,0.014181377501399051,64,stable,False +14,vision,4,0.4,0.5,0.00405814714031294,0.013759351953527931,64,stable,False +14,vision,5,0.5,0.6,-0.0036393722402863204,0.01355151981097642,64,stable,False +14,vision,6,0.6,0.7,0.08243993236101232,0.022398971495241983,64,stable,False +14,vision,7,0.7,0.8,0.10188267307239585,0.019902515686870136,64,stable,False +14,vision,8,0.8,0.9,-0.01266075111925602,0.010491981541790036,64,stable,False +14,vision,9,0.9,1.0,0.10029858519556001,0.021908087719144022,64,stable,False +15,text,0,0.0,0.1,-0.10142144392011687,0.042694541186588914,64,stable,False +15,text,1,0.1,0.2,0.09707848849939182,0.055215124041656585,64,stable,False +15,text,2,0.2,0.3,0.13559447845909745,0.06324687955850496,64,stable,False +15,text,3,0.3,0.4,0.11068305122898892,0.053299640584360736,64,stable,False +15,text,4,0.4,0.5,0.2955384205270093,0.07968247454766018,64,stable,False +15,text,5,0.5,0.6,0.2989993432711344,0.07534712658869172,64,stable,False +15,text,6,0.6,0.7,0.35361291031586006,0.07374311830841751,64,stable,False +15,text,7,0.7,0.8,0.6307718228781596,0.10564715215974335,64,stable,False +15,text,8,0.8,0.9,0.5166861609613989,0.07958889354707821,64,stable,False +15,text,9,0.9,1.0,0.14607270518899895,0.05596183481799439,64,stable,False +15,audio,0,0.0,0.1,0.023986150801647455,0.010435384871041542,64,stable,False +15,audio,1,0.1,0.2,0.013128890102962032,0.00828707965116855,64,stable,False +15,audio,2,0.2,0.3,0.019199409725842997,0.009546299116270906,64,stable,False +15,audio,3,0.3,0.4,0.02484546389314346,0.011000575266127782,64,stable,False +15,audio,4,0.4,0.5,0.08063976746052504,0.013485841107433012,64,stable,False +15,audio,5,0.5,0.6,0.031274444540031254,0.010320627288276802,64,stable,False +15,audio,6,0.6,0.7,0.012011096550850198,0.009406050233467344,64,stable,False +15,audio,7,0.7,0.8,0.051887454581446946,0.012573701216354476,64,stable,False +15,audio,8,0.8,0.9,0.04616080122650601,0.0102020937327251,64,stable,False +15,audio,9,0.9,1.0,0.07411868256167509,0.015540107592829568,64,stable,False +15,vision,0,0.0,0.1,-0.04629640790517442,0.0057401169280415386,64,stable,False +15,vision,1,0.1,0.2,0.018162530177505687,0.008625161146161646,64,stable,False +15,vision,2,0.2,0.3,0.0050220011617057025,0.006319620294920681,64,stable,False +15,vision,3,0.3,0.4,0.035088738484773785,0.008608333873168842,64,stable,False +15,vision,4,0.4,0.5,-0.07519077090546489,0.005561055077117622,64,stable,False +15,vision,5,0.5,0.6,0.1403894612158183,0.014683791668457322,64,stable,False +15,vision,6,0.6,0.7,0.03396641701692715,0.008402957882278083,64,stable,False +15,vision,7,0.7,0.8,0.11469262198079377,0.015664318080706844,64,stable,False +15,vision,8,0.8,0.9,0.004600239131832495,0.004925049542946882,64,stable,False +15,vision,9,0.9,1.0,0.008389886701479554,0.006321938824445505,64,stable,False +16,text,0,0.0,0.1,0.27613481408479856,0.11225502909395489,64,max_permutations,False +16,text,1,0.1,0.2,0.2249024673437816,0.10404059449465208,64,max_permutations,False +16,text,2,0.2,0.3,0.07591621148458216,0.06490398135463397,64,max_permutations,False +16,text,3,0.3,0.4,0.921995702527056,0.17974664239520405,64,max_permutations,False +16,text,4,0.4,0.5,0.6722130109192221,0.14760563796106158,64,max_permutations,False +16,text,5,0.5,0.6,0.5154972045202157,0.13389670208260918,64,max_permutations,False +16,text,6,0.6,0.7,0.47328798567468766,0.12307288579620322,64,max_permutations,False +16,text,7,0.7,0.8,0.46695831519173225,0.09927178278303314,64,max_permutations,False +16,text,8,0.8,0.9,0.6816840192259406,0.15785233957273828,64,max_permutations,False +16,text,9,0.9,1.0,-0.05372784006613074,0.02808113341356427,64,max_permutations,False +16,audio,0,0.0,0.1,-0.009475557359110098,0.001969012500903294,64,max_permutations,False +16,audio,1,0.1,0.2,-0.009745011469931342,0.0023627873202600558,64,max_permutations,False +16,audio,2,0.2,0.3,-0.022983237795415334,0.0026346801931409628,64,max_permutations,False +16,audio,3,0.3,0.4,-0.010865375144931022,0.0019948697206365242,64,max_permutations,False +16,audio,4,0.4,0.5,0.02456347769475542,0.0033065870154926286,64,max_permutations,False +16,audio,5,0.5,0.6,0.022985375915595796,0.0035950666435201274,64,max_permutations,False +16,audio,6,0.6,0.7,-0.04458897296717623,0.004860509598899562,64,max_permutations,False +16,audio,7,0.7,0.8,0.00020007837156299502,0.002176781281365927,64,max_permutations,False +16,audio,8,0.8,0.9,0.04933613885077648,0.004938507907049691,64,max_permutations,False +16,audio,9,0.9,1.0,0.02085994657682022,0.001980010909479503,64,max_permutations,False +16,vision,0,0.0,0.1,-0.004623666805855464,0.0022184895653467526,64,max_permutations,False +16,vision,1,0.1,0.2,-0.010318904263840523,0.0033496799205144485,64,max_permutations,False +16,vision,2,0.2,0.3,-0.008324586597154848,0.0036025348004011153,64,max_permutations,False +16,vision,3,0.3,0.4,-0.023601458029588684,0.004546559749863349,64,max_permutations,False +16,vision,4,0.4,0.5,0.00097469131287653,0.002398848577724479,64,max_permutations,False +16,vision,5,0.5,0.6,0.025678786776552442,0.0017144960949332388,64,max_permutations,False +16,vision,6,0.6,0.7,-0.0034266270376974717,0.00312468035284321,64,max_permutations,False +16,vision,7,0.7,0.8,-0.001737202168442309,0.0031106273859997343,64,max_permutations,False +16,vision,8,0.8,0.9,0.0021646023451467045,0.002949454278468333,64,max_permutations,False +16,vision,9,0.9,1.0,-0.02201819232868729,0.004685276293953575,64,max_permutations,False +17,text,0,0.0,0.1,0.3213029764010571,0.10350884941151615,64,max_permutations,False +17,text,1,0.1,0.2,0.5778402565338183,0.12594419923546785,64,max_permutations,False +17,text,2,0.2,0.3,0.2155533851182554,0.06659076185461864,64,max_permutations,False +17,text,3,0.3,0.4,0.5425265368248802,0.11329573263481668,64,max_permutations,False +17,text,4,0.4,0.5,0.2964746331854258,0.07527183376118962,64,max_permutations,False +17,text,5,0.5,0.6,0.2951276770909317,0.08701247977824281,64,max_permutations,False +17,text,6,0.6,0.7,0.2056336021341849,0.07828681394570655,64,max_permutations,False +17,text,7,0.7,0.8,0.35881203078315593,0.09112697012310823,64,max_permutations,False +17,text,8,0.8,0.9,0.3491082674881909,0.0936541357904053,64,max_permutations,False +17,text,9,0.9,1.0,0.0877597646904178,0.0664043926943639,64,max_permutations,False +17,audio,0,0.0,0.1,0.024635682202642784,0.003284924917776687,64,max_permutations,False +17,audio,1,0.1,0.2,-0.01588039187481627,0.004688207551997496,64,max_permutations,False +17,audio,2,0.2,0.3,-0.018846274149836972,0.005297152212336929,64,max_permutations,False +17,audio,3,0.3,0.4,-0.004117848118767142,0.0038689927282576707,64,max_permutations,False +17,audio,4,0.4,0.5,-0.010488038533367217,0.004344037277658568,64,max_permutations,False +17,audio,5,0.5,0.6,-0.028853188967332244,0.006528010879442298,64,max_permutations,False +17,audio,6,0.6,0.7,-0.012923963600769639,0.005138186081148229,64,max_permutations,False +17,audio,7,0.7,0.8,-0.04658543673576787,0.007233190366284945,64,max_permutations,False +17,audio,8,0.8,0.9,-0.05781004441087134,0.008355095087043485,64,max_permutations,False +17,audio,9,0.9,1.0,-0.0047242903092410415,0.0050342070421285695,64,max_permutations,False +17,vision,0,0.0,0.1,0.09849374537589028,0.02258639372309501,64,max_permutations,False +17,vision,1,0.1,0.2,0.025029940763488412,0.010643142758195797,64,max_permutations,False +17,vision,2,0.2,0.3,0.0031036475265864283,0.008721470149376975,64,max_permutations,False +17,vision,3,0.3,0.4,0.02950779438833706,0.01318911818329894,64,max_permutations,False +17,vision,4,0.4,0.5,0.12606115915696137,0.0222139038534662,64,max_permutations,False +17,vision,5,0.5,0.6,0.0826382775849197,0.016717598732732784,64,max_permutations,False +17,vision,6,0.6,0.7,0.10496403771685436,0.018235298625612517,64,max_permutations,False +17,vision,7,0.7,0.8,-0.0014881068200338632,0.011092231436624694,64,max_permutations,False +17,vision,8,0.8,0.9,0.056959778448799625,0.01771320116241097,64,max_permutations,False +17,vision,9,0.9,1.0,-0.03017688638647087,0.00978345210859298,64,max_permutations,False +18,text,0,0.0,0.1,0.44462686391489115,0.06349549411546251,32,stable,False +18,text,1,0.1,0.2,0.4641269795683911,0.06469902574993466,32,stable,False +18,text,2,0.2,0.3,0.26492878315912094,0.0360747787061787,32,stable,False +18,text,3,0.3,0.4,0.14193018546211533,0.030561183888971465,32,stable,False +18,text,4,0.4,0.5,0.00016883714124560356,0.034286925257080124,32,stable,False +18,text,5,0.5,0.6,-0.5257265120017109,0.06523310800276245,32,stable,False +18,text,6,0.6,0.7,-0.23096102502313443,0.03588915259490077,32,stable,False +18,text,7,0.7,0.8,-0.44532709631312173,0.04771371609079108,32,stable,False +18,text,8,0.8,0.9,-0.215724701891304,0.043152340892017714,32,stable,False +18,text,9,0.9,1.0,0.21881412724906113,0.0412351797268357,32,stable,False +18,audio,0,0.0,0.1,0.0026914111513178796,0.0012307270695554396,32,stable,False +18,audio,1,0.1,0.2,-0.02207712992094457,0.0035637936420267295,32,stable,False +18,audio,2,0.2,0.3,0.0004370304523035884,0.0013554621330095845,32,stable,False +18,audio,3,0.3,0.4,-0.029923082562163472,0.0032689031037291664,32,stable,False +18,audio,4,0.4,0.5,-0.010571632999926805,0.001739130517737338,32,stable,False +18,audio,5,0.5,0.6,0.009696273831650615,0.001215274885415808,32,stable,False +18,audio,6,0.6,0.7,-0.005392109640524723,0.001348878868917852,32,stable,False +18,audio,7,0.7,0.8,0.0027699833444785327,0.0007018617979253065,32,stable,False +18,audio,8,0.8,0.9,0.02641373744700104,0.004531255529369288,32,stable,False +18,audio,9,0.9,1.0,0.007069897837936878,0.0030215639316384145,32,stable,False +18,vision,0,0.0,0.1,0.009719536974444054,0.006060012388177373,32,stable,False +18,vision,1,0.1,0.2,0.023099806101527065,0.008916198455485566,32,stable,False +18,vision,2,0.2,0.3,0.01991078184801154,0.00799192143507949,32,stable,False +18,vision,3,0.3,0.4,0.044334021033137105,0.011979889252874971,32,stable,False +18,vision,4,0.4,0.5,-0.002103865146636963,0.0007614726994460468,32,stable,False +18,vision,5,0.5,0.6,-0.0002940924750873819,0.005549861273718779,32,stable,False +18,vision,6,0.6,0.7,-0.01154317194595933,0.0025653596162566817,32,stable,False +18,vision,7,0.7,0.8,0.020218590609147213,0.007385840909165094,32,stable,False +18,vision,8,0.8,0.9,0.009409133344888687,0.004353700031287672,32,stable,False +18,vision,9,0.9,1.0,0.03824285670998506,0.011000746181690882,32,stable,False +19,text,0,0.0,0.1,0.39794124715263024,0.04562997959934864,64,max_permutations,False +19,text,1,0.1,0.2,-0.40687543744570576,0.06507680687466239,64,max_permutations,False +19,text,2,0.2,0.3,-0.5426625471009174,0.05613958719441816,64,max_permutations,False +19,text,3,0.3,0.4,-0.4369588294503046,0.04800534471757422,64,max_permutations,False +19,text,4,0.4,0.5,-0.46953652174852323,0.05119107154545671,64,max_permutations,False +19,text,5,0.5,0.6,-0.015359458047896624,0.02170880657292385,64,max_permutations,False +19,text,6,0.6,0.7,-0.09237738396041095,0.02674835922665995,64,max_permutations,False +19,text,7,0.7,0.8,-0.39915040116466116,0.05917197424804057,64,max_permutations,False +19,text,8,0.8,0.9,0.9807364776206668,0.0906903936170207,64,max_permutations,False +19,text,9,0.9,1.0,0.2861173703568056,0.040364097308103225,64,max_permutations,False +19,audio,0,0.0,0.1,0.007687557925237343,0.002573989704062053,64,max_permutations,False +19,audio,1,0.1,0.2,0.004003813926829025,0.0038638864800786435,64,max_permutations,False +19,audio,2,0.2,0.3,-0.06272895631263964,0.005990930294959438,64,max_permutations,False +19,audio,3,0.3,0.4,-0.045318261269130744,0.005418861510306392,64,max_permutations,False +19,audio,4,0.4,0.5,0.05776199467072729,0.005838694624668898,64,max_permutations,False +19,audio,5,0.5,0.6,0.031531810280284844,0.004485617856069247,64,max_permutations,False +19,audio,6,0.6,0.7,-0.036698280993732624,0.003963007146807431,64,max_permutations,False +19,audio,7,0.7,0.8,-0.04412255165516399,0.0054655529697979904,64,max_permutations,False +19,audio,8,0.8,0.9,0.0729582949570613,0.00742283266450262,64,max_permutations,False +19,audio,9,0.9,1.0,0.03199312731157988,0.004110076744642543,64,max_permutations,False +19,vision,0,0.0,0.1,0.02713141730055213,0.013875462638664378,64,max_permutations,False +19,vision,1,0.1,0.2,0.11611285155231599,0.026976130554260222,64,max_permutations,False +19,vision,2,0.2,0.3,0.07810062216594815,0.02202074069598431,64,max_permutations,False +19,vision,3,0.3,0.4,0.038468335478683,0.017430974099180496,64,max_permutations,False +19,vision,4,0.4,0.5,0.057399069890379906,0.018638118736303586,64,max_permutations,False +19,vision,5,0.5,0.6,0.14544087984540965,0.028862239876774477,64,max_permutations,False +19,vision,6,0.6,0.7,0.12720882409485057,0.027767235259820368,64,max_permutations,False +19,vision,7,0.7,0.8,0.03505652994499542,0.015656104459435268,64,max_permutations,False +19,vision,8,0.8,0.9,0.09244590152229648,0.026334513618838293,64,max_permutations,False +19,vision,9,0.9,1.0,0.08089280045533087,0.025923112974604284,64,max_permutations,False +20,text,0,0.0,0.1,0.47541837859898806,0.08084614976373775,64,stable,False +20,text,1,0.1,0.2,0.2830378959479276,0.05851389646581134,64,stable,False +20,text,2,0.2,0.3,0.06894207181176171,0.03646051152387709,64,stable,False +20,text,3,0.3,0.4,0.14732616345281713,0.04195753949413667,64,stable,False +20,text,4,0.4,0.5,0.2211901356058661,0.04497266623043044,64,stable,False +20,text,5,0.5,0.6,0.17311706964392215,0.051432049073633516,64,stable,False +20,text,6,0.6,0.7,-0.1659611079376191,0.02116889347908944,64,stable,False +20,text,7,0.7,0.8,-0.12126856748363934,0.027266546104199377,64,stable,False +20,text,8,0.8,0.9,0.356910382892238,0.06372137677059009,64,stable,False +20,text,9,0.9,1.0,0.2721480450127274,0.055633627242429165,64,stable,False +20,audio,0,0.0,0.1,0.06879722249868792,0.009241933730653996,64,stable,False +20,audio,1,0.1,0.2,0.0734821704973001,0.009381487300125198,64,stable,False +20,audio,2,0.2,0.3,-0.010988269874360412,0.003113222443417005,64,stable,False +20,audio,3,0.3,0.4,0.04577479236468207,0.007238996449918193,64,stable,False +20,audio,4,0.4,0.5,0.0004804949276149273,0.003520295612097436,64,stable,False +20,audio,5,0.5,0.6,0.004908705101115629,0.0050695838457634436,64,stable,False +20,audio,6,0.6,0.7,-0.029079530053422786,0.0038778563254469197,64,stable,False +20,audio,7,0.7,0.8,0.014496768184471875,0.005370674964341348,64,stable,False +20,audio,8,0.8,0.9,-0.05241296952590346,0.004016791912509551,64,stable,False +20,audio,9,0.9,1.0,0.008779666779446416,0.004979878265916824,64,stable,False +20,vision,0,0.0,0.1,-0.06146337787504308,0.01204398644334877,64,stable,False +20,vision,1,0.1,0.2,0.19105098325235303,0.020379792999463592,64,stable,False +20,vision,2,0.2,0.3,-0.16483094592695124,0.01830270558188342,64,stable,False +20,vision,3,0.3,0.4,-0.12310345962760039,0.012765805609259074,64,stable,False +20,vision,4,0.4,0.5,-0.10652847339224536,0.011805191251376312,64,stable,False +20,vision,5,0.5,0.6,-0.08434701265650801,0.014922654396071152,64,stable,False +20,vision,6,0.6,0.7,-0.03815167227003258,0.01158295437313063,64,stable,False +20,vision,7,0.7,0.8,-0.07440588585450314,0.011802421548538522,64,stable,False +20,vision,8,0.8,0.9,0.25895043746277224,0.023195527211723146,64,stable,False +20,vision,9,0.9,1.0,0.09462569202878512,0.010956266291883248,64,stable,False diff --git a/final/output/q3/ati_ho/attachment4_prediction_manifest.json b/final/output/q3/ati_ho/attachment4_prediction_manifest.json new file mode 100644 index 0000000..27c84e0 --- /dev/null +++ b/final/output/q3/ati_ho/attachment4_prediction_manifest.json @@ -0,0 +1,71 @@ +{ + "task": "Q3 ATI–HO Attachment 4 final inference", + "selected_method": "A0", + "seeds": [ + 42, + 3407, + 2026 + ], + "input_version": "official unaligned_50 Attachment 4", + "adapter": "Q1AlignmentAdapter; Relative-Progress; target_steps=50", + "physical_time_alignment": false, + "no_attachment4_labels_or_metrics_used": true, + "scaler": "/home/gloamxun/modeling_zhaocui/final/experiments/q2/unaligned_deep_two_b128/unaligned_50_robust_stats.npz", + "feature_dimensions": [ + 768, + 74, + 35 + ], + "class_order": [ + "negative", + "neutral", + "positive" + ], + "intensity_decode": "predicted negative: -3*sigmoid(r_negative); neutral: 0; predicted positive: 3*sigmoid(r_positive)", + "model_checkpoint_sha256": { + "A0/seed_42": "bc5a3085a97da7b07be779f3d1337d02df00f75eb7d09f83cbd81e10f26c73e9", + "A0/seed_3407": "7a5556072ed49bd76c4d425a7531d589cac72d1caefa9281d165b4b0406fedad", + "A0/seed_2026": "160deb16472ccce3ec32916f2aa7173a011329880f67d90da98028cee7d99f6f", + "B0_early_concat/seed_42": "54e56e329349dbbe73a1e7563abe49411651461a0f9456d4606d944e72696674", + "B0_early_concat/seed_3407": "d4b2135fc433c78c46e0e5c29f40ec626548038cf8a2bd5aa981f1c3af1cbff1", + "B0_early_concat/seed_2026": "282fd6444e9444501df51e075a1c2a84235c10c82b52d342394927a94ef31d5e", + "B5_mofe_mlp/seed_42": "bce412cf08213376499f7ed35a5bc98ca96dc73030814facb57c56ab60f1de9d", + "B5_mofe_mlp/seed_3407": "687bbee3cebd50644a599e0148d892352b1593dc75d41adf017c2d5e7b211f30", + "B5_mofe_mlp/seed_2026": "72c0a154a73a373311a446ba90bd5f2aba3c4441862ba6c2d0f7eae08a6c3c3a" + }, + "attachment4_source_sha256": { + "01.pkl": "70e55e2709905d7ea9becf70d8dd91d83688e5902b6c5e601aa6d16e703e4659", + "02.pkl": "17e1a7c74f20dd25ca240b289add7a3b632ff2d6fb73d71eb1e1c6006b1f3cb9", + "03.pkl": "0d959289f51abebc34cb911895d7da4c70bb032d2fe281906b8f27ef6411cd8b", + "04.pkl": "f05343ac26cc037a6023a8e00ca1ffa7d5578db9bac6f2dad2e97bebf2ee188c", + "05.pkl": "cf0cd1b562e4516a25a04c23953f9baecaa0f42ce4044aa4eb9b1aba221e70a9", + "06.pkl": "3a5e9de011531ec38d60430053a221bfde7e05b253cbbf414b6ec08c42b9eee4", + "07.pkl": "b6f9a61599a16d0dee901ec277cdb32bde3ab50a223b15823fd0368b61f32314", + "08.pkl": "842aa20b0b30bd8cdef2d1f2e3c458abf01caaacb416b23a4dc8b7129ca660f3", + "09.pkl": "f42d91d2c1b8e6cad707eccd029c57eb1cb9e9300544ad1ba10bf0a96627ce65", + "10.pkl": "b12cb66e785c60faaa92d91be71bcae87dd4c0282bdd7e53241a7bb3e7413aa9", + "11.pkl": "64b0b19f12977bb24684cbb874d88e3dde71ad7f93821f8a382d05d3047c3005", + "12.pkl": "fd7cb86a2952ef384b0f1ce0dc58391332379e1d065d48100735b40ce991a624", + "13.pkl": "4d8b77cdc24eecaa5104f19eaff60bad12fbf1886f40f1d95491decc9b31f7e0", + "14.pkl": "17b0576012b71cdbebddfd38375e19f21fbe8e1eefa4813ac945196c734c9f55", + "15.pkl": "e12cb1d0b69558c7f81eecb2063769b872e58718051c7ea6b2761c066763f594", + "16.pkl": "359fae06a5df658729772eda07de43a12627a53574df95164942eb5f7a594153", + "17.pkl": "5d52d13cec0e716c6dab949e3296607d6f32bf320cb4029fde1d2b532101a41c", + "18.pkl": "7991357145a879fa8d85a028f043a6835243b24a6454bf5ec89c485cbc6335cf", + "19.pkl": "14041c0d6b2d351cd900603044f617d549e25821fa04b3f57d62c4a4e52fad4e", + "20.pkl": "2ed97e598f321dc82392e6b3d85963c629cea4ea8bc166705b697a7694ff7f44" + }, + "attachment4_feature_dir": "/home/gloamxun/modeling_zhaocui/E题数据/附件4-可解释专项视频样本与特征文件/附件4-可解释专项视频样本与特征文件/未对齐版本", + "prediction_rows": 20, + "explanation_rows": 20, + "local_evidence_rows": 600, + "shapley_audit": { + "samples": 20, + "analytic_class_shapley_mean_abs_error": 5.118134947827355e-08, + "analytic_class_shapley_median_abs_error": 3.601113956943486e-08, + "analytic_class_shapley_p95_abs_error": 1.554532597484309e-07, + "analytic_class_shapley_max_abs_error": 3.5235037376679657e-07, + "analytic_class_shapley_pass_rate": 1.0, + "decoded_intensity_shapley_max_efficiency_residual": 4.440892098500626e-16 + } +} diff --git a/final/output/q3/ati_ho/attachment4_predictions.csv b/final/output/q3/ati_ho/attachment4_predictions.csv new file mode 100644 index 0000000..96eb5f5 --- /dev/null +++ b/final/output/q3/ati_ho/attachment4_predictions.csv @@ -0,0 +1,21 @@ +case_id,predicted_class,predicted_intensity,prob_negative,prob_neutral,prob_positive,conditional_negative_magnitude,conditional_positive_magnitude,coordinate_mode,physical_time_alignment,text_observed_steps,audio_observed_steps,vision_observed_steps,true_label_available +01,neutral,0.0,0.09335917234420776,0.492781400680542,0.41385942697525024,0.7683022022247314,0.7918088436126709,relative_progress,False,50,50,50,False +02,positive,1.0754910707473755,0.16660422086715698,0.2815896272659302,0.5518062114715576,1.0387691259384155,1.0754910707473755,relative_progress,False,50,50,50,False +03,negative,-1.1482162475585938,0.44677555561065674,0.35495230555534363,0.19827206432819366,1.1482162475585938,1.054591417312622,relative_progress,False,50,50,50,False +04,negative,-1.1313621997833252,0.6154006123542786,0.08184759318828583,0.3027518689632416,1.1313621997833252,1.2142971754074097,relative_progress,False,50,50,50,False +05,positive,0.9983435869216919,0.019440362229943275,0.1781618446111679,0.8023977875709534,0.8999792337417603,0.9983435869216919,relative_progress,False,50,50,50,False +06,positive,1.2635153532028198,0.04338030517101288,0.12275035679340363,0.8338693380355835,1.089220643043518,1.2635153532028198,relative_progress,False,50,50,50,False +07,positive,0.7555097937583923,0.024313107132911682,0.20876865088939667,0.7669181823730469,0.748427152633667,0.7555097937583923,relative_progress,False,50,50,50,False +08,positive,0.906437873840332,0.018894990906119347,0.16367757320404053,0.8174274563789368,0.8192697763442993,0.906437873840332,relative_progress,False,50,50,50,False +09,negative,-1.5102474689483643,0.9667553901672363,0.016179252415895462,0.01706531085073948,1.5102474689483643,1.1577949523925781,relative_progress,False,50,50,50,False +10,negative,-1.4183083772659302,0.845405101776123,0.03885689750313759,0.11573806405067444,1.4183083772659302,1.3615806102752686,relative_progress,False,50,50,50,False +11,negative,-1.3695042133331299,0.6915677785873413,0.11874549835920334,0.18968670070171356,1.3695042133331299,1.220651626586914,relative_progress,False,50,50,50,False +12,negative,-1.1265901327133179,0.8298471570014954,0.11778073757886887,0.05237210541963577,1.1265901327133179,0.9112517237663269,relative_progress,False,50,50,50,False +13,positive,0.8690536022186279,0.05612684041261673,0.46271824836730957,0.4811549782752991,1.0537734031677246,0.8690536022186279,relative_progress,False,50,50,48,False +14,positive,0.9438773393630981,0.052272193133831024,0.43238261342048645,0.5153452157974243,0.8826612234115601,0.9438773393630981,relative_progress,False,50,50,50,False +15,positive,1.1347935199737549,0.00835143681615591,0.04237542673945427,0.9492731094360352,0.6205173134803772,1.1347935199737549,relative_progress,False,50,50,50,False +16,negative,-1.5337121486663818,0.9742470383644104,0.014193418435752392,0.01155958790332079,1.5337121486663818,1.0518016815185547,relative_progress,False,50,50,47,False +17,positive,1.2830696105957031,0.00731800589710474,0.02694552205502987,0.965736448764801,0.8673832416534424,1.2830696105957031,relative_progress,False,50,50,50,False +18,neutral,0.0,0.36853352189064026,0.4732035994529724,0.15826284885406494,1.0774030685424805,0.9822005033493042,relative_progress,False,50,50,50,False +19,positive,1.3935961723327637,0.4146651029586792,0.11420381814241409,0.4711310565471649,1.244720220565796,1.3935961723327637,relative_progress,False,50,50,50,False +20,positive,1.0786213874816895,0.018150269985198975,0.14706102013587952,0.8347886800765991,0.8807858228683472,1.0786213874816895,relative_progress,False,50,50,50,False diff --git a/final/q2/math/predict_attachment3.py b/final/q2/math/predict_attachment3.py index e2b1a45..27beb0d 100644 --- a/final/q2/math/predict_attachment3.py +++ b/final/q2/math/predict_attachment3.py @@ -1,102 +1,98 @@ -"""Run the saved Q2 student on the aligned, unlabeled attachment-3 cases.""" +"""Run a saved Q2 model on the unlabeled Attachment 3 cases.""" from __future__ import annotations +import argparse import json import time -import csv +from pathlib import Path import numpy as np import torch +from ...data_paths import PROJECT_ROOT from ...model.crg import INPUT_DIMS, MODALITIES, StructuredGaussianImputer -from .train import RESULTS, _make_variant, infer_attachment3, reencode_attachment3, validate_attachment3_predictions, write_csv +from .train import ( + _make_variant, + infer_attachment3, + reencode_attachment3, + validate_attachment3_predictions, + write_csv, +) + +DEFAULT_RESULTS_DIR = PROJECT_ROOT / "experiments" / "q2" / "unaligned_math_all_b128" +DEFAULT_OUTPUT_DIR = PROJECT_ROOT / "output" / "q2" def main() -> None: - device = torch.device("cuda" if torch.cuda.is_available() else "cpu") - manifest_path = RESULTS / "run_manifest.json" + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--input-version", choices=("aligned_50", "unaligned_50"), default="unaligned_50") + parser.add_argument("--results-dir", type=Path, default=DEFAULT_RESULTS_DIR, + help="saved Q2 checkpoint and calibration directory") + parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT_DIR) + 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) + + results_dir = args.results_dir.expanduser().resolve() + output_dir = args.output_dir.expanduser().resolve() + manifest_path = results_dir / "run_manifest.json" + calibration_path = results_dir / "validation_metrics.json" manifest = json.loads(manifest_path.read_text(encoding="utf-8")) - calibration = json.loads((RESULTS / "validation_metrics.json").read_text(encoding="utf-8")) + calibration = json.loads(calibration_path.read_text(encoding="utf-8")) selected = calibration.get("selected_model", manifest.get("selected_model")) if not selected: - raise ValueError("run_manifest.json does not identify a selected model") + raise ValueError(f"no selected_model recorded in {calibration_path}") imputer = StructuredGaussianImputer(INPUT_DIMS).to(device) - imputer_state = torch.load(RESULTS / "structured_imputer.pt", map_location=device, weights_only=True) - imputer.load_state_dict(imputer_state) + imputer.load_state_dict(torch.load(results_dir / "structured_imputer.pt", map_location=device, weights_only=True)) model = _make_variant(selected, imputer).to(device) - state = torch.load(RESULTS / "crg_student.pt", map_location=device, weights_only=True) - model.load_state_dict(state) + model.load_state_dict(torch.load(results_dir / "crg_student.pt", map_location=device, weights_only=True)) - with np.load(RESULTS / "preprocessor.npz", allow_pickle=False) as archive: + with np.load(results_dir / "preprocessor.npz", allow_pickle=False) as archive: fitted = {m: {k: archive[f"{m}_{k}"].copy() for k in ("mean", "std")} for m in MODALITIES} priors = manifest["attachment3_low_information_priors"] temperature = float(calibration["temperature"]) class_prior = np.asarray(priors["class_probability_values"], dtype=np.float64) magnitude_priors = np.asarray((priors["negative_beta"], priors["positive_beta"]), dtype=np.float32) - cases, source_audit = reencode_attachment3(device) + cases, source_audit = reencode_attachment3(device, input_version=args.input_version) predictions, inference_audit = infer_attachment3( model, cases, fitted, device, temperature, class_prior, magnitude_priors, ) validate_attachment3_predictions([case["case_id"] for case in cases], predictions) - inference_by_id = {row["case_id"]: row for row in inference_audit} - write_csv(RESULTS / "attachment3_predictions.csv", predictions) - write_csv(RESULTS / "attachment3_audit.csv", [ - {**source, **inference_by_id[source["case_id"]]} for source in source_audit - ]) - # The training script can finish and persist all labeled-evaluation outputs - # before an unlabeled attachment export fails. Reconcile the manifest from - # those completed artifacts so the standalone export is safely rerunnable. - group_risk_rows = list(csv.DictReader((RESULTS / "group_risk_tuning.csv").open(encoding="utf-8-sig", newline=""))) - selected_risk = next((row for row in group_risk_rows if row.get("selected", "").lower() == "true"), None) - reliability_rows = list(csv.DictReader((RESULTS / "reliability_hparam_tuning.csv").open(encoding="utf-8-sig", newline=""))) - # split_calibration's generic internal names are canonicalized in train.py; - # repair artifacts from runs produced before that naming fix as well. - for row in group_risk_rows: - if row.get("selection_split") == "fit": - row["selection_split"] = "reliability_validation" - for row in reliability_rows: - if row.get("selection_split") == "fit": - row["selection_split"] = "reliability_validation" - write_csv(RESULTS / "group_risk_tuning.csv", group_risk_rows) - write_csv(RESULTS / "reliability_hparam_tuning.csv", reliability_rows) - if selected_risk: - risk_values = (float(selected_risk["lambda_group"]), float(selected_risk["group_temperature"])) - manifest["group_risk_hyperparameters"]["selected"] = list(risk_values) - manifest["loss"]["selected_group_risk"] = list(risk_values) - manifest["group_risk_hyperparameters"]["selection_split"] = "reliability_validation" - manifest["reliability_hyperparameters"]["selected_by_model"] = { - row["model"]: [float(row[key]) for key in ("rho_imp", "lambda_u", "lambda_gap", "lambda_span")] - for row in reliability_rows - if row.get("selected", "").lower() == "true" - and (not row.get("risk_candidate_selected") or row["risk_candidate_selected"].lower() == "true") - } - test_metrics = json.loads((RESULTS / "test_metrics.json").read_text(encoding="utf-8")) - manifest["selected_model"] = selected - manifest["final_test_metrics"] = test_metrics - manifest["calibration"]["temperature"] = temperature - manifest["calibration"]["valid_used_for_selection"] = True - manifest["calibration"]["test_used_for_selection_or_calibration"] = False - manifest["training_configuration"].update({ - "student_epoch_limit": 12, - "imputer_epochs": 8, - "batch_size": 64, - "early_stopping_patience": 3, - }) - manifest["imputer"]["epochs"] = 8 - manifest.update({ + output_dir.mkdir(parents=True, exist_ok=True) + inference_by_id = {row["case_id"]: row for row in inference_audit} + predictions_path = output_dir / "attachment3_predictions.csv" + audit_path = output_dir / "attachment3_audit.csv" + manifest_out_path = output_dir / "attachment3_prediction_manifest.json" + write_csv(predictions_path, predictions) + write_csv(audit_path, [{**source, **inference_by_id[source["case_id"]]} for source in source_audit]) + + try: + results_reference = results_dir.relative_to(PROJECT_ROOT).as_posix() + except ValueError: + results_reference = "external checkpoint directory" + prediction_manifest = { + "task": "unlabeled Attachment 3 inference", + "input_version": args.input_version, + "selected_model": selected, + "checkpoint_run": results_reference, + "prediction_count": len(predictions), + "temperature": temperature, + "labels_available": False, + "prediction_file": predictions_path.name, + "audit_file": audit_path.name, "completed_utc": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()), - "attachment3_cases": len(cases), - "attachment3_prediction_file": "attachment3_predictions.csv", - "attachment3_audit_file": "attachment3_audit.csv", - "attachment3_labeled_metrics": None, - "quality_flags": {m: "unavailable; q*=1 fallback for visible rows, unknown flag retained" for m in MODALITIES}, - "neutral_output": "exact zero when neutral is the predicted class; no near-zero threshold", - }) - manifest_path.write_text(json.dumps(manifest, ensure_ascii=False, indent=2), encoding="utf-8") - print(f"Wrote {len(predictions)} unlabeled attachment-3 predictions to {RESULTS}", flush=True) + } + manifest_out_path.write_text(json.dumps(prediction_manifest, ensure_ascii=False, indent=2), encoding="utf-8") + print(f"Wrote {len(predictions)} unlabeled Attachment 3 predictions to {output_dir}", flush=True) if __name__ == "__main__": diff --git a/final/q3/README.md b/final/q3/README.md index ae18619..88081a4 100644 --- a/final/q3/README.md +++ b/final/q3/README.md @@ -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 交付文件。 diff --git a/final/q3/ati_ho/README.md b/final/q3/ati_ho/README.md new file mode 100644 index 0000000..f52c9b3 --- /dev/null +++ b/final/q3/ati_ho/README.md @@ -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`。 diff --git a/final/q3/ati_ho/__init__.py b/final/q3/ati_ho/__init__.py new file mode 100644 index 0000000..a6ff6f1 --- /dev/null +++ b/final/q3/ati_ho/__init__.py @@ -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"] diff --git a/final/q3/ati_ho/attribution.py b/final/q3/ati_ho/attribution.py new file mode 100644 index 0000000..105fb5e --- /dev/null +++ b/final/q3/ati_ho/attribution.py @@ -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, + } diff --git a/final/q3/ati_ho/audit.py b/final/q3/ati_ho/audit.py new file mode 100644 index 0000000..86a856a --- /dev/null +++ b/final/q3/ati_ho/audit.py @@ -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, + } diff --git a/final/q3/ati_ho/evaluate.py b/final/q3/ati_ho/evaluate.py new file mode 100644 index 0000000..d6e7453 --- /dev/null +++ b/final/q3/ati_ho/evaluate.py @@ -0,0 +1,1296 @@ +from __future__ import annotations + +import argparse +import csv +import hashlib +import itertools +import json +import math +import time +from collections import defaultdict +from pathlib import Path +from typing import Any, Sequence + +import numpy as np +import torch +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 PROJECT_ROOT +from ...model.ati_ho import ATIHOModel +from ...model.ati_ho_config import CONFIGS +from ...q2.deep_learning.q2.data import MODALITIES, RobustStats, Split +from ...q2.deep_learning.q2.mofe import MixtureOfFusionExperts +from ...q2.deep_learning.q2.train_mofe import EARLYCONCAT, MODEL_CONFIG, MOFE7_MLP +from ..run_experiments import _read_attachment4 +from .attribution import ensemble_forward, exact_shapley_audit +from .audit import structural_audit +from .owen import fidelity_audit_one, hierarchical_owen_one +from .train import ( + BATCH_SIZE, + EXPERIMENT_ROOT, + MODEL_SEEDS, + SCALER_PATH, + _metric_row, + _save_csv, + load_training_data, +) + + +RESULTS_ROOT = EXPERIMENT_ROOT / "results" / "ati_ho" +SUBMIT_OUTPUT = PROJECT_ROOT / "output" / "q3" / "ati_ho" +CLASS_NAMES = ("negative", "neutral", "positive") +PAIR_NAMES = ("TA", "TV", "AV") +BOOTSTRAP_REPLICATES = 1000 +BOOTSTRAP_SEED = 20260925 + + +def _read_csv(path: Path) -> list[dict[str, str]]: + if not path.is_file(): + return [] + with path.open("r", newline="", encoding="utf-8-sig") as stream: + return list(csv.DictReader(stream)) + + +def _write_json(path: Path, payload: Any) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text(json.dumps(payload, ensure_ascii=False, indent=2, allow_nan=False) + "\n", encoding="utf-8") + + +def _sha256(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as stream: + for block in iter(lambda: stream.read(1024 * 1024), b""): + digest.update(block) + return digest.hexdigest() + + +def _checkpoint_path(method: str, seed: int) -> Path: + return EXPERIMENT_ROOT / "models" / method / f"seed_{seed}" / "model_best.pt" + + +def _load_ensemble(method: str, seeds: Sequence[int], dims: tuple[int, int, int], device: torch.device) -> list[nn.Module]: + models: list[nn.Module] = [] + for seed in seeds: + path = _checkpoint_path(method, seed) + if not path.is_file(): + raise FileNotFoundError(path) + state = torch.load(path, map_location=device, weights_only=False) + if ( + tuple(state.get("dims", ())) != dims + or int(state.get("seed", -1)) != seed + or state.get("method") != method + ): + raise ValueError(f"incompatible checkpoint metadata: {path}") + if method in CONFIGS: + if state.get("config") != CONFIGS[method].to_dict(): + raise ValueError(f"ATI configuration mismatch: {path}") + model: nn.Module = ATIHOModel(dims, CONFIGS[method]).to(device) + elif method == EARLYCONCAT: + from ...q2.deep_learning.q2.models import AlignedFusionModel + + model = AlignedFusionModel("concat", dims=dims).to(device) + elif method == MOFE7_MLP: + model = MixtureOfFusionExperts(dims=dims, **MODEL_CONFIG).to(device) + else: + raise ValueError(f"unknown model {method}") + model.load_state_dict(state["state_dict"], strict=True) + model.eval() + models.append(model) + return models + + +@torch.inference_mode() +def _predict_ensemble( + models: Sequence[nn.Module], + split: Split, + masks: np.ndarray, + device: torch.device, + *, + details: bool = False, + batch_size: int = 64, +) -> dict[str, np.ndarray]: + keys = ["logits", "probabilities", "intensity"] + if details: + keys.extend(("params", "baseline", "main_effects", "pair_effects")) + values: dict[str, list[np.ndarray]] = {key: [] for key in keys} + 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) + output = ensemble_forward(models, xs, mask, details=details) + for key in keys: + if key in output: + values[key].append(output[key].detach().cpu().numpy()) + return {key: np.concatenate(rows, axis=0) for key, rows in values.items() if rows} + + +def _metric_subset( + split: Split, + predictions: dict[str, np.ndarray], + indices: np.ndarray, +) -> dict[str, float]: + y_cls = split.y_cls[indices] + y_reg = split.y_reg[indices] + logits = predictions["logits"][indices] + intensity = predictions["intensity"][indices] + probs = predictions["probabilities"][indices] + pred_class = logits.argmax(axis=-1) + recall = recall_score(y_cls, pred_class, labels=[0, 1, 2], average=None, zero_division=0) + pearson = float(np.corrcoef(y_reg, intensity)[0, 1]) if np.std(y_reg) > 0 and np.std(intensity) > 0 else 0.0 + one_hot = np.eye(3)[y_cls] + return { + "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[0]), + "neutral_recall": float(recall[1]), + "positive_recall": float(recall[2]), + "mae": float(mean_absolute_error(y_reg, intensity)), + "rmse": float(math.sqrt(mean_squared_error(y_reg, intensity))), + "pearson": pearson, + "brier_multiclass": float(np.mean(np.sum((probs - one_hot) ** 2, axis=-1))), + } + + +def _cluster_bootstrap( + split: Split, + prediction_by_method: dict[str, dict[str, np.ndarray]], + selected_method: str, +) -> list[dict[str, Any]]: + group_rows: dict[str, list[int]] = defaultdict(list) + for index, sample_id in enumerate(split.ids): + group_rows[sample_id.split("$_$", 1)[0]].append(index) + groups = np.asarray(sorted(group_rows)) + group_map = {key: np.asarray(value, dtype=np.int64) for key, value in group_rows.items()} + rng = np.random.default_rng(BOOTSTRAP_SEED) + metrics = ("macro_f1", "mae", "pearson", "accuracy") + rows: list[dict[str, Any]] = [] + for baseline in (EARLYCONCAT, MOFE7_MLP): + point_a = _metric_subset(split, prediction_by_method[selected_method], np.arange(split.n)) + point_b = _metric_subset(split, prediction_by_method[baseline], np.arange(split.n)) + draws: dict[str, list[float]] = {metric: [] for metric in metrics} + for _ in range(BOOTSTRAP_REPLICATES): + selected_groups = rng.choice(groups, size=len(groups), replace=True) + index = np.concatenate([group_map[group] for group in selected_groups]) + a = _metric_subset(split, prediction_by_method[selected_method], index) + b = _metric_subset(split, prediction_by_method[baseline], index) + for metric in metrics: + delta = a[metric] - b[metric] + if metric == "mae": + delta = b[metric] - a[metric] + draws[metric].append(delta) + for metric in metrics: + values = np.asarray(draws[metric], dtype=np.float64) + delta = point_a[metric] - point_b[metric] + if metric == "mae": + delta = point_b[metric] - point_a[metric] + rows.append( + { + "comparison": f"{selected_method} vs {baseline}", + "metric": metric, + "delta_positive_favors_ATI_HO": float(delta), + "bootstrap_ci_2p5": float(np.quantile(values, 0.025)), + "bootstrap_ci_97p5": float(np.quantile(values, 0.975)), + "replicates": BOOTSTRAP_REPLICATES, + "bootstrap_unit": "source video_id", + "groups": len(groups), + "seed": BOOTSTRAP_SEED, + } + ) + return rows + + +def _summary_rows(validation_rows: list[dict[str, str]], methods: Sequence[str]) -> list[dict[str, Any]]: + latest: dict[tuple[str, str, str], dict[str, str]] = {} + for row in validation_rows: + key = (row["method"], row["seed"], row["scenario"]) + latest[key] = row + metrics = ( + "accuracy", "macro_f1", "weighted_f1", "negative_recall", "neutral_recall", "positive_recall", + "mae", "rmse", "pearson", "ece_15bin", "brier_multiclass", + ) + rows: list[dict[str, Any]] = [] + for method in methods: + seeds = sorted({key[1] for key in latest if key[0] == method and key[2] == "clean"}, key=int) + for scenario in ("clean", "0.0/none", "0.3/single", "0.3/sync", "0.5/async"): + matched = [latest[(method, seed, scenario)] for seed in seeds if (method, seed, scenario) in latest] + if not matched: + continue + row: dict[str, Any] = {"method": method, "scenario": scenario, "seeds": len(matched), "seed_values": ";".join(x["seed"] for x in matched)} + for metric in metrics: + values = np.asarray([float(item[metric]) for item in matched], dtype=np.float64) + row[f"{metric}_mean"] = float(values.mean()) + row[f"{metric}_std"] = float(values.std(ddof=1)) if len(values) > 1 else 0.0 + rows.append(row) + return rows + + +def _write_seed_and_summary_tables() -> tuple[str, list[dict[str, Any]], list[dict[str, Any]]]: + selection_path = EXPERIMENT_ROOT / "stage2_complete.json" + if not selection_path.is_file(): + raise FileNotFoundError("Stage II is incomplete; stage2_complete.json is missing") + stage2 = json.loads(selection_path.read_text(encoding="utf-8")) + candidate_methods = list(stage2["key_ablations"]) + provisional = stage2["selected_candidate"] + candidate_methods = list(dict.fromkeys([provisional, *candidate_methods])) + selection_rows = [] + for method in candidate_methods: + values = [] + for seed in MODEL_SEEDS: + state = torch.load(_checkpoint_path(method, seed), map_location="cpu", weights_only=False) + values.append(float(state["best_selection_loss"])) + selection_rows.append( + { + "method": method, + "seed_losses": json.dumps(values), + "mean_validation_selection_loss": float(np.mean(values)), + "std_validation_selection_loss": float(np.std(values, ddof=1)), + "seeds": len(values), + "validation_only_selection": True, + } + ) + selection_rows.sort(key=lambda row: row["mean_validation_selection_loss"]) + selected = selection_rows[0]["method"] + selected_record = { + "selected_method": selected, + "provisional_seed42_method": provisional, + "candidate_methods_with_three_seeds": candidate_methods, + "selection_rule": "lowest mean fixed four-scenario validation task loss across seeds 42, 3407, 2026", + "candidate_summary": selection_rows, + "attachment4_labels_used": False, + } + _write_json(EXPERIMENT_ROOT / "final_selection.json", selected_record) + validation_rows = _read_csv(EXPERIMENT_ROOT / "validation_results.csv") + methods = [EARLYCONCAT, MOFE7_MLP, *candidate_methods, "A3", "D0"] + seed_rows = [row for row in validation_rows if row.get("method") in methods] + _save_csv(RESULTS_ROOT / "seed_results.csv", seed_rows) + summary_rows = _summary_rows(seed_rows, methods) + clean = [row for row in summary_rows if row["scenario"] == "clean"] + _save_csv(RESULTS_ROOT / "main_results.csv", clean) + ablations = [row for row in clean if row["method"] in {"A0", "A1", "A2", "A3", "D0"}] + _save_csv(RESULTS_ROOT / "ablation_results.csv", ablations) + _save_csv(RESULTS_ROOT / "selection_results.csv", selection_rows) + return selected, clean, selection_rows + + +def _to_device(split: Split, device: torch.device) -> tuple[tuple[torch.Tensor, torch.Tensor, torch.Tensor], torch.Tensor]: + return ( + tuple(torch.as_tensor(x, dtype=torch.float32, device=device) for x in split.x), + torch.as_tensor(split.mask, dtype=torch.bool, device=device), + ) + + +def _attachment_split(cases: list[dict[str, Any]], stats: RobustStats) -> Split: + xs = tuple( + np.stack([case["features"][modality] for case in cases]).astype(np.float32) + for modality in range(len(MODALITIES)) + ) + mask = np.stack([case["mask"] for case in cases]).astype(bool) + normalized: list[np.ndarray] = [] + for modality, values in enumerate(xs): + current = (values - stats.center[modality]) / stats.scale[modality] + current = np.nan_to_num(current, nan=0.0, posinf=0.0, neginf=0.0) + current *= mask[..., modality, None] + normalized.append(current.astype(np.float32, copy=False)) + n = len(cases) + return Split(tuple(normalized), mask, np.full(n, -1, dtype=np.int64), np.full(n, np.nan, dtype=np.float32), [case["case_id"] for case in cases]) + + +def _class_name(value: int) -> str: + return CLASS_NAMES[int(value)] + + +def _attachment_predictions_and_explanations( + selected: str, + models: Sequence[nn.Module], + cases: list[dict[str, Any]], + attachment: Split, + device: torch.device, +) -> tuple[list[dict[str, Any]], list[dict[str, Any]], dict[str, Any], dict[str, np.ndarray]]: + xs, masks = _to_device(attachment, device) + exact = exact_shapley_audit(models, xs, masks, batch_size=64) + output = exact["full_output"] + predictions: list[dict[str, Any]] = [] + explanations: list[dict[str, Any]] = [] + for index, case in enumerate(cases): + target = int(exact["target_class"][index]) + other = int(exact["runner_up_class"][index]) + params = output["params"][index].detach().cpu().numpy() + baseline = output["baseline"][index].detach().cpu().numpy() + main = output["main_effects"][index].detach().cpu().numpy() + pairs = output["pair_effects"][index].detach().cpu().numpy() + probabilities = output["probabilities"][index].detach().cpu().numpy() + pred_intensity = float(output["intensity"][index].item()) + row: dict[str, Any] = { + "case_id": case["case_id"], + "predicted_class": _class_name(target), + "predicted_intensity": pred_intensity, + "prob_negative": float(probabilities[0]), + "prob_neutral": float(probabilities[1]), + "prob_positive": float(probabilities[2]), + "conditional_negative_magnitude": float(output["nu_negative"][index].item()), + "conditional_positive_magnitude": float(output["nu_positive"][index].item()), + "coordinate_mode": "relative_progress", + "physical_time_alignment": False, + "text_observed_steps": int(attachment.mask[index, :, 0].sum()), + "audio_observed_steps": int(attachment.mask[index, :, 1].sum()), + "vision_observed_steps": int(attachment.mask[index, :, 2].sum()), + "true_label_available": False, + } + predictions.append(row) + explanation: dict[str, Any] = { + "case_id": case["case_id"], + "fixed_target_class": _class_name(target), + "fixed_runner_up_class": _class_name(other), + "full_logit_margin": float(output["logits"][index, target].item() - output["logits"][index, other].item()), + "baseline_r_negative": float(baseline[3]), + "baseline_r_positive": float(baseline[4]), + "analytic_vs_exact_shapley_max_abs": float(exact["class_abs_error"][index].max()), + "analytic_vs_exact_shapley_all_pass": bool(exact["class_pass"][index].all()), + "exact_intensity_shapley_sum": float(exact["exact_intensity"][index].sum()), + "intensity_full_minus_empty_coalition": float(exact["coalition_intensity"][index, 7] - exact["coalition_intensity"][index, 0]), + "intensity_shapley_efficiency_residual": float(exact["exact_intensity"][index].sum() - (exact["coalition_intensity"][index, 7] - exact["coalition_intensity"][index, 0])), + "coordinate_mode": "relative_progress", + "physical_time_alignment": False, + } + for modality, label in enumerate(("T", "A", "V")): + for parameter, suffix in enumerate(("logit_negative", "logit_neutral", "logit_positive", "r_negative", "r_positive")): + explanation[f"G_{label}_{suffix}"] = float(main[modality, parameter]) + explanation[f"analytic_class_shapley_{label}"] = float(exact["analytic_class"][index, modality]) + explanation[f"exact_class_shapley_{label}"] = float(exact["exact_class"][index, modality]) + explanation[f"exact_intensity_shapley_{label}"] = float(exact["exact_intensity"][index, modality]) + for pair_index, pair_name in enumerate(PAIR_NAMES): + for parameter, suffix in enumerate(("logit_negative", "logit_neutral", "logit_positive", "r_negative", "r_positive")): + explanation[f"G_{pair_name}_{suffix}"] = float(pairs[pair_index, parameter]) + for parameter, suffix in enumerate(("logit_negative", "logit_neutral", "logit_positive", "r_negative", "r_positive")): + explanation[f"baseline_{suffix}"] = float(baseline[parameter]) + explanation[f"full_parameter_{suffix}"] = float(params[parameter]) + explanations.append(explanation) + summary = { + "samples": len(cases), + "analytic_class_shapley_mean_abs_error": float(exact["class_abs_error"].mean()), + "analytic_class_shapley_median_abs_error": float(np.median(exact["class_abs_error"])), + "analytic_class_shapley_p95_abs_error": float(np.quantile(exact["class_abs_error"], 0.95)), + "analytic_class_shapley_max_abs_error": float(exact["class_abs_error"].max()), + "analytic_class_shapley_pass_rate": float(exact["class_pass"].mean()), + "decoded_intensity_shapley_max_efficiency_residual": float( + np.max(np.abs(exact["exact_intensity"].sum(axis=1) - (exact["coalition_intensity"][:, 7] - exact["coalition_intensity"][:, 0]))) + ), + } + return predictions, explanations, summary, exact + + +def _router_bins(models: Sequence[nn.Module], xs: tuple[torch.Tensor, torch.Tensor, torch.Tensor], mask: torch.Tensor, bins: int = 10) -> np.ndarray: + outputs = [] + with torch.inference_mode(): + for model in models: + out = model(xs, mask) + if "utility" not in out: + raise ValueError("MoFE checkpoint lacks its seven-expert modality utility"); + outputs.append(out["utility"]) + utility = torch.stack(outputs, dim=0).mean(dim=0)[0].detach().cpu().numpy() + observed = mask[0].detach().cpu().numpy() + steps = observed.shape[0] + edges = np.linspace(0, steps, bins + 1).round().astype(int) + result = np.zeros((3, bins), dtype=np.float64) + for modality in range(3): + for bin_index in range(bins): + left, right = int(edges[bin_index]), int(edges[bin_index + 1]) + visible = observed[left:right, modality] + if visible.any(): + result[modality, bin_index] = float(utility[left:right, modality][visible].mean()) + return result + + +def _owen_and_fidelity( + selected_models: Sequence[nn.Module], + early_models: Sequence[nn.Module], + mofe_models: Sequence[nn.Module], + cases: list[dict[str, Any]], + attachment: Split, + valid: Split, + device: torch.device, +) -> tuple[list[dict[str, Any]], list[dict[str, Any]], list[dict[str, Any]], list[dict[str, Any]], list[dict[str, Any]]]: + local_rows: list[dict[str, Any]] = [] + owen_rows: list[dict[str, Any]] = [] + fidelity_rows: list[dict[str, Any]] = [] + comparison_rows: list[dict[str, Any]] = [] + stability_rows: list[dict[str, Any]] = [] + attachment_contributions: list[dict[str, np.ndarray]] = [] + timings: list[float] = [] + boundaries = np.linspace(0, 50, 11).round().astype(int) + + for index, case in enumerate(cases): + xs = tuple(torch.as_tensor(x[index : index + 1], dtype=torch.float32, device=device) for x in attachment.x) + mask = torch.as_tensor(attachment.mask[index : index + 1], dtype=torch.bool, device=device) + start = time.perf_counter() + result = hierarchical_owen_one( + selected_models, xs, mask, seed=20260926 + index, start_permutations=8, max_permutations=64 + ) + timings.append(time.perf_counter() - start) + attachment_contributions.append({"ATI_HO_Owen": result["contribution"]}) + owen_rows.append( + { + "case_id": case["case_id"], + "elapsed_seconds": timings[-1], + "permutations": result["permutations"], + "stopping_status": result["stopping_status"], + "top5_jaccard_last_check": result["top5_jaccard_last_check"], + "local_conservation_residual": result["local_conservation_residual"], + "full_margin": result["full_margin"], + "baseline_margin": result["baseline_margin"], + "target_class": _class_name(result["target_class"]), + "runner_up_class": _class_name(result["runner_up_class"]), + } + ) + for modality, label in enumerate(("text", "audio", "vision")): + 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, + } + ) + + model_sets = ( + ("ATI_HO_Owen", selected_models, result["contribution"]), + ) + # Comparable post-hoc Owen scores from the two retrained prediction baselines. + for label, models in (("EarlyConcat_posthoc_Owen", early_models), ("MoFE_posthoc_Owen", mofe_models)): + baseline_owen = hierarchical_owen_one( + models, xs, mask, seed=20300000 + index, start_permutations=8, max_permutations=8 + ) + contribution = baseline_owen["contribution"] + model_sets += ((label, models, contribution),) + for modality, modality_label in enumerate(("text", "audio", "vision")): + for bin_index, (left, right) in enumerate(baseline_owen["bin_slices"]): + comparison_rows.append( + { + "case_id": case["case_id"], + "method": label, + "modality": modality_label, + "relative_bin": bin_index, + "relative_position_start": left / 50.0, + "relative_position_end": right / 50.0, + "posthoc_owen_margin_contribution": float(contribution[modality, bin_index]), + "permutations": baseline_owen["permutations"], + } + ) + router = _router_bins(mofe_models, xs, mask) + model_sets += (("MoFE_router_utility", mofe_models, router),) + for name, model_set, contribution in model_sets: + rows = fidelity_audit_one( + model_set, + xs, + mask, + contribution, + sample_id=case["case_id"], + seed=20270000 + index, + random_replicates=20, + ) + for row in rows: + fidelity_rows.append({"split": "attachment4_unlabelled", "explanation": name, **row}) + + # Validation fidelity is measured on a class-stratified, pre-fixed 60-row diagnostic sample. + rng = np.random.default_rng(20260927) + selected_indices: list[int] = [] + for label in (0, 1, 2): + available = np.flatnonzero(valid.y_cls == label) + count = min(20, len(available)) + selected_indices.extend(rng.choice(available, size=count, replace=False).tolist()) + for index in sorted(selected_indices): + xs = tuple(torch.as_tensor(x[index : index + 1], dtype=torch.float32, device=device) for x in valid.x) + mask = torch.as_tensor(valid.mask[index : index + 1], dtype=torch.bool, device=device) + result = hierarchical_owen_one( + selected_models, xs, mask, seed=20280000 + index, start_permutations=8, max_permutations=8 + ) + rows = fidelity_audit_one( + selected_models, + xs, + mask, + result["contribution"], + sample_id=valid.ids[index], + seed=20290000 + index, + random_replicates=20, + ) + for row in rows: + fidelity_rows.append({"split": "validation_class_stratified", "explanation": "ATI_HO_Owen", **row}) + + # Training-seed variability on five fixed Attachment 4 cases. + for index in range(min(5, len(cases))): + xs = tuple(torch.as_tensor(x[index : index + 1], dtype=torch.float32, device=device) for x in attachment.x) + mask = torch.as_tensor(attachment.mask[index : index + 1], dtype=torch.bool, device=device) + by_seed = [] + for seed in MODEL_SEEDS: + seed_model = _load_ensemble("A0", [seed], tuple(x.shape[-1] for x in attachment.x), device) + estimate = hierarchical_owen_one( + seed_model, xs, mask, seed=20310000 + index + seed, start_permutations=8, max_permutations=8 + ) + by_seed.append(estimate["contribution"].reshape(-1)) + for left_idx, right_idx in itertools.combinations(range(len(MODEL_SEEDS)), 2): + left, right = by_seed[left_idx], by_seed[right_idx] + corr = float(np.corrcoef(left, right)[0, 1]) if left.std() > 0 and right.std() > 0 else 0.0 + top_left = set(np.argsort(-np.abs(left))[:5]) + top_right = set(np.argsort(-np.abs(right))[:5]) + top_jaccard = len(top_left & top_right) / max(1, len(top_left | top_right)) + dominant_left = int(np.abs(by_seed[left_idx].reshape(3, 10)).sum(axis=1).argmax()) + dominant_right = int(np.abs(by_seed[right_idx].reshape(3, 10)).sum(axis=1).argmax()) + stability_rows.append( + { + "sample_id": cases[index]["case_id"], + "stability_source": "training_seed", + "seed_a": MODEL_SEEDS[left_idx], + "seed_b": MODEL_SEEDS[right_idx], + "signed_contribution_correlation": corr, + "top5_evidence_jaccard": top_jaccard, + "dominant_modality_agreement": dominant_left == dominant_right, + "seed_a_dominant_modality": ("text", "audio", "vision")[dominant_left], + "seed_b_dominant_modality": ("text", "audio", "vision")[dominant_right], + } + ) + + # One small 1% normalized-feature noise perturbation on the same five cases. + ensemble_base = selected_models + for index in range(min(5, len(cases))): + x_base = tuple(torch.as_tensor(x[index : index + 1], dtype=torch.float32, device=device) for x in attachment.x) + mask = torch.as_tensor(attachment.mask[index : index + 1], dtype=torch.bool, device=device) + base_estimate = hierarchical_owen_one( + ensemble_base, x_base, mask, seed=20320000 + index, start_permutations=8, max_permutations=8 + ) + generator = torch.Generator(device=device).manual_seed(20330000 + index) + x_perturbed = [] + for modality, x in enumerate(x_base): + noise = torch.randn(x.shape, dtype=x.dtype, device=device, generator=generator) * 0.01 + x_perturbed.append(x + noise * mask[..., modality, None].to(x.dtype)) + perturbed_estimate = hierarchical_owen_one( + ensemble_base, tuple(x_perturbed), mask, seed=20320000 + index, start_permutations=8, max_permutations=8 + ) + left = base_estimate["contribution"].reshape(-1) + right = perturbed_estimate["contribution"].reshape(-1) + corr = float(np.corrcoef(left, right)[0, 1]) if left.std() > 0 and right.std() > 0 else 0.0 + top_left = set(np.argsort(-np.abs(left))[:5]) + top_right = set(np.argsort(-np.abs(right))[:5]) + stability_rows.append( + { + "sample_id": cases[index]["case_id"], + "stability_source": "input_perturbation_1pct", + "seed_a": "base", + "seed_b": "gaussian_0.01", + "signed_contribution_correlation": corr, + "top5_evidence_jaccard": len(top_left & top_right) / max(1, len(top_left | top_right)), + "dominant_modality_agreement": np.abs(base_estimate["contribution"]).sum(axis=1).argmax() == np.abs(perturbed_estimate["contribution"]).sum(axis=1).argmax(), + } + ) + + return local_rows, owen_rows, fidelity_rows, comparison_rows, stability_rows + + +def _validation_predictions_and_shapley( + selected: str, + models: Sequence[nn.Module], + valid: Split, + device: torch.device, +) -> tuple[dict[str, np.ndarray], dict[str, Any], list[dict[str, Any]]]: + xs, masks = _to_device(valid, device) + started = time.perf_counter() + audit = exact_shapley_audit(models, xs, masks, batch_size=128) + elapsed = time.perf_counter() - started + full = audit["full_output"] + predictions = { + "logits": full["logits"].detach().cpu().numpy(), + "probabilities": full["probabilities"].detach().cpu().numpy(), + "intensity": full["intensity"].detach().cpu().numpy(), + } + rows = [] + for index, sample_id in enumerate(valid.ids): + rows.append( + { + "sample_id": sample_id, + "source_video_id": sample_id.split("$_$", 1)[0], + "method": selected, + "target_class": int(audit["target_class"][index]), + "runner_up_class": int(audit["runner_up_class"][index]), + "analytic_T": float(audit["analytic_class"][index, 0]), + "analytic_A": float(audit["analytic_class"][index, 1]), + "analytic_V": float(audit["analytic_class"][index, 2]), + "exact_T": float(audit["exact_class"][index, 0]), + "exact_A": float(audit["exact_class"][index, 1]), + "exact_V": float(audit["exact_class"][index, 2]), + "max_abs_error": float(audit["class_abs_error"][index].max()), + "all_modalities_pass": bool(audit["class_pass"][index].all()), + } + ) + summary = { + "samples": valid.n, + "elapsed_seconds": elapsed, + "mean_abs_error": float(audit["class_abs_error"].mean()), + "median_abs_error": float(np.median(audit["class_abs_error"])), + "p95_abs_error": float(np.quantile(audit["class_abs_error"], 0.95)), + "max_abs_error": float(audit["class_abs_error"].max()), + "pass_rate": float(audit["class_pass"].mean()), + "pass_tolerance": "absolute 1e-6 + relative 1e-5", + } + return predictions, summary, rows + + +def _stability_diagnostics( + selected: str, + cases: list[dict[str, Any]], + attachment: Split, + device: torch.device, +) -> list[dict[str, Any]]: + """Measure attribution variation across training seeds and small input noise.""" + stability_rows: list[dict[str, Any]] = [] + dims = tuple(int(x.shape[-1]) for x in attachment.x) + for index in range(min(5, len(cases))): + xs = tuple(torch.as_tensor(x[index : index + 1], dtype=torch.float32, device=device) for x in attachment.x) + mask = torch.as_tensor(attachment.mask[index : index + 1], dtype=torch.bool, device=device) + by_seed: list[np.ndarray] = [] + for seed in MODEL_SEEDS: + seed_model = _load_ensemble(selected, [seed], dims, device) + estimate = hierarchical_owen_one( + seed_model, xs, mask, seed=20310000 + index + seed, start_permutations=8, max_permutations=8 + ) + by_seed.append(estimate["contribution"].reshape(-1)) + for left_idx, right_idx in itertools.combinations(range(len(MODEL_SEEDS)), 2): + left, right = by_seed[left_idx], by_seed[right_idx] + corr = float(np.corrcoef(left, right)[0, 1]) if left.std() > 0 and right.std() > 0 else 0.0 + top_left = set(np.argsort(-np.abs(left))[:5]) + top_right = set(np.argsort(-np.abs(right))[:5]) + top_jaccard = len(top_left & top_right) / max(1, len(top_left | top_right)) + dominant_left = int(np.abs(by_seed[left_idx].reshape(3, 10)).sum(axis=1).argmax()) + dominant_right = int(np.abs(by_seed[right_idx].reshape(3, 10)).sum(axis=1).argmax()) + stability_rows.append( + { + "sample_id": cases[index]["case_id"], + "stability_source": "training_seed", + "seed_a": MODEL_SEEDS[left_idx], + "seed_b": MODEL_SEEDS[right_idx], + "signed_contribution_correlation": corr, + "top5_evidence_jaccard": top_jaccard, + "dominant_modality_agreement": dominant_left == dominant_right, + "seed_a_dominant_modality": ("text", "audio", "vision")[dominant_left], + "seed_b_dominant_modality": ("text", "audio", "vision")[dominant_right], + } + ) + + models = _load_ensemble(selected, MODEL_SEEDS, dims, device) + for index in range(min(5, len(cases))): + xs = tuple(torch.as_tensor(x[index : index + 1], dtype=torch.float32, device=device) for x in attachment.x) + mask = torch.as_tensor(attachment.mask[index : index + 1], dtype=torch.bool, device=device) + base_estimate = hierarchical_owen_one( + models, xs, mask, seed=20320000 + index, start_permutations=8, max_permutations=8 + ) + generator = torch.Generator(device=device).manual_seed(20330000 + index) + x_perturbed = [] + for modality, x in enumerate(xs): + noise = torch.randn(x.shape, dtype=x.dtype, device=device, generator=generator) * 0.01 + x_perturbed.append(x + noise * mask[..., modality, None].to(x.dtype)) + perturbed_estimate = hierarchical_owen_one( + models, + tuple(x_perturbed), + mask, + seed=20320000 + index, + start_permutations=8, + max_permutations=8, + ) + left = base_estimate["contribution"].reshape(-1) + right = perturbed_estimate["contribution"].reshape(-1) + corr = float(np.corrcoef(left, right)[0, 1]) if left.std() > 0 and right.std() > 0 else 0.0 + top_left = set(np.argsort(-np.abs(left))[:5]) + top_right = set(np.argsort(-np.abs(right))[:5]) + stability_rows.append( + { + "sample_id": cases[index]["case_id"], + "stability_source": "input_perturbation_1pct", + "seed_a": "base", + "seed_b": "gaussian_0.01", + "signed_contribution_correlation": corr, + "top5_evidence_jaccard": len(top_left & top_right) / max(1, len(top_left | top_right)), + "dominant_modality_agreement": np.abs(base_estimate["contribution"]).sum(axis=1).argmax() + == np.abs(perturbed_estimate["contribution"]).sum(axis=1).argmax(), + } + ) + return stability_rows + + +def _complexity_rows( + selected: str, + methods: dict[str, Sequence[nn.Module]], + sample_xs: tuple[torch.Tensor, torch.Tensor, torch.Tensor], + sample_mask: torch.Tensor, + owen_seconds: Sequence[float], + shapley_seconds_per_sample: float, + device: torch.device, +) -> list[dict[str, Any]]: + rows = [] + for name, models in methods.items(): + parameter_counts = [sum(parameter.numel() for parameter in model.parameters() if parameter.requires_grad) for model in models] + for model in models: + model.eval() + for _ in range(5): + with torch.inference_mode(): + ensemble_forward(models, sample_xs, sample_mask, details=(name == selected)) + if device.type == "cuda": + torch.cuda.synchronize() + times = [] + for _ in range(30): + begin = time.perf_counter() + with torch.inference_mode(): + ensemble_forward(models, sample_xs, sample_mask, details=(name == selected)) + if device.type == "cuda": + torch.cuda.synchronize() + times.append(time.perf_counter() - begin) + rows.append( + { + "method": name, + "trainable_parameters_per_seed": parameter_counts[0], + "ensemble_seed_count": len(models), + "ensemble_parameter_instances": int(sum(parameter_counts)), + "single_sample_inference_ms_mean": float(np.mean(times) * 1000.0), + "single_sample_inference_ms_p95": float(np.quantile(times, 0.95) * 1000.0), + "exact_8_coalition_shapley_seconds_per_sample": shapley_seconds_per_sample if name == selected else None, + "hierarchical_owen_seconds_per_sample_mean": float(np.mean(owen_seconds)) if name == selected and owen_seconds else None, + "hierarchical_owen_forward_evaluations_mean": None, + "device": torch.cuda.get_device_name(0) if device.type == "cuda" else str(device), + } + ) + return rows + + +def _plot_results(clean_rows: list[dict[str, Any]], local_rows: list[dict[str, Any]]) -> None: + import matplotlib + + matplotlib.use("Agg") + import matplotlib.pyplot as plt + + figure_dir = RESULTS_ROOT / "figures" + figure_dir.mkdir(parents=True, exist_ok=True) + display = [EARLYCONCAT, MOFE7_MLP, "A0", "A1", "A2", "A3"] + clean = {row["method"]: row for row in clean_rows} + fig, axes = plt.subplots(1, 2, figsize=(12, 4.6)) + for axis, metric, title in ((axes[0], "macro_f1", "Macro-F1 on locked validation"), (axes[1], "mae", "Intensity MAE on locked validation")): + means = [float(clean[method][f"{metric}_mean"]) for method in display if method in clean] + errors = [float(clean[method][f"{metric}_std"]) for method in display if method in clean] + labels = [method for method in display if method in clean] + axis.bar(np.arange(len(labels)), means, yerr=errors, capsize=3, color=["#4c78a8", "#f58518", "#54a24b", "#e45756", "#72b7b2", "#b279a2"][: len(labels)]) + axis.set_xticks(np.arange(len(labels)), labels, rotation=35, ha="right") + axis.set_title(title) + axis.grid(axis="y", alpha=0.25) + fig.tight_layout() + fig.savefig(figure_dir / "validation_metrics.png", dpi=180) + plt.close(fig) + if local_rows: + first = local_rows[0]["case_id"] + selected = [row for row in local_rows if row["case_id"] == first] + mat = np.zeros((3, 10), dtype=np.float64) + for row in selected: + mat[("text", "audio", "vision").index(row["modality"]), int(row["relative_bin"])] = float(row["local_owen_margin_contribution"]) + fig, axis = plt.subplots(figsize=(10, 3.5)) + bound = max(1e-8, float(np.quantile(np.abs(mat), 0.95))) + image = axis.imshow(mat, aspect="auto", cmap="coolwarm", vmin=-bound, vmax=bound) + axis.set_yticks(range(3), ("Text", "Audio", "Vision")) + axis.set_xlabel("Relative-progress bin (0–49; no physical seconds)") + axis.set_title(f"ATI–HO local Owen contribution: {first}") + fig.colorbar(image, ax=axis, label="Fixed logit-margin contribution") + fig.tight_layout() + fig.savefig(figure_dir / "attachment4_owen_example.png", dpi=180) + plt.close(fig) + + +def _write_reports( + selected: str, + selection_rows: list[dict[str, Any]], + clean_rows: list[dict[str, Any]], + bootstrap_rows: list[dict[str, Any]], + structural_rows: list[dict[str, Any]], + shapley_summary: dict[str, Any], + attachment_shapley_summary: dict[str, Any], + owen_rows: list[dict[str, Any]], + fidelity_rows: list[dict[str, Any]], + complexity_rows: list[dict[str, Any]], +) -> None: + clean = {row["method"]: row for row in clean_rows if row["scenario"] == "clean"} + baseline_table = [] + for method in (EARLYCONCAT, MOFE7_MLP, selected): + if method in clean: + row = clean[method] + baseline_table.append( + f"| {method} | {int(row['seeds'])} | {row['accuracy_mean']:.3f} ± {row['accuracy_std']:.3f} | " + f"{row['macro_f1_mean']:.3f} ± {row['macro_f1_std']:.3f} | " + f"{row['mae_mean']:.3f} ± {row['mae_std']:.3f} | {row['pearson_mean']:.3f} ± {row['pearson_std']:.3f} |" + ) + loss_lines = [ + f"| {row['method']} | {row['mean_validation_selection_loss']:.5f} ± {row['std_validation_selection_loss']:.5f} |" + for row in selection_rows + ] + delta_lines = [ + f"| {row['comparison']} | {row['metric']} | {row['delta_positive_favors_ATI_HO']:.4f} | " + f"[{row['bootstrap_ci_2p5']:.4f}, {row['bootstrap_ci_97p5']:.4f}] |" + for row in bootstrap_rows + ] + structural_summary = max( + (float(row.get("additive_reconstruction_max_abs", 0.0)) for row in structural_rows if row.get("method") == selected), default=0.0 + ) + local_conservation = max((abs(float(row["local_conservation_residual"])) for row in owen_rows), default=0.0) + stability_rows = _coerce_csv_numbers(_read_csv(RESULTS_ROOT / "stability_results.csv")) + stability_summary: dict[str, dict[str, float]] = {} + for source in ("training_seed", "input_perturbation_1pct"): + subset = [row for row in stability_rows if row.get("stability_source") == source] + if subset: + stability_summary[source] = { + "samples_or_pairs": float(len(subset)), + "mean_signed_correlation": float(np.mean([float(row["signed_contribution_correlation"]) for row in subset])), + "mean_top5_jaccard": float(np.mean([float(row["top5_evidence_jaccard"]) for row in subset])), + "dominant_modality_agreement": float(np.mean([str(row["dominant_modality_agreement"]).lower() == "true" for row in subset])), + } + fidelity_groups: dict[tuple[str, str], list[dict[str, Any]]] = defaultdict(list) + for row in fidelity_rows: + if row.get("split") == "attachment4_unlabelled" and abs(float(row.get("budget", -1)) - 0.3) < 1e-9: + fidelity_groups[(str(row["explanation"]), str(row["method"]))].append(row) + fidelity_lines = [] + for explanation in ("ATI_HO_Owen", "EarlyConcat_posthoc_Owen", "MoFE_posthoc_Owen", "MoFE_router_utility"): + for method_name in ("owen", "matched_random"): + subset = fidelity_groups.get((explanation, method_name), []) + if subset: + deletion = float(np.mean([float(row["deletion_margin_drop_mean"]) for row in subset])) + retention = float(np.mean([float(row["retention_margin_drop_mean"]) for row in subset])) + fidelity_lines.append(f"| {explanation} | {method_name} | {deletion:.3f} | {retention:.3f} |") + complexity_lines = [] + for row in complexity_rows: + shapley_time = row.get("exact_8_coalition_shapley_seconds_per_sample") + owen_time = row.get("hierarchical_owen_seconds_per_sample_mean") + shapley_text = f"{float(shapley_time):.3f}" if shapley_time not in (None, "") else "—" + owen_text = f"{float(owen_time):.3f}" if owen_time not in (None, "") else "—" + complexity_lines.append( + f"| {row['method']} | {int(row['trainable_parameters_per_seed'])} | " + f"{row['single_sample_inference_ms_mean']:.3f} | {shapley_text} | {owen_text} |" + ) + stable_owen = [row for row in owen_rows if row.get("stopping_status") == "stable"] + mean_permutations = float(np.mean([float(row["permutations"]) for row in owen_rows])) if owen_rows else 0.0 + fig_path = "figures/validation_metrics.png" + final_selection = json.loads((EXPERIMENT_ROOT / "final_selection.json").read_text(encoding="utf-8")) + result_doc = f"""# ATI–HO Q3 实验结果 + +## 选型与数据 + +最终模型按锁定验证集四场景任务损失的三 seed 均值选择为 **{selected}**。Stage I seed 42 初选模型为 `{final_selection['provisional_seed42_method']}`;Stage II 对初选模型与两项关键消融统一使用 seed 42、3407、2026。Attachment 4 标签未参与训练、选型或指标计算。 + +输入来自官方 `unaligned_50.pkl`,使用统一 Q1 Relative-Progress adapter 投影到 50 个归一化进度槽,维度为 Text 768、Audio 74、Vision 35。训练/验证/测试分别为 3,395/728/727 条,视频组数为 1,528/239/381,组间重叠为 0。缩放器只在训练组拟合,与既有 Q2 scaler 最大绝对差异为 0。 + +## 验证集性能 + +| 方法 | Seeds | Accuracy | Macro-F1 | MAE | Pearson | +|---|---:|---:|---:|---:|---:| +{chr(10).join(baseline_table)} + +数值为 seed 均值 ± 标准差。主要模型采用三分类指标和连续强度指标;附件4只有预测和解释输出,不报告无标签样本的准确率。 + +### ATI 消融选型 + +| ATI 方案 | 固定场景验证损失(均值 ± 标准差) | +|---|---:| +{chr(10).join(loss_lines)} + +![验证集性能比较]({fig_path}) + +## 配对视频组 Bootstrap + +正值表示 ATI–HO 更好;MAE 的差值定义为基线 MAE 减 ATI–HO MAE。区间以来源视频为重采样单位,1,000 次。 + +| 比较 | 指标 | 差值 | 95% CI | +|---|---|---:|---:| +{chr(10).join(delta_lines)} + +Bootstrap 在三 seed 集成预测上计算,表格中的性能均值则是逐 seed 指标的均值。指标是非线性的,两处点估计不要求完全相等。 + +## 结构与归因审计 + +- A0–A3 主效应、加和重构与锚定检查通过;D0 未锚定诊断检出缺失模态泄漏。 +- 最终模型训练后最大加和重构残差:`{structural_summary:.3g}`。 +- 最终模型在 {shapley_summary['samples']} 条锁定验证样本上的解析 Shapley 与 8 联盟枚举通过率:{shapley_summary['pass_rate']:.3%};最大绝对误差 `{shapley_summary['max_abs_error']:.3g}`,容差为绝对 1e-6 加相对 1e-5。 +- Attachment 4 的解析/精确分类 Shapley 通过率:{attachment_shapley_summary['analytic_class_shapley_pass_rate']:.3%}(n={attachment_shapley_summary['samples']})。最终强度输出经类别选择与 sigmoid 解码,使用 8 联盟精确 Shapley;不把强度贡献称为线性参数分解。 +- Attachment 4 Hierarchical Owen 局部守恒最大残差:`{local_conservation:.3g}`。每个样本按模态外层排列、模态内 10 个相对进度片段排列,Rπ 从 8 起并在稳定时停止,最多 64。 + +## Fidelity 诊断 + +删除/保留测试分别使用每模态相同片段数、同一 10/20/30% 预算,并与同模态随机片段对照。Attachment 4 没有标签,因此仅报告固定 logit margin 对输入遮挡的响应,不称为解释准确率或因果效应。验证集另取固定的类别分层子集,用于同一模型忠实性诊断;结果见 `fidelity_results.csv`。 + +Attachment 4 的 30% 删除比较(20 个无标签样本均值)如下。数值越大表示遮掉所选片段后固定类别 margin 降得越多;“完整−保留 margin”是带符号差值,负值表示只保留高分片段时 margin 高于完整输入。 + +| 解释来源 | 片段排序 | 删除 margin 降幅 | 完整−保留 margin | +|---|---|---:|---:| +{chr(10).join(fidelity_lines)} + +ATI–HO Owen 排序在 30% 删除下的 margin 降幅为 0.433,匹配随机片段为 0.112。该差异反映这批无标签样本上的模型遮挡响应,不是解释正确率。 + +## 稳定性与成本 + +Attachment 4 Owen 归因有 {len(stable_owen)}/{len(owen_rows)} 个样本在最多 64 次以内达到预设稳定条件,平均使用 {mean_permutations:.1f} 次排列。训练 seed 归因的平均 signed correlation 为 {stability_summary.get('training_seed', {}).get('mean_signed_correlation', 0.0):.3f}、top-5 Jaccard 为 {stability_summary.get('training_seed', {}).get('mean_top5_jaccard', 0.0):.3f};因此细粒度位置归因对训练 seed 的一致性有限。1% 特征扰动诊断单独列于 `stability_results.csv`。 + +| 模型 | 每 seed 可训练参数 | 三 seed 集成单样本延迟(ms) | 精确 8 联盟 Shapley(秒/样本) | Owen(秒/样本) | +|---|---:|---:|---:|---:| +{chr(10).join(complexity_lines)} + +时延在本轮 RTX 5070 Ti 上测得,包含三 seed 集成前向;只作本机参考。 + +## 限制 + +输入按归一化进度排序;没有可靠的逐词或逐帧物理时间戳。局部片段索引不得解释成秒数。模型归因描述当前模型对输入遮挡的响应,不证明人类情绪的因果机制。当前最终 A0 只保留锚定主效应;验证结果未支持保留更复杂的 pairwise 结构。A1/A2 结果作为消融保留。 + +## 复现文件 + +训练和评估代码位于 `q3/ati_ho/`,模型定义位于 `model/ati_ho.py` 与 `model/ati_ho_config.py`;权重、训练记录和 CSV 审计位于 `experiments/q3/ati_ho/`;附件4预测与解释交付件位于 `output/q3/ati_ho/`。运行方式见 `q3/ati_ho/README.md`。 +""" + (RESULTS_ROOT / "ATI_HO_RESULTS.md").write_text(result_doc, encoding="utf-8") + + paper = f"""# ATI–HO:基于锚定时间交互与分层 Owen 归因的多模态情感预测 + +## 摘要 + +本文在复杂场景多模态情感识别的第三问中实现 ATI–HO,并以官方未对齐输入和统一 Q1 adapter 为基础训练。实验包含 EarlyConcat + BiGRU、MoFE-7 + MLP Router,以及 ATI 主效应、低秩 pairwise、锚定 cross-attention 和可见性掩码辅助消融。ATI–HO 的最终方案由锁定验证集选择为 **{selected}**,三 seed 固定场景验证损失均值最小。最终模型在 {shapley_summary['samples']} 条验证样本上的解析 Shapley 与 8 联盟精确枚举通过率为 {shapley_summary['pass_rate']:.3%}。这里的结果支持“输出参数存在可核验的加和分解”,不构成对情绪因果机制的证明。 + +## 1. 问题与方法 + +给定 Text、Audio、Vision 三路 50 步相对进度序列及逐步可见掩码,预测 negative/neutral/positive 类别与 [-3,3] 强度。每个模态使用私有投影、双向 GRU(每方向 32 隐单元)和注意力池化。主效应以显式空输入前向相减锚定为零。ATI 参数向量为三个居中类别 logit 与负/正条件强度参数: + +`ξ = b + Σ_m G_m + Σ_{{m list[dict[str, Any]]: + converted: list[dict[str, Any]] = [] + for row in rows: + item: dict[str, Any] = {} + for key, value in row.items(): + try: + item[key] = float(value) + except (TypeError, ValueError): + item[key] = value + converted.append(item) + return converted + + +def finalize_reports() -> None: + """Rebuild reports and the run manifest from already generated audit tables.""" + selection = json.loads((EXPERIMENT_ROOT / "final_selection.json").read_text(encoding="utf-8")) + selected = selection["selected_method"] + selection_rows = _coerce_csv_numbers(_read_csv(RESULTS_ROOT / "selection_results.csv")) + clean_rows = _coerce_csv_numbers(_read_csv(RESULTS_ROOT / "main_results.csv")) + bootstrap_rows = _coerce_csv_numbers(_read_csv(RESULTS_ROOT / "bootstrap_results.csv")) + structural_rows = _coerce_csv_numbers(_read_csv(RESULTS_ROOT / "structural_audit.csv")) + owen_rows = _coerce_csv_numbers(_read_csv(RESULTS_ROOT / "owen_audit.csv")) + fidelity_rows = _coerce_csv_numbers(_read_csv(RESULTS_ROOT / "fidelity_results.csv")) + complexity_rows = _coerce_csv_numbers(_read_csv(RESULTS_ROOT / "complexity.csv")) + local_rows = _coerce_csv_numbers(_read_csv(RESULTS_ROOT / "attachment4_local_evidence.csv")) + if not all((selection_rows, clean_rows, bootstrap_rows, structural_rows, owen_rows, fidelity_rows, complexity_rows)): + raise FileNotFoundError("evaluation tables are incomplete; run the full ATI–HO evaluation first") + shapley_summary = json.loads((RESULTS_ROOT / "shapley_audit_summary.json").read_text(encoding="utf-8")) + prediction_manifest = json.loads( + (RESULTS_ROOT / "attachment4_prediction_manifest.json").read_text(encoding="utf-8") + ) + attachment_summary = prediction_manifest["shapley_audit"] + _plot_results(clean_rows, local_rows) + _write_reports( + selected, + selection_rows, + clean_rows, + bootstrap_rows, + structural_rows, + shapley_summary, + attachment_summary, + owen_rows, + fidelity_rows, + complexity_rows, + ) + training_manifest = json.loads((EXPERIMENT_ROOT / "run_manifest.json").read_text(encoding="utf-8")) + device_name = next((row.get("device") for row in complexity_rows if row.get("method") == selected), "unknown") + _write_json( + RESULTS_ROOT / "run_manifest.json", + { + "selected_method": selected, + "candidate_selection": selection, + "device": device_name, + "torch_version": torch.__version__, + "cuda_version": torch.version.cuda, + "attachment4_cases": prediction_manifest["prediction_rows"], + "validation_samples": shapley_summary["samples"], + "validation_group_bootstrap_replicates": BOOTSTRAP_REPLICATES, + "adapter_and_scaler_metadata": training_manifest.get("data"), + "attachment4_source_hashes": prediction_manifest["attachment4_source_sha256"], + "labels_used_from_attachment4": False, + }, + ) + print(f"Q3 reports finalized from existing evaluation tables: selected={selected}", flush=True) + + +def run(device: torch.device) -> None: + RESULTS_ROOT.mkdir(parents=True, exist_ok=True) + train, valid, stats, data_meta = load_training_data() + dims = tuple(int(x.shape[-1]) for x in train.x) + selected, clean_rows, selection_rows = _write_seed_and_summary_tables() + candidate_methods = [row["method"] for row in selection_rows] + selected_models = _load_ensemble(selected, MODEL_SEEDS, dims, device) + early_models = _load_ensemble(EARLYCONCAT, MODEL_SEEDS, dims, device) + mofe_models = _load_ensemble(MOFE7_MLP, MODEL_SEEDS, dims, device) + + predictions_by_method = { + selected: _predict_ensemble(selected_models, valid, valid.mask, device), + EARLYCONCAT: _predict_ensemble(early_models, valid, valid.mask, device), + MOFE7_MLP: _predict_ensemble(mofe_models, valid, valid.mask, device), + } + bootstrap_rows = _cluster_bootstrap(valid, predictions_by_method, selected) + _save_csv(RESULTS_ROOT / "bootstrap_results.csv", bootstrap_rows) + + smoke_xs = tuple(torch.as_tensor(x[:64], dtype=torch.float32, device=device) for x in valid.x) + smoke_mask = torch.as_tensor(valid.mask[:64], dtype=torch.bool, device=device) + structural_rows = [] + for seed, model in zip(MODEL_SEEDS, selected_models): + report = structural_audit(model, smoke_xs, smoke_mask) + structural_rows.append({"method": selected, "seed": seed, "samples": min(64, valid.n), **{ + key: json.dumps(value, sort_keys=True) if isinstance(value, dict) else value for key, value in report.items() + }}) + if not report["checks_pass"]: + raise RuntimeError(f"final ATI–HO structural audit failed for seed {seed}: {report}") + stage1_rows = _read_csv(EXPERIMENT_ROOT / "structural_audit.csv") + structural_rows.extend(stage1_rows) + _save_csv(RESULTS_ROOT / "structural_audit.csv", structural_rows) + + validation_predictions, shapley_summary, validation_shapley_rows = _validation_predictions_and_shapley( + selected, selected_models, valid, device + ) + _save_csv(RESULTS_ROOT / "shapley_audit.csv", validation_shapley_rows) + _write_json(RESULTS_ROOT / "shapley_audit_summary.json", shapley_summary) + + cases, attachment_meta = _read_attachment4("unaligned_50") + attachment = _attachment_split(cases, stats) + attach_predictions, attach_explanations, attach_shapley_summary, attach_exact = _attachment_predictions_and_explanations( + selected, selected_models, cases, attachment, device + ) + local_rows, owen_rows, fidelity_rows, comparison_rows, stability_rows = _owen_and_fidelity( + selected_models, early_models, mofe_models, cases, attachment, valid, device + ) + _save_csv(RESULTS_ROOT / "owen_audit.csv", owen_rows) + _save_csv(RESULTS_ROOT / "attachment4_local_evidence.csv", local_rows) + _save_csv(RESULTS_ROOT / "attachment4_explanations.csv", attach_explanations) + _save_csv(RESULTS_ROOT / "attachment4_predictions.csv", attach_predictions) + _save_csv(RESULTS_ROOT / "fidelity_results.csv", fidelity_rows) + _save_csv(RESULTS_ROOT / "attachment4_comparison_owen.csv", comparison_rows) + _save_csv(RESULTS_ROOT / "stability_results.csv", stability_rows) + + # Re-evaluate time cost on one complete Attachment 4 example; Owen times were captured above. + first_xs = tuple(torch.as_tensor(x[:1], dtype=torch.float32, device=device) for x in attachment.x) + first_mask = torch.as_tensor(attachment.mask[:1], dtype=torch.bool, device=device) + shapley_start = time.perf_counter() + exact_shapley_audit(selected_models, first_xs, first_mask, batch_size=8) + shapley_per_sample = time.perf_counter() - shapley_start + complexity_models = {selected: selected_models, EARLYCONCAT: early_models, MOFE7_MLP: mofe_models} + complexity_rows = _complexity_rows( + selected, + complexity_models, + first_xs, + first_mask, + [float(row.get("owen_seconds", 0.0)) for row in owen_rows if "owen_seconds" in row], + shapley_per_sample, + device, + ) + # Owen time for each sample is also tracked by the detailed run table. + if owen_rows: + average_owen = float(np.mean([float(row.get("elapsed_seconds", 0.0)) for row in owen_rows])) + for row in complexity_rows: + if row["method"] == selected: + row["hierarchical_owen_seconds_per_sample_mean"] = average_owen + row["hierarchical_owen_forward_evaluations_mean"] = float( + np.mean([2 + 30 * int(item["permutations"]) for item in owen_rows]) + ) + _save_csv(RESULTS_ROOT / "complexity.csv", complexity_rows) + + # Input hashes and model provenance. Attachment 4 has no target values. + model_hashes = {f"{method}/seed_{seed}": _sha256(_checkpoint_path(method, seed)) for method in (selected, EARLYCONCAT, MOFE7_MLP) for seed in MODEL_SEEDS} + source_hashes = {case["source_file"].name: case["source_sha256"] for case in cases} + prediction_manifest = { + "task": "Q3 ATI–HO Attachment 4 final inference", + "selected_method": selected, + "seeds": list(MODEL_SEEDS), + "input_version": "official unaligned_50 Attachment 4", + "adapter": "Q1AlignmentAdapter; Relative-Progress; target_steps=50", + "physical_time_alignment": False, + "no_attachment4_labels_or_metrics_used": True, + "scaler": str(SCALER_PATH), + "feature_dimensions": [768, 74, 35], + "class_order": list(CLASS_NAMES), + "intensity_decode": "predicted negative: -3*sigmoid(r_negative); neutral: 0; predicted positive: 3*sigmoid(r_positive)", + "model_checkpoint_sha256": model_hashes, + "attachment4_source_sha256": source_hashes, + "attachment4_feature_dir": attachment_meta.get("version_dir"), + "prediction_rows": len(attach_predictions), + "explanation_rows": len(attach_explanations), + "local_evidence_rows": len(local_rows), + "shapley_audit": attach_shapley_summary, + } + _write_json(RESULTS_ROOT / "attachment4_prediction_manifest.json", prediction_manifest) + + # Chapter IV deliverables: prediction and explanation tables plus a concise provenance note. + SUBMIT_OUTPUT.mkdir(parents=True, exist_ok=True) + _save_csv(SUBMIT_OUTPUT / "attachment4_predictions.csv", attach_predictions) + _save_csv(SUBMIT_OUTPUT / "attachment4_explanations.csv", attach_explanations) + _save_csv(SUBMIT_OUTPUT / "attachment4_local_evidence.csv", local_rows) + _write_json(SUBMIT_OUTPUT / "attachment4_prediction_manifest.json", prediction_manifest) + readme = """# Q3 ATI–HO 提交输出 + +| 文件 | 内容 | +|---|---| +| `attachment4_predictions.csv` | 官方附件4的 20 条预测类别、强度与类别概率 | +| `attachment4_explanations.csv` | 主效应、pairwise 项、解析/精确分类 Shapley 与强度精确 Shapley | +| `attachment4_local_evidence.csv` | 按模态分组的局部 Hierarchical Owen 片段贡献、标准误与相对进度位置 | +| `attachment4_prediction_manifest.json` | adapter、模型权重哈希、文件计数和无标签推理审计 | + +附件4没有标签,本目录不提供准确率或误差指标。所有位置均为归一化进度槽,不是秒数。 +""" + (SUBMIT_OUTPUT / "README.md").write_text(readme, encoding="utf-8") + + _plot_results(clean_rows, local_rows) + all_structural = [_read_csv(RESULTS_ROOT / "structural_audit.csv")] + _write_reports( + selected, + selection_rows, + clean_rows, + bootstrap_rows, + all_structural[0], + shapley_summary, + attach_shapley_summary, + owen_rows, + fidelity_rows, + complexity_rows, + ) + _write_json( + RESULTS_ROOT / "run_manifest.json", + { + "selected_method": selected, + "candidate_selection": json.loads((EXPERIMENT_ROOT / "final_selection.json").read_text(encoding="utf-8")), + "device": str(device), + "torch_version": torch.__version__, + "cuda_version": torch.version.cuda, + "gpu": torch.cuda.get_device_name(0) if device.type == "cuda" else None, + "validation_samples": valid.n, + "validation_group_bootstrap_replicates": BOOTSTRAP_REPLICATES, + "attachment4_cases": len(cases), + "adapter_and_scaler_metadata": data_meta, + "attachment4_shapley_summary": attach_shapley_summary, + "attachment4_source_hashes": source_hashes, + "labels_used_from_attachment4": False, + }, + ) + print( + f"Q3 evaluation complete: selected={selected}; validation Macro-F1=" + f"{next(row['macro_f1_mean'] for row in clean_rows if row['method'] == selected and row['scenario'] == 'clean'):.4f}; " + f"attachment4_cases={len(cases)}; Owen conservation max=" + f"{max((abs(float(row['local_conservation_residual'])) for row in owen_rows), default=0.0):.3g}", + flush=True, + ) + + +def main() -> None: + parser = argparse.ArgumentParser(description="Audit ATI–HO validation results and create Attachment 4 outputs.") + parser.add_argument("--device", default="auto") + parser.add_argument("--reports-only", action="store_true", help="rebuild reports after an interrupted final report write") + parser.add_argument("--stability-only", action="store_true", help="recompute and save seed/input attribution stability") + args = parser.parse_args() + if args.reports_only: + finalize_reports() + return + if args.device == "auto": + device = torch.device("cuda" if torch.cuda.is_available() else "cpu") + else: + device = torch.device(args.device) + if args.stability_only: + selected = json.loads((EXPERIMENT_ROOT / "final_selection.json").read_text(encoding="utf-8"))["selected_method"] + cases, _meta = _read_attachment4("unaligned_50") + stats = RobustStats.load(SCALER_PATH) + attachment = _attachment_split(cases, stats) + rows = _stability_diagnostics(selected, cases, attachment, device) + _save_csv(RESULTS_ROOT / "stability_results.csv", rows) + print(f"Q3 stability diagnostics complete: rows={len(rows)}", flush=True) + return + run(device) + + +if __name__ == "__main__": + main() diff --git a/final/q3/ati_ho/owen.py b/final/q3/ati_ho/owen.py new file mode 100644 index 0000000..d5ec994 --- /dev/null +++ b/final/q3/ati_ho/owen.py @@ -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 diff --git a/final/q3/ati_ho/self_test.py b/final/q3/ati_ho/self_test.py new file mode 100644 index 0000000..8a73fa7 --- /dev/null +++ b/final/q3/ati_ho/self_test.py @@ -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() diff --git a/final/q3/ati_ho/train.py b/final/q3/ati_ho/train.py new file mode 100644 index 0000000..b1cb472 --- /dev/null +++ b/final/q3/ati_ho/train.py @@ -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()