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()