Files

55 lines
1.8 KiB
Python

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