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, }