Train ATI-HO and finalize project outputs
This commit is contained in:
@@ -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,
|
||||
}
|
||||
Reference in New Issue
Block a user