Download validation/transfer.py from HWresearch/GNN4Colliders: direct link, hf CLI and curl.
- Browser
- Download file 1.88 kB
-
https://huggingface.co/HWresearch/GNN4Colliders/resolve/main/validation/transfer.py
- Command line
-
hf download hf://HWresearch/GNN4Colliders/validation/transfer.py
-
curl -L -o transfer.py https://huggingface.co/HWresearch/GNN4Colliders/resolve/main/validation/transfer.py
1.88 kB
| #!/usr/bin/env python3 | |
| """Run canonical ROOT-GNN transfer-learning forward with a fixed binary head.""" | |
| from __future__ import annotations | |
| import argparse | |
| from pathlib import Path | |
| import dgl | |
| import numpy as np | |
| import torch | |
| from validation.artifacts import load_artifact | |
| from validation.forward import _graph | |
| def main() -> int: | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument("--artifact", type=Path, required=True) | |
| parser.add_argument("--checkpoint", type=Path, required=True) | |
| parser.add_argument("--output", type=Path, required=True) | |
| parser.add_argument("--events", type=int, default=96) | |
| parser.add_argument("--trainable-backbone", action="store_true") | |
| args = parser.parse_args() | |
| artifact = load_artifact(args.artifact) | |
| count = min(args.events, artifact.event_count) | |
| graphs = [_graph(artifact, index) for index in range(count)] | |
| batch_graph = dgl.batch(graphs) | |
| empty_globals = torch.empty((1, 0), dtype=torch.float32) | |
| from gnn4colliders.models.root_gnn import ( | |
| EdgeNetwork, | |
| FineTunedEdgeNetwork, | |
| load_legacy_edge_network_state_dict, | |
| ) | |
| backbone = EdgeNetwork(graphs[0], empty_globals, 64, 12, 4, 4, dropout=0.0) | |
| payload = torch.load(args.checkpoint, map_location="cpu", weights_only=False) | |
| load_legacy_edge_network_state_dict(backbone, payload) | |
| model = FineTunedEdgeNetwork( | |
| backbone, 1, freeze_backbone=not args.trainable_backbone | |
| ) | |
| classifier = model.classifier | |
| torch.manual_seed(20260818) | |
| classifier.load_state_dict(torch.nn.Linear(64, 1).state_dict()) | |
| model.eval() | |
| with torch.inference_mode(): | |
| logits = model(batch_graph, None).cpu().numpy() | |
| args.output.parent.mkdir(parents=True, exist_ok=True) | |
| np.savez_compressed(args.output, logits=logits) | |
| return 0 | |
| if __name__ == "__main__": | |
| raise SystemExit(main()) | |