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