Download validation/compare_transfer.py from HWresearch/GNN4Colliders: direct link, hf CLI and curl.
- Browser
- Download file 1.89 kB
-
https://huggingface.co/HWresearch/GNN4Colliders/resolve/main/validation/compare_transfer.py
- Command line
-
hf download hf://HWresearch/GNN4Colliders/validation/compare_transfer.py
-
curl -L -o compare_transfer.py https://huggingface.co/HWresearch/GNN4Colliders/resolve/main/validation/compare_transfer.py
1.89 kB
| #!/usr/bin/env python3 | |
| """Compare frozen/trainable transfer-forward artifacts.""" | |
| from __future__ import annotations | |
| import argparse | |
| import json | |
| from pathlib import Path | |
| import numpy as np | |
| def compare(left: np.ndarray, right: np.ndarray) -> dict[str, object]: | |
| delta = np.abs(left.astype(np.float64) - right.astype(np.float64)) | |
| passed = np.isclose(left, right, atol=1e-5, rtol=1e-5) | |
| return { | |
| "status": "PASS" if bool(np.all(passed)) else "FAIL", | |
| "shape": list(left.shape), | |
| "max_abs": float(delta.max()), | |
| "outside_tolerance": int((~passed).sum()), | |
| } | |
| def main() -> int: | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument("reference_frozen", type=Path) | |
| parser.add_argument("candidate_frozen", type=Path) | |
| parser.add_argument("reference_trainable", type=Path) | |
| parser.add_argument("candidate_trainable", type=Path) | |
| parser.add_argument("output", type=Path) | |
| args = parser.parse_args() | |
| reports = {} | |
| for name, left_path, right_path in ( | |
| ("frozen_backbone", args.reference_frozen, args.candidate_frozen), | |
| ("trainable_backbone", args.reference_trainable, args.candidate_trainable), | |
| ): | |
| with np.load(left_path) as left, np.load(right_path) as right: | |
| reports[name] = compare(left["logits"], right["logits"]) | |
| report = { | |
| "schema_version": 1, | |
| "stages": reports, | |
| "overall_status": "PASS" | |
| if all(item["status"] == "PASS" for item in reports.values()) | |
| else "FAIL", | |
| } | |
| args.output.mkdir(parents=True, exist_ok=True) | |
| (args.output / "transfer.json").write_text( | |
| json.dumps(report, indent=2, sort_keys=True) + "\n" | |
| ) | |
| print(json.dumps({"overall_status": report["overall_status"]}, indent=2)) | |
| return 0 if report["overall_status"] == "PASS" else 1 | |
| if __name__ == "__main__": | |
| raise SystemExit(main()) | |