55 lines
1.8 KiB
Python
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()
|