Flatten submit package structure
This commit is contained in:
@@ -0,0 +1 @@
|
||||
"""Q3 interpretable emotion-recognition pipeline."""
|
||||
@@ -0,0 +1,46 @@
|
||||
# ATI–HO 训练与评估
|
||||
|
||||
ATI–HO 是当前 Q3 方案。模型定义位于 `model/ati_ho.py` 和 `model/ati_ho_config.py`。
|
||||
|
||||
## 输入准备
|
||||
|
||||
设置 `FINAL_DATA_DIR` 指向官方数据根目录。完整训练和评估需要附件 2 的 `unaligned_50.pkl`、附件 4 的未对齐特征文件和训练集 robust scaler:
|
||||
|
||||
`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 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 轮,固定验证情景任务损失早停。训练记录与检查点保存在 `experiments/q3/ati_ho/`。只有明确要覆盖检查点时才加 `--force`。
|
||||
|
||||
## 评估命令
|
||||
|
||||
```bash
|
||||
export FINAL_DATA_DIR="/path/to/E题数据"
|
||||
python -m q3.ati_ho.evaluate --device auto
|
||||
```
|
||||
|
||||
评估会比较 ATI 消融、重算官方验证指标和按来源视频组 Bootstrap、审计解析/精确 Shapley、运行附件 4 局部 Owen 和删除/保留诊断,并生成 Q3 输出。最终 ATI 方案按三 seed、四个固定验证情景的平均任务损失选出。附件 4 标签不会读取或用于报告;附件 4 输出没有准确率。
|
||||
|
||||
若完整评估已写完 CSV,但报告阶段中断,可运行:
|
||||
|
||||
```bash
|
||||
python -m q3.ati_ho.evaluate --reports-only
|
||||
```
|
||||
|
||||
若只需补算训练 seed 与 1% 输入扰动下的 attribution 稳定性:
|
||||
|
||||
```bash
|
||||
export FINAL_DATA_DIR="/path/to/E题数据"
|
||||
python -m q3.ati_ho.evaluate --device auto --stability-only
|
||||
```
|
||||
|
||||
题目交付文件写入 `output/q3/ati_ho/`。完整结果表、审计和论文式记录位于 `experiments/q3/ati_ho/results/ati_ho/`。验证指标解释和结果边界见 `REPORTS.md`。
|
||||
@@ -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"]
|
||||
@@ -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,
|
||||
}
|
||||
@@ -0,0 +1,68 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
|
||||
PAIR_MODES = ((0, 1), (0, 2), (1, 2))
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def structural_audit(
|
||||
model: nn.Module,
|
||||
xs: tuple[torch.Tensor, torch.Tensor, torch.Tensor],
|
||||
masks: torch.Tensor,
|
||||
*,
|
||||
atol: float = 1e-6,
|
||||
) -> dict[str, Any]:
|
||||
model.eval()
|
||||
output = model(xs, masks, return_details=True)
|
||||
reconstructed = output["baseline"] + output["main_effects"].sum(dim=1) + output["pair_effects"].sum(dim=1)
|
||||
additive_residual = torch.max(torch.abs(reconstructed - output["params"])).item()
|
||||
|
||||
absent_xs = tuple(torch.zeros_like(x) for x in xs)
|
||||
absent_mask = torch.zeros_like(masks, dtype=torch.bool)
|
||||
absent = model(absent_xs, absent_mask, return_details=True)
|
||||
main_zero_residual = torch.max(torch.abs(absent["main_effects"])).item()
|
||||
full_zero_residual = torch.max(torch.abs(absent["params"] - absent["baseline"])).item()
|
||||
pair_zero_residual = torch.max(torch.abs(absent["pair_effects"])).item()
|
||||
|
||||
pair_anchor_residuals: dict[str, float] = {}
|
||||
for pair_index, (left, right) in enumerate(PAIR_MODES):
|
||||
maxima = []
|
||||
for hidden in (left, right):
|
||||
altered = masks.clone()
|
||||
altered[..., hidden] = False
|
||||
out = model(xs, altered, return_details=True)
|
||||
maxima.append(torch.max(torch.abs(out["pair_effects"][:, pair_index])).item())
|
||||
pair_anchor_residuals[f"{left}{right}"] = max(maxima)
|
||||
|
||||
finite_count = 0
|
||||
nonfinite_count = 0
|
||||
for key in ("params", "logits", "intensity", "main_effects", "pair_effects"):
|
||||
tensor = output[key]
|
||||
finite_count += int(torch.isfinite(tensor).sum().item())
|
||||
nonfinite_count += int((~torch.isfinite(tensor)).sum().item())
|
||||
|
||||
anchored = bool(getattr(getattr(model, "config", None), "anchored", False))
|
||||
main_additive_pass = max(additive_residual, main_zero_residual) <= atol
|
||||
pair_pass = max(pair_anchor_residuals.values(), default=0.0) <= atol
|
||||
baseline_pass = full_zero_residual <= atol and pair_zero_residual <= atol
|
||||
checks_pass = main_additive_pass and nonfinite_count == 0 and (baseline_pass if anchored else True)
|
||||
return {
|
||||
"anchored": anchored,
|
||||
"main_effect_zero_anchor_max_abs": main_zero_residual,
|
||||
"pair_effect_zero_anchor_max_abs": pair_zero_residual,
|
||||
"pair_single_missing_anchor_max_abs": pair_anchor_residuals,
|
||||
"additive_reconstruction_max_abs": additive_residual,
|
||||
"all_modalities_missing_equals_baseline_max_abs": full_zero_residual,
|
||||
"finite_value_count": finite_count,
|
||||
"nonfinite_value_count": nonfinite_count,
|
||||
"main_and_additivity_pass": main_additive_pass,
|
||||
"full_baseline_anchor_pass": baseline_pass if anchored else None,
|
||||
"pair_anchor_pass": pair_pass if anchored else None,
|
||||
"unanchored_control_detected_leakage": ((not pair_pass) or not baseline_pass) if not anchored else False,
|
||||
"checks_pass": checks_pass,
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,240 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Sequence
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from .attribution import ensemble_forward
|
||||
|
||||
|
||||
def _margin_values(
|
||||
models: Sequence[nn.Module],
|
||||
xs: tuple[torch.Tensor, torch.Tensor, torch.Tensor],
|
||||
masks: np.ndarray,
|
||||
target: int,
|
||||
other: int,
|
||||
*,
|
||||
batch_size: int = 64,
|
||||
) -> np.ndarray:
|
||||
device = xs[0].device
|
||||
values: list[np.ndarray] = []
|
||||
with torch.inference_mode():
|
||||
for start in range(0, len(masks), batch_size):
|
||||
end = min(len(masks), start + batch_size)
|
||||
current_mask = torch.as_tensor(masks[start:end], dtype=torch.bool, device=device)
|
||||
repeated = tuple(x.expand(end - start, -1, -1).contiguous() for x in xs)
|
||||
output = ensemble_forward(models, repeated, current_mask, details=False)
|
||||
margin = output["logits"][:, target] - output["logits"][:, other]
|
||||
values.append(margin.detach().cpu().numpy().astype(np.float64))
|
||||
return np.concatenate(values) if values else np.empty(0, dtype=np.float64)
|
||||
|
||||
|
||||
def _top_stability(previous: np.ndarray, current: np.ndarray, k: int = 5) -> float:
|
||||
old = set(np.argsort(-np.abs(previous), kind="stable")[:k].tolist())
|
||||
new = set(np.argsort(-np.abs(current), kind="stable")[:k].tolist())
|
||||
return float(len(old & new) / max(1, len(old | new)))
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def hierarchical_owen_one(
|
||||
models: Sequence[nn.Module],
|
||||
xs: tuple[torch.Tensor, torch.Tensor, torch.Tensor],
|
||||
mask: torch.Tensor,
|
||||
*,
|
||||
seed: int,
|
||||
bins_per_modality: int = 10,
|
||||
start_permutations: int = 8,
|
||||
max_permutations: int = 64,
|
||||
batch_size: int = 64,
|
||||
) -> dict[str, Any]:
|
||||
"""Estimate a three-group Owen allocation over relative-progress segments.
|
||||
|
||||
The outer permutation orders modalities. Each modality's 10 temporal bins
|
||||
are then added in a random inner permutation. This is a sampled Owen value,
|
||||
not a perturbation of arbitrary individual feature dimensions.
|
||||
"""
|
||||
model_output = ensemble_forward(models, xs, mask, details=False)
|
||||
logits = model_output["logits"][0]
|
||||
target = int(logits.argmax().item())
|
||||
other = int(logits.argsort(descending=True)[1].item())
|
||||
full_margin = float((logits[target] - logits[other]).item())
|
||||
no_features = torch.zeros_like(mask, dtype=torch.bool)
|
||||
baseline = ensemble_forward(models, xs, no_features, details=False)["logits"][0]
|
||||
base_margin = float((baseline[target] - baseline[other]).item())
|
||||
|
||||
base_mask = mask[0].detach().cpu().numpy().astype(bool)
|
||||
steps = base_mask.shape[0]
|
||||
boundaries = np.linspace(0, steps, bins_per_modality + 1).round().astype(int)
|
||||
slices = [(int(boundaries[k]), int(boundaries[k + 1])) for k in range(bins_per_modality)]
|
||||
rng = np.random.default_rng(seed)
|
||||
draws: list[np.ndarray] = []
|
||||
previous_mean: np.ndarray | None = None
|
||||
stability = float("nan")
|
||||
schedule = [start_permutations]
|
||||
while schedule[-1] < max_permutations:
|
||||
schedule.append(min(max_permutations, schedule[-1] * 2))
|
||||
next_target = schedule[0]
|
||||
stopping_status = "max_permutations"
|
||||
|
||||
while len(draws) < max_permutations:
|
||||
needed = min(8, max_permutations - len(draws))
|
||||
all_masks: list[np.ndarray] = []
|
||||
player_order: list[list[tuple[int, int]]] = []
|
||||
for _ in range(needed):
|
||||
current = np.zeros_like(base_mask, dtype=bool)
|
||||
order: list[tuple[int, int]] = []
|
||||
outer = rng.permutation(3)
|
||||
for modality in outer:
|
||||
for bin_index in rng.permutation(bins_per_modality):
|
||||
left, right = slices[int(bin_index)]
|
||||
current[left:right, int(modality)] = base_mask[left:right, int(modality)]
|
||||
order.append((int(modality), int(bin_index)))
|
||||
all_masks.append(current.copy())
|
||||
player_order.append(order)
|
||||
|
||||
scores = _margin_values(
|
||||
models,
|
||||
xs,
|
||||
np.stack(all_masks),
|
||||
target,
|
||||
other,
|
||||
batch_size=batch_size,
|
||||
)
|
||||
cursor = 0
|
||||
for order in player_order:
|
||||
previous_score = base_margin
|
||||
draw = np.zeros((3, bins_per_modality), dtype=np.float64)
|
||||
for modality, bin_index in order:
|
||||
current_score = float(scores[cursor])
|
||||
cursor += 1
|
||||
draw[modality, bin_index] = current_score - previous_score
|
||||
previous_score = current_score
|
||||
draws.append(draw)
|
||||
|
||||
if len(draws) >= next_target:
|
||||
current_mean = np.mean(np.stack(draws), axis=0)
|
||||
flat_mean = current_mean.reshape(-1)
|
||||
standard_error = np.std(np.stack(draws), axis=0, ddof=1) / np.sqrt(len(draws))
|
||||
se_mean = float(np.mean(standard_error))
|
||||
if previous_mean is not None:
|
||||
stability = _top_stability(previous_mean, flat_mean)
|
||||
se_limit = max(0.02, 0.10 * abs(full_margin - base_margin))
|
||||
if stability >= 0.8 and se_mean <= se_limit:
|
||||
stopping_status = "stable"
|
||||
break
|
||||
previous_mean = flat_mean
|
||||
next_idx = next((i for i, value in enumerate(schedule) if value > len(draws)), None)
|
||||
if next_idx is None:
|
||||
break
|
||||
next_target = schedule[next_idx]
|
||||
|
||||
draw_array = np.stack(draws)
|
||||
mean = draw_array.mean(axis=0)
|
||||
standard_error = draw_array.std(axis=0, ddof=1) / np.sqrt(len(draws)) if len(draws) > 1 else np.full_like(mean, np.nan)
|
||||
conservation = float(mean.sum() - (full_margin - base_margin))
|
||||
return {
|
||||
"target_class": target,
|
||||
"runner_up_class": other,
|
||||
"full_margin": full_margin,
|
||||
"baseline_margin": base_margin,
|
||||
"contribution": mean,
|
||||
"standard_error": standard_error,
|
||||
"permutations": len(draws),
|
||||
"stopping_status": stopping_status,
|
||||
"top5_jaccard_last_check": stability,
|
||||
"local_conservation_residual": conservation,
|
||||
"bin_slices": slices,
|
||||
}
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def fidelity_audit_one(
|
||||
models: Sequence[nn.Module],
|
||||
xs: tuple[torch.Tensor, torch.Tensor, torch.Tensor],
|
||||
mask: torch.Tensor,
|
||||
contribution: np.ndarray,
|
||||
*,
|
||||
sample_id: str,
|
||||
seed: int,
|
||||
bins_per_modality: int = 10,
|
||||
random_replicates: int = 20,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Deletion/retention at 10/20/30%, matched by modality and segment count."""
|
||||
base_mask = mask[0].detach().cpu().numpy().astype(bool)
|
||||
steps = base_mask.shape[0]
|
||||
boundaries = np.linspace(0, steps, bins_per_modality + 1).round().astype(int)
|
||||
slices = [(int(boundaries[k]), int(boundaries[k + 1])) for k in range(bins_per_modality)]
|
||||
full = ensemble_forward(models, xs, mask, details=False)
|
||||
ranked = full["logits"][0].argsort(descending=True)
|
||||
target, other = int(ranked[0].item()), int(ranked[1].item())
|
||||
full_margin = float((full["logits"][0, target] - full["logits"][0, other]).item())
|
||||
rng = np.random.default_rng(seed)
|
||||
candidates = [
|
||||
[index for index, (left, right) in enumerate(slices) if base_mask[left:right, m].any()]
|
||||
for m in range(3)
|
||||
]
|
||||
masks_to_score: list[np.ndarray] = []
|
||||
row_specs: list[tuple[float, str, int]] = []
|
||||
|
||||
for rate in (0.1, 0.2, 0.3):
|
||||
counts = [max(1, int(np.ceil(rate * len(indices)))) if indices else 0 for indices in candidates]
|
||||
method_choices: list[tuple[str, list[list[int]]]] = []
|
||||
top_choices: list[list[int]] = []
|
||||
random_choices: list[list[int]] = []
|
||||
for modality in range(3):
|
||||
available = candidates[modality]
|
||||
count = min(counts[modality], len(available))
|
||||
score_order = sorted(available, key=lambda index: (-abs(contribution[modality, index]), index))
|
||||
top_choices.append(score_order[:count])
|
||||
random_choices.append(list(rng.choice(available, size=count, replace=False)) if count else [])
|
||||
method_choices.append(("owen", top_choices))
|
||||
method_choices.append(("matched_random", random_choices))
|
||||
|
||||
for method_name, choices in method_choices:
|
||||
reps = 1 if method_name == "owen" else random_replicates
|
||||
for rep in range(reps):
|
||||
if method_name == "matched_random" and rep > 0:
|
||||
choices = [
|
||||
list(rng.choice(candidates[m], size=min(counts[m], len(candidates[m])), replace=False))
|
||||
if counts[m]
|
||||
else []
|
||||
for m in range(3)
|
||||
]
|
||||
deleted = base_mask.copy()
|
||||
retained = np.zeros_like(base_mask, dtype=bool)
|
||||
for modality, selected in enumerate(choices):
|
||||
for bin_index in selected:
|
||||
left, right = slices[bin_index]
|
||||
deleted[left:right, modality] = False
|
||||
retained[left:right, modality] = base_mask[left:right, modality]
|
||||
masks_to_score.extend((deleted, retained))
|
||||
row_specs.extend(((rate, method_name, rep), (rate, method_name, rep)))
|
||||
|
||||
score_values = _margin_values(models, xs, np.stack(masks_to_score), target, other)
|
||||
output: list[dict[str, Any]] = []
|
||||
cursor = 0
|
||||
grouped: dict[tuple[float, str], list[tuple[float, float]]] = {}
|
||||
for rate, method_name, _rep in row_specs[::2]:
|
||||
deletion_margin = float(score_values[cursor])
|
||||
retention_margin = float(score_values[cursor + 1])
|
||||
cursor += 2
|
||||
grouped.setdefault((rate, method_name), []).append(
|
||||
(full_margin - deletion_margin, retention_margin)
|
||||
)
|
||||
for (rate, method_name), values in sorted(grouped.items()):
|
||||
arr = np.asarray(values, dtype=np.float64)
|
||||
output.append(
|
||||
{
|
||||
"sample_id": sample_id,
|
||||
"budget": rate,
|
||||
"method": method_name,
|
||||
"replicates": len(values),
|
||||
"deletion_margin_drop_mean": float(arr[:, 0].mean()),
|
||||
"retention_margin_mean": float(arr[:, 1].mean()),
|
||||
"retention_margin_drop_mean": float((full_margin - arr[:, 1]).mean()),
|
||||
"full_margin": full_margin,
|
||||
}
|
||||
)
|
||||
return output
|
||||
@@ -0,0 +1,86 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import csv
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from data_paths import PROJECT_ROOT
|
||||
from q2.deep_learning.q2.data import RobustStats
|
||||
from ..run_experiments import _read_attachment4
|
||||
from .evaluate import (
|
||||
MODEL_SEEDS,
|
||||
_attachment_predictions_and_explanations,
|
||||
_attachment_split,
|
||||
_load_ensemble,
|
||||
)
|
||||
from .owen import hierarchical_owen_one
|
||||
|
||||
SCALER_PATH = PROJECT_ROOT / "experiments" / "q2" / "unaligned_deep_two_b128" / "unaligned_50_robust_stats.npz"
|
||||
DEFAULT_OUTPUT = PROJECT_ROOT / "output" / "q3" / "ati_ho"
|
||||
|
||||
|
||||
def _write_csv(path: Path, rows: list[dict[str, Any]]) -> None:
|
||||
if not rows:
|
||||
raise ValueError(f"no rows to write: {path}")
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
with path.open("w", encoding="utf-8-sig", newline="") as stream:
|
||||
writer = csv.DictWriter(stream, fieldnames=list(rows[0]))
|
||||
writer.writeheader()
|
||||
writer.writerows(rows)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description="Regenerate the official Q3 Attachment 4 predictions and explanations.")
|
||||
parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT)
|
||||
parser.add_argument("--device", choices=("auto", "cpu", "cuda"), default="auto")
|
||||
args = parser.parse_args()
|
||||
if args.device == "cuda" and not torch.cuda.is_available():
|
||||
parser.error("CUDA was requested but is not available")
|
||||
device_name = "cuda" if args.device == "auto" and torch.cuda.is_available() else args.device
|
||||
if device_name == "auto":
|
||||
device_name = "cpu"
|
||||
device = torch.device(device_name)
|
||||
|
||||
cases, _ = _read_attachment4("unaligned_50")
|
||||
stats = RobustStats.load(SCALER_PATH)
|
||||
attachment = _attachment_split(cases, stats)
|
||||
dims = tuple(int(values.shape[-1]) for values in attachment.x)
|
||||
models = _load_ensemble("A0", MODEL_SEEDS, dims, device)
|
||||
predictions, explanations, _, _ = _attachment_predictions_and_explanations(
|
||||
"A0", models, cases, attachment, device
|
||||
)
|
||||
|
||||
local_rows: list[dict[str, Any]] = []
|
||||
for index, case in enumerate(cases):
|
||||
xs = tuple(torch.as_tensor(values[index:index + 1], dtype=torch.float32, device=device) for values in attachment.x)
|
||||
mask = torch.as_tensor(attachment.mask[index:index + 1], dtype=torch.bool, device=device)
|
||||
result = hierarchical_owen_one(
|
||||
models, xs, mask, seed=20260926 + index, start_permutations=8, max_permutations=64
|
||||
)
|
||||
for modality, label in enumerate(("T", "A", "V")):
|
||||
for bin_index, (left, right) in enumerate(result["bin_slices"]):
|
||||
local_rows.append({
|
||||
"case_id": case["case_id"],
|
||||
"modality": label,
|
||||
"relative_bin": bin_index,
|
||||
"relative_position_start": left / 50.0,
|
||||
"relative_position_end": right / 50.0,
|
||||
"local_owen_margin_contribution": float(result["contribution"][modality, bin_index]),
|
||||
"owen_standard_error": float(result["standard_error"][modality, bin_index]),
|
||||
"permutations": result["permutations"],
|
||||
"stopping_status": result["stopping_status"],
|
||||
"physical_time_alignment": False,
|
||||
})
|
||||
|
||||
_write_csv(args.output_dir / "attachment4_predictions.csv", predictions)
|
||||
_write_csv(args.output_dir / "attachment4_explanations.csv", explanations)
|
||||
_write_csv(args.output_dir / "attachment4_local_evidence.csv", local_rows)
|
||||
print(f"Wrote {len(predictions)} predictions, {len(explanations)} explanations, and {len(local_rows)} local-evidence rows to {args.output_dir}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -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": "adapter.adapt_official_split; shared train-only robust scaler retained from Q2 V2",
|
||||
"dimensions": [int(x.shape[-1]) for x in train.x],
|
||||
"train_samples": train.n,
|
||||
"valid_samples": valid.n,
|
||||
"train_source_video_groups": len(group_sets["train"]),
|
||||
"valid_source_video_groups": len(group_sets["valid"]),
|
||||
"test_samples": sample_counts["test"],
|
||||
"test_source_video_groups": len(group_sets["test"]),
|
||||
"source_video_overlap_counts": overlap,
|
||||
"adapter_audit": adapter_audit,
|
||||
}
|
||||
return train, valid, stats, metadata
|
||||
|
||||
|
||||
def build_baseline(method: str, dims: tuple[int, int, int], device: torch.device) -> nn.Module:
|
||||
if method == EARLYCONCAT:
|
||||
return AlignedFusionModel("concat", dims=dims).to(device)
|
||||
if method == MOFE7_MLP:
|
||||
return MixtureOfFusionExperts(dims=dims, **MODEL_CONFIG).to(device)
|
||||
raise ValueError(f"unknown baseline: {method}")
|
||||
|
||||
|
||||
def _training_masks(split: Split, seed: int, epoch: int) -> tuple[np.ndarray, Counter[str]]:
|
||||
rows: list[np.ndarray] = []
|
||||
counts: Counter[str] = Counter()
|
||||
for sample_id, observed in zip(split.ids, split.mask):
|
||||
rng = np.random.default_rng(scenario_seed(TRAIN_MASK_SEED + seed, sample_id, f"train/{epoch}"))
|
||||
rate = float(rng.choice(TRAIN_RATES))
|
||||
mode = str(rng.choice(TRAIN_MODES))
|
||||
counts[f"{rate:.1f}/{mode}"] += 1
|
||||
rows.append(continuous_mask(observed, rate, mode, rng))
|
||||
return np.stack(rows), counts
|
||||
|
||||
|
||||
def _metric_row(
|
||||
split: Split,
|
||||
logits: np.ndarray,
|
||||
intensity: np.ndarray,
|
||||
probabilities: np.ndarray,
|
||||
*,
|
||||
method: str,
|
||||
seed: int,
|
||||
scenario: str,
|
||||
) -> dict[str, Any]:
|
||||
pred_class = logits.argmax(axis=-1)
|
||||
y_cls = split.y_cls
|
||||
y_reg = split.y_reg
|
||||
confidence = probabilities.max(axis=-1)
|
||||
correct = (pred_class == y_cls).astype(np.float64)
|
||||
ece = 0.0
|
||||
for left in np.linspace(0.0, 1.0, 16)[:-1]:
|
||||
right = left + 1.0 / 15.0
|
||||
hit = (confidence >= left) & (confidence < right if right < 1.0 else confidence <= right)
|
||||
if hit.any():
|
||||
ece += float(hit.mean() * abs(confidence[hit].mean() - correct[hit].mean()))
|
||||
one_hot = np.eye(3, dtype=np.float64)[y_cls]
|
||||
pearson = float(np.corrcoef(y_reg, intensity)[0, 1]) if np.std(y_reg) > 0 and np.std(intensity) > 0 else 0.0
|
||||
cm = confusion_matrix(y_cls, pred_class, labels=[0, 1, 2]).tolist()
|
||||
return {
|
||||
"method": method,
|
||||
"seed": seed,
|
||||
"scenario": scenario,
|
||||
"samples": len(y_cls),
|
||||
"accuracy": float(accuracy_score(y_cls, pred_class)),
|
||||
"macro_f1": float(f1_score(y_cls, pred_class, labels=[0, 1, 2], average="macro", zero_division=0)),
|
||||
"weighted_f1": float(f1_score(y_cls, pred_class, average="weighted", zero_division=0)),
|
||||
"negative_recall": float(recall_score(y_cls, pred_class, labels=[0, 1, 2], average=None, zero_division=0)[0]),
|
||||
"neutral_recall": float(recall_score(y_cls, pred_class, labels=[0, 1, 2], average=None, zero_division=0)[1]),
|
||||
"positive_recall": float(recall_score(y_cls, pred_class, labels=[0, 1, 2], average=None, zero_division=0)[2]),
|
||||
"mae": float(mean_absolute_error(y_reg, intensity)),
|
||||
"rmse": float(math.sqrt(mean_squared_error(y_reg, intensity))),
|
||||
"pearson": pearson,
|
||||
"ece_15bin": ece,
|
||||
"brier_multiclass": float(np.mean(np.sum((probabilities - one_hot) ** 2, axis=-1))),
|
||||
"confusion_matrix_0_1_2": json.dumps(cm),
|
||||
}
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def _predict_arrays(
|
||||
model: nn.Module,
|
||||
split: Split,
|
||||
masks: np.ndarray,
|
||||
device: torch.device,
|
||||
*,
|
||||
batch_size: int = BATCH_SIZE,
|
||||
ati: bool,
|
||||
details: bool = False,
|
||||
) -> dict[str, np.ndarray]:
|
||||
model.eval()
|
||||
outputs: dict[str, list[np.ndarray]] = {"logits": [], "intensity": [], "probabilities": []}
|
||||
if details:
|
||||
outputs.update({"params": [], "baseline": [], "main_effects": [], "pair_effects": []})
|
||||
for start in range(0, split.n, batch_size):
|
||||
end = min(split.n, start + batch_size)
|
||||
xs = tuple(torch.as_tensor(x[start:end], dtype=torch.float32, device=device) for x in split.x)
|
||||
mask = torch.as_tensor(masks[start:end], dtype=torch.bool, device=device)
|
||||
result = model(xs, mask, return_details=details) if ati else model(xs, mask)
|
||||
logits = result["logits"]
|
||||
if ati:
|
||||
probs = result["probabilities"]
|
||||
intensity = result["intensity"]
|
||||
if details:
|
||||
for key in ("params", "baseline", "main_effects", "pair_effects"):
|
||||
outputs[key].append(result[key].detach().cpu().numpy())
|
||||
else:
|
||||
probs = torch.softmax(logits, dim=-1)
|
||||
intensity = result["intensity"].clamp(-3.0, 3.0)
|
||||
outputs["logits"].append(logits.detach().cpu().numpy())
|
||||
outputs["probabilities"].append(probs.detach().cpu().numpy())
|
||||
outputs["intensity"].append(intensity.detach().cpu().numpy())
|
||||
return {key: np.concatenate(values, axis=0) for key, values in outputs.items()}
|
||||
|
||||
|
||||
def _loss_on_masks(
|
||||
model: nn.Module,
|
||||
split: Split,
|
||||
masks: np.ndarray,
|
||||
device: torch.device,
|
||||
*,
|
||||
ati: bool,
|
||||
lambda_interaction: float,
|
||||
lambda_mask: float,
|
||||
) -> float:
|
||||
model.eval()
|
||||
losses: list[float] = []
|
||||
counts: list[int] = []
|
||||
with torch.inference_mode():
|
||||
for start in range(0, split.n, BATCH_SIZE):
|
||||
end = min(split.n, start + BATCH_SIZE)
|
||||
xs = tuple(torch.as_tensor(x[start:end], dtype=torch.float32, device=device) for x in split.x)
|
||||
mb = torch.as_tensor(masks[start:end], dtype=torch.bool, device=device)
|
||||
y_cls = torch.as_tensor(split.y_cls[start:end], dtype=torch.long, device=device)
|
||||
y_reg = torch.as_tensor(split.y_reg[start:end], dtype=torch.float32, device=device)
|
||||
output = model(xs, mb, return_details=False) if ati else model(xs, mb)
|
||||
if ati:
|
||||
loss, _ = task_loss(
|
||||
output,
|
||||
y_cls,
|
||||
y_reg,
|
||||
lambda_interaction=lambda_interaction,
|
||||
lambda_mask=0.0,
|
||||
)
|
||||
else:
|
||||
loss = baseline_loss(output, y_cls, y_reg)
|
||||
losses.append(float(loss.item()))
|
||||
counts.append(end - start)
|
||||
return float(np.average(losses, weights=counts))
|
||||
|
||||
|
||||
def _selection_loss(
|
||||
model: nn.Module,
|
||||
valid: Split,
|
||||
scenarios: dict[str, np.ndarray],
|
||||
device: torch.device,
|
||||
*,
|
||||
ati: bool,
|
||||
config: ATIConfig | None,
|
||||
) -> float:
|
||||
return float(
|
||||
np.mean(
|
||||
[
|
||||
_loss_on_masks(
|
||||
model,
|
||||
valid,
|
||||
scenarios[key],
|
||||
device,
|
||||
ati=ati,
|
||||
lambda_interaction=config.lambda_interaction if config else 0.0,
|
||||
lambda_mask=0.0,
|
||||
)
|
||||
for key in SELECTION_SCENARIOS
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _save_csv(path: Path, rows: list[dict[str, Any]], *, append: bool = False) -> None:
|
||||
if not rows:
|
||||
return
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
fields = list(dict.fromkeys(key for row in rows for key in row))
|
||||
write_header = not (append and path.exists() and path.stat().st_size > 0)
|
||||
mode = "a" if append else "w"
|
||||
with path.open(mode, newline="", encoding="utf-8-sig") as stream:
|
||||
writer = csv.DictWriter(stream, fieldnames=fields, extrasaction="ignore")
|
||||
if write_header:
|
||||
writer.writeheader()
|
||||
writer.writerows(rows)
|
||||
|
||||
|
||||
def _train_one(
|
||||
method: str,
|
||||
seed: int,
|
||||
train: Split,
|
||||
valid: Split,
|
||||
valid_scenarios: dict[str, np.ndarray],
|
||||
device: torch.device,
|
||||
*,
|
||||
epochs: int,
|
||||
force: bool,
|
||||
) -> tuple[nn.Module, dict[str, Any]]:
|
||||
ati = method in CONFIGS
|
||||
config = CONFIGS[method] if ati else None
|
||||
run_dir = EXPERIMENT_ROOT / "models" / method / f"seed_{seed}"
|
||||
run_dir.mkdir(parents=True, exist_ok=True)
|
||||
checkpoint_path = run_dir / "model_best.pt"
|
||||
if checkpoint_path.is_file() and not force:
|
||||
saved = torch.load(checkpoint_path, map_location=device, weights_only=False)
|
||||
if saved.get("method") != method or int(saved.get("seed", -1)) != seed:
|
||||
raise ValueError(f"stale or mismatched checkpoint: {checkpoint_path}")
|
||||
model = ATIHOModel(tuple(x.shape[-1] for x in train.x), config).to(device) if ati else build_baseline(method, tuple(x.shape[-1] for x in train.x), device)
|
||||
model.load_state_dict(saved["state_dict"])
|
||||
model.eval()
|
||||
return model, saved
|
||||
|
||||
_seed_everything(seed)
|
||||
dims = tuple(int(x.shape[-1]) for x in train.x)
|
||||
model = ATIHOModel(dims, config).to(device) if ati else build_baseline(method, dims, device)
|
||||
optimizer = torch.optim.AdamW(model.parameters(), lr=LEARNING_RATE, weight_decay=WEIGHT_DECAY)
|
||||
train_x = tuple(torch.as_tensor(x, dtype=torch.float32, device=device) for x in train.x)
|
||||
train_cls = torch.as_tensor(train.y_cls, dtype=torch.long, device=device)
|
||||
train_reg = torch.as_tensor(train.y_reg, dtype=torch.float32, device=device)
|
||||
train_base_mask = torch.as_tensor(train.mask, dtype=torch.bool, device=device)
|
||||
order_rng = np.random.default_rng(seed + 809)
|
||||
orders = [order_rng.permutation(train.n) for _ in range(epochs)]
|
||||
history: list[dict[str, Any]] = []
|
||||
mask_counts: Counter[str] = Counter()
|
||||
best_selection = math.inf
|
||||
best_epoch = 0
|
||||
stale = 0
|
||||
start_time = time.perf_counter()
|
||||
|
||||
for epoch in range(1, epochs + 1):
|
||||
model.train()
|
||||
current_masks, current_counts = _training_masks(train, seed, epoch)
|
||||
mask_counts.update(current_counts)
|
||||
losses: list[float] = []
|
||||
order = orders[epoch - 1]
|
||||
for start in range(0, train.n, BATCH_SIZE):
|
||||
index_np = order[start : start + BATCH_SIZE]
|
||||
index = torch.as_tensor(index_np, dtype=torch.long, device=device)
|
||||
mb = torch.as_tensor(current_masks[index_np], dtype=torch.bool, device=device)
|
||||
xs = tuple(x.index_select(0, index) for x in train_x)
|
||||
output = model(xs, mb, return_details=False) if ati else model(xs, mb)
|
||||
if ati:
|
||||
loss, loss_parts = task_loss(
|
||||
output,
|
||||
train_cls.index_select(0, index),
|
||||
train_reg.index_select(0, index),
|
||||
lambda_interaction=config.lambda_interaction,
|
||||
lambda_mask=config.lambda_mask,
|
||||
mask_target=train_base_mask.index_select(0, index),
|
||||
)
|
||||
else:
|
||||
loss = baseline_loss(output, train_cls.index_select(0, index), train_reg.index_select(0, index))
|
||||
loss_parts = {"total": loss}
|
||||
optimizer.zero_grad(set_to_none=True)
|
||||
loss.backward()
|
||||
nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
|
||||
optimizer.step()
|
||||
losses.append(float(loss.detach().item()))
|
||||
|
||||
selection = _selection_loss(
|
||||
model, valid, valid_scenarios, device, ati=ati, config=config
|
||||
)
|
||||
clean = _loss_on_masks(
|
||||
model,
|
||||
valid,
|
||||
valid.mask,
|
||||
device,
|
||||
ati=ati,
|
||||
lambda_interaction=config.lambda_interaction if config else 0.0,
|
||||
lambda_mask=0.0,
|
||||
)
|
||||
row = {
|
||||
"method": method,
|
||||
"seed": seed,
|
||||
"epoch": epoch,
|
||||
"train_loss": float(np.mean(losses)),
|
||||
"valid_selection_loss": selection,
|
||||
"valid_clean_loss": clean,
|
||||
"lambda_interaction": config.lambda_interaction if config else 0.0,
|
||||
"lambda_mask": config.lambda_mask if config else 0.0,
|
||||
}
|
||||
history.append(row)
|
||||
print(
|
||||
f"[{method} seed={seed}] epoch={epoch:02d} train={row['train_loss']:.4f} "
|
||||
f"valid_selection={selection:.4f} clean={clean:.4f}",
|
||||
flush=True,
|
||||
)
|
||||
if selection < best_selection - 1e-4:
|
||||
best_selection = selection
|
||||
best_epoch = epoch
|
||||
stale = 0
|
||||
state = {
|
||||
"method": method,
|
||||
"seed": seed,
|
||||
"dims": dims,
|
||||
"steps": train.steps,
|
||||
"config": config.to_dict() if config else None,
|
||||
"state_dict": model.state_dict(),
|
||||
"best_epoch": best_epoch,
|
||||
"best_selection_loss": best_selection,
|
||||
"protocol": "official Q2 V2 unaligned_50 Relative-Progress; train-only robust scaler; video-disjoint validation",
|
||||
}
|
||||
torch.save(state, checkpoint_path)
|
||||
else:
|
||||
stale += 1
|
||||
if stale >= PATIENCE:
|
||||
break
|
||||
|
||||
saved = torch.load(checkpoint_path, map_location=device, weights_only=False)
|
||||
model.load_state_dict(saved["state_dict"])
|
||||
model.eval()
|
||||
_save_csv(run_dir / "training_history.csv", history)
|
||||
(run_dir / "training_manifest.json").write_text(
|
||||
json.dumps(
|
||||
{
|
||||
"method": method,
|
||||
"seed": seed,
|
||||
"best_epoch": best_epoch,
|
||||
"best_selection_loss": best_selection,
|
||||
"elapsed_seconds": time.perf_counter() - start_time,
|
||||
"batch_size": BATCH_SIZE,
|
||||
"epoch_limit": epochs,
|
||||
"patience": PATIENCE,
|
||||
"optimizer": "AdamW",
|
||||
"learning_rate": LEARNING_RATE,
|
||||
"weight_decay": WEIGHT_DECAY,
|
||||
"gradient_clip_norm": 1.0,
|
||||
"training_mask_rates": list(TRAIN_RATES),
|
||||
"training_mask_patterns": list(TRAIN_MODES),
|
||||
"training_mask_seed_base": TRAIN_MASK_SEED,
|
||||
"same_orders_and_masks_across_methods_for_same_seed": True,
|
||||
"config": config.to_dict() if config else {"model_config": MODEL_CONFIG},
|
||||
"history": history,
|
||||
"mask_counts": dict(mask_counts),
|
||||
},
|
||||
ensure_ascii=False,
|
||||
indent=2,
|
||||
),
|
||||
encoding="utf-8",
|
||||
)
|
||||
return model, saved
|
||||
|
||||
|
||||
def _evaluate_job(
|
||||
method: str,
|
||||
seed: int,
|
||||
model: nn.Module,
|
||||
valid: Split,
|
||||
scenarios: dict[str, np.ndarray],
|
||||
device: torch.device,
|
||||
) -> list[dict[str, Any]]:
|
||||
ati = method in CONFIGS
|
||||
rows: list[dict[str, Any]] = []
|
||||
selected = {"clean": valid.mask}
|
||||
selected.update({key: scenarios[key] for key in SELECTION_SCENARIOS if key != "0.0/none"})
|
||||
for name, masks in selected.items():
|
||||
predictions = _predict_arrays(model, valid, masks, device, ati=ati)
|
||||
rows.append(
|
||||
_metric_row(
|
||||
valid,
|
||||
predictions["logits"],
|
||||
predictions["intensity"],
|
||||
predictions["probabilities"],
|
||||
method=method,
|
||||
seed=seed,
|
||||
scenario=name,
|
||||
)
|
||||
)
|
||||
return rows
|
||||
|
||||
|
||||
def _write_root_manifest(data_meta: dict[str, Any], device: torch.device, epochs: int) -> None:
|
||||
EXPERIMENT_ROOT.mkdir(parents=True, exist_ok=True)
|
||||
try:
|
||||
git_sha = subprocess.check_output(
|
||||
["git", "rev-parse", "HEAD"], cwd=PROJECT_ROOT, text=True, stderr=subprocess.DEVNULL
|
||||
).strip()
|
||||
except Exception:
|
||||
git_sha = None
|
||||
info: dict[str, Any] = {
|
||||
"experiment": "ATI–HO Q3 staged training and structural attribution audit",
|
||||
"created_utc": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()),
|
||||
"device": str(device),
|
||||
"torch_version": torch.__version__,
|
||||
"cuda_available": torch.cuda.is_available(),
|
||||
"cuda_version": torch.version.cuda,
|
||||
"gpu": torch.cuda.get_device_name(0) if torch.cuda.is_available() else None,
|
||||
"python": platform.python_version(),
|
||||
"seeds": list(MODEL_SEEDS),
|
||||
"epochs_max": epochs,
|
||||
"training_protocol": {
|
||||
"batch_size": BATCH_SIZE,
|
||||
"early_stopping_patience": PATIENCE,
|
||||
"optimizer": "AdamW",
|
||||
"learning_rate": LEARNING_RATE,
|
||||
"weight_decay": WEIGHT_DECAY,
|
||||
"gradient_clip_norm": 1.0,
|
||||
"training_mask_rates": list(TRAIN_RATES),
|
||||
"training_mask_patterns": list(TRAIN_MODES),
|
||||
"validation_selection_scenarios": list(SELECTION_SCENARIOS),
|
||||
"held_out_attachment4_touched_during_training": False,
|
||||
},
|
||||
"ati_output": {
|
||||
"parameter_vector": "3 centered class logits + r_negative + r_positive",
|
||||
"intensity": "negative/positive magnitudes are 3*sigmoid(r); neutral class is exactly zero",
|
||||
"loss": "cross entropy + conditional magnitude SmoothL1 + 0.2*Huber(delta=0.25) + configured regularizers",
|
||||
"baseline_checkpoint_reuse": "No: retrain B0 and B1 on the fixed ATI split/mask schedule because existing Q2 checkpoints differ in seeds, batch size, and schedule.",
|
||||
"calibration_temperature": 1.0,
|
||||
},
|
||||
"data": data_meta,
|
||||
}
|
||||
(EXPERIMENT_ROOT / "run_manifest.json").write_text(
|
||||
json.dumps(info, ensure_ascii=False, indent=2), encoding="utf-8"
|
||||
)
|
||||
|
||||
|
||||
def run_stage1(train: Split, valid: Split, scenarios: dict[str, np.ndarray], device: torch.device, epochs: int, force: bool) -> None:
|
||||
(EXPERIMENT_ROOT / "stage1_complete.json").unlink(missing_ok=True)
|
||||
rows: list[dict[str, Any]] = []
|
||||
models: dict[str, nn.Module] = {}
|
||||
for method in ("A0", "A1", "A2", "A3", "D0"):
|
||||
model, saved = _train_one(method, 42, train, valid, scenarios, device, epochs=epochs, force=force)
|
||||
models[method] = model
|
||||
rows.extend(_evaluate_job(method, 42, model, valid, scenarios, device))
|
||||
print(f"[{method}] best_epoch={saved['best_epoch']} selected_loss={saved['best_selection_loss']:.5f}", flush=True)
|
||||
_save_csv(EXPERIMENT_ROOT / "validation_results.csv", rows)
|
||||
candidates = []
|
||||
for method in ("A0", "A1", "A2", "A3"):
|
||||
checkpoint = torch.load(EXPERIMENT_ROOT / "models" / method / "seed_42" / "model_best.pt", map_location="cpu", weights_only=False)
|
||||
candidates.append({
|
||||
"method": method,
|
||||
"best_selection_loss": float(checkpoint["best_selection_loss"]),
|
||||
"best_epoch": int(checkpoint["best_epoch"]),
|
||||
})
|
||||
candidates.sort(key=lambda row: row["best_selection_loss"])
|
||||
choice = candidates[0]["method"]
|
||||
candidate_doc = {
|
||||
"stage1_candidates": candidates,
|
||||
"provisional_selected_candidate": choice,
|
||||
"selection_rule": "lowest fixed four-scenario ATI task loss on the locked official validation split; seed 42 only in Stage I",
|
||||
}
|
||||
(EXPERIMENT_ROOT / "provisional_candidate.json").write_text(
|
||||
json.dumps(candidate_doc, ensure_ascii=False, indent=2), encoding="utf-8"
|
||||
)
|
||||
smoke_count = min(16, valid.n)
|
||||
smoke_xs = tuple(torch.as_tensor(x[:smoke_count], dtype=torch.float32, device=device) for x in valid.x)
|
||||
smoke_mask = torch.as_tensor(valid.mask[:smoke_count], dtype=torch.bool, device=device)
|
||||
structural_rows = []
|
||||
for method, model in models.items():
|
||||
report = structural_audit(model, smoke_xs, smoke_mask)
|
||||
row = {"method": method, "seed": 42, "samples": smoke_count, **report}
|
||||
row["pair_single_missing_anchor_max_abs"] = json.dumps(
|
||||
report["pair_single_missing_anchor_max_abs"], sort_keys=True
|
||||
)
|
||||
structural_rows.append(row)
|
||||
if not report["checks_pass"]:
|
||||
raise RuntimeError(f"Stage I structural audit failed for {method}: {report}")
|
||||
if method == "D0" and not report["unanchored_control_detected_leakage"]:
|
||||
raise RuntimeError("D0 unanchored diagnostic did not expose the expected missing-modality leakage")
|
||||
_save_csv(EXPERIMENT_ROOT / "structural_audit.csv", structural_rows)
|
||||
shapley = exact_shapley_audit([models[choice]], smoke_xs, smoke_mask, batch_size=32)
|
||||
if not np.asarray(shapley["class_pass"]).all():
|
||||
raise RuntimeError(
|
||||
f"Stage I analytic-vs-exact Shapley audit failed for {choice}: "
|
||||
f"max_abs={float(np.max(shapley['class_abs_error'])):.8g}"
|
||||
)
|
||||
shapley_rows = []
|
||||
for index in range(smoke_count):
|
||||
shapley_rows.append(
|
||||
{
|
||||
"sample_index": index,
|
||||
"method": choice,
|
||||
"target_class": int(shapley["target_class"][index]),
|
||||
"runner_up_class": int(shapley["runner_up_class"][index]),
|
||||
"analytic_T": float(shapley["analytic_class"][index, 0]),
|
||||
"analytic_A": float(shapley["analytic_class"][index, 1]),
|
||||
"analytic_V": float(shapley["analytic_class"][index, 2]),
|
||||
"exact_T": float(shapley["exact_class"][index, 0]),
|
||||
"exact_A": float(shapley["exact_class"][index, 1]),
|
||||
"exact_V": float(shapley["exact_class"][index, 2]),
|
||||
"max_abs_error": float(shapley["class_abs_error"][index].max()),
|
||||
"pass": bool(shapley["class_pass"][index].all()),
|
||||
}
|
||||
)
|
||||
_save_csv(EXPERIMENT_ROOT / "shapley_audit_seed42_smoke.csv", shapley_rows)
|
||||
(EXPERIMENT_ROOT / "stage1_complete.json").write_text(
|
||||
json.dumps(
|
||||
{
|
||||
"models": ["A0", "A1", "A2", "A3", "D0"],
|
||||
"seed": 42,
|
||||
"structural_audit_pass": True,
|
||||
"analytic_vs_exact_shapley_pass": True,
|
||||
"shapley_smoke_samples": smoke_count,
|
||||
"selected_candidate": choice,
|
||||
},
|
||||
indent=2,
|
||||
),
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
|
||||
def run_stage2(train: Split, valid: Split, scenarios: dict[str, np.ndarray], device: torch.device, epochs: int, force: bool) -> None:
|
||||
provisional_path = EXPERIMENT_ROOT / "provisional_candidate.json"
|
||||
if not provisional_path.is_file():
|
||||
raise FileNotFoundError("run Stage I before Stage II; provisional_candidate.json is missing")
|
||||
selected = json.loads(provisional_path.read_text(encoding="utf-8"))["provisional_selected_candidate"]
|
||||
key_ablations = {
|
||||
"A0": ["A1", "A2"],
|
||||
"A1": ["A0", "A2"],
|
||||
"A2": ["A0", "A1"],
|
||||
"A3": ["A2", "A1"],
|
||||
}[selected]
|
||||
methods = list(dict.fromkeys([selected, *key_ablations]))
|
||||
rows: list[dict[str, Any]] = []
|
||||
for method in (EARLYCONCAT, MOFE7_MLP, *methods):
|
||||
for seed in MODEL_SEEDS:
|
||||
model, saved = _train_one(method, seed, train, valid, scenarios, device, epochs=epochs, force=force)
|
||||
rows.extend(_evaluate_job(method, seed, model, valid, scenarios, device))
|
||||
print(
|
||||
f"[{method} seed={seed}] best_epoch={saved['best_epoch']} "
|
||||
f"selected_loss={saved['best_selection_loss']:.5f}",
|
||||
flush=True,
|
||||
)
|
||||
_save_csv(EXPERIMENT_ROOT / "validation_results.csv", rows, append=True)
|
||||
stage2 = {
|
||||
"selected_candidate": selected,
|
||||
"key_ablations": key_ablations,
|
||||
"baseline_methods": [EARLYCONCAT, MOFE7_MLP],
|
||||
"seeds": list(MODEL_SEEDS),
|
||||
"baseline_checkpoints_retrained": True,
|
||||
"all_selection_uses_locked_validation_only": True,
|
||||
}
|
||||
(EXPERIMENT_ROOT / "stage2_complete.json").write_text(
|
||||
json.dumps(stage2, ensure_ascii=False, indent=2), encoding="utf-8"
|
||||
)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description="Train ATI–HO and compatible Q3 baselines.")
|
||||
parser.add_argument("--phase", choices=("stage1", "stage2", "all"), default="all")
|
||||
parser.add_argument("--device", default="auto")
|
||||
parser.add_argument("--epochs", type=int, default=EPOCH_LIMIT)
|
||||
parser.add_argument("--force", action="store_true")
|
||||
args = parser.parse_args()
|
||||
device = torch.device("cuda" if args.device == "auto" and torch.cuda.is_available() else "cpu" if args.device == "auto" else args.device)
|
||||
train, valid, _stats, data_meta = load_training_data()
|
||||
scenarios = make_scenarios(valid, SCENARIO_SEED)
|
||||
_write_root_manifest(data_meta, device, args.epochs)
|
||||
print(
|
||||
f"ATI–HO protocol: train={train.n} valid={valid.n} groups="
|
||||
f"{data_meta['train_source_video_groups']}/{data_meta['valid_source_video_groups']} "
|
||||
f"dims={data_meta['dimensions']} device={device}",
|
||||
flush=True,
|
||||
)
|
||||
if args.phase in {"stage1", "all"}:
|
||||
run_stage1(train, valid, scenarios, device, args.epochs, args.force)
|
||||
if args.phase in {"stage2", "all"}:
|
||||
run_stage2(train, valid, scenarios, device, args.epochs, args.force)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user