Train ATI-HO and finalize project outputs

This commit is contained in:
2026-09-26 16:05:44 +08:00
parent a86560da64
commit 9cdd604117
358 changed files with 10540 additions and 173 deletions
+68
View File
@@ -0,0 +1,68 @@
from __future__ import annotations
from typing import Any
import torch
from torch import nn
PAIR_MODES = ((0, 1), (0, 2), (1, 2))
@torch.inference_mode()
def structural_audit(
model: nn.Module,
xs: tuple[torch.Tensor, torch.Tensor, torch.Tensor],
masks: torch.Tensor,
*,
atol: float = 1e-6,
) -> dict[str, Any]:
model.eval()
output = model(xs, masks, return_details=True)
reconstructed = output["baseline"] + output["main_effects"].sum(dim=1) + output["pair_effects"].sum(dim=1)
additive_residual = torch.max(torch.abs(reconstructed - output["params"])).item()
absent_xs = tuple(torch.zeros_like(x) for x in xs)
absent_mask = torch.zeros_like(masks, dtype=torch.bool)
absent = model(absent_xs, absent_mask, return_details=True)
main_zero_residual = torch.max(torch.abs(absent["main_effects"])).item()
full_zero_residual = torch.max(torch.abs(absent["params"] - absent["baseline"])).item()
pair_zero_residual = torch.max(torch.abs(absent["pair_effects"])).item()
pair_anchor_residuals: dict[str, float] = {}
for pair_index, (left, right) in enumerate(PAIR_MODES):
maxima = []
for hidden in (left, right):
altered = masks.clone()
altered[..., hidden] = False
out = model(xs, altered, return_details=True)
maxima.append(torch.max(torch.abs(out["pair_effects"][:, pair_index])).item())
pair_anchor_residuals[f"{left}{right}"] = max(maxima)
finite_count = 0
nonfinite_count = 0
for key in ("params", "logits", "intensity", "main_effects", "pair_effects"):
tensor = output[key]
finite_count += int(torch.isfinite(tensor).sum().item())
nonfinite_count += int((~torch.isfinite(tensor)).sum().item())
anchored = bool(getattr(getattr(model, "config", None), "anchored", False))
main_additive_pass = max(additive_residual, main_zero_residual) <= atol
pair_pass = max(pair_anchor_residuals.values(), default=0.0) <= atol
baseline_pass = full_zero_residual <= atol and pair_zero_residual <= atol
checks_pass = main_additive_pass and nonfinite_count == 0 and (baseline_pass if anchored else True)
return {
"anchored": anchored,
"main_effect_zero_anchor_max_abs": main_zero_residual,
"pair_effect_zero_anchor_max_abs": pair_zero_residual,
"pair_single_missing_anchor_max_abs": pair_anchor_residuals,
"additive_reconstruction_max_abs": additive_residual,
"all_modalities_missing_equals_baseline_max_abs": full_zero_residual,
"finite_value_count": finite_count,
"nonfinite_value_count": nonfinite_count,
"main_and_additivity_pass": main_additive_pass,
"full_baseline_anchor_pass": baseline_pass if anchored else None,
"pair_anchor_pass": pair_pass if anchored else None,
"unanchored_control_detected_leakage": ((not pair_pass) or not baseline_pass) if not anchored else False,
"checks_pass": checks_pass,
}