Download validation/forward.py from HWresearch/GNN4Colliders: direct link, hf CLI and curl.
- Browser
- Download file 2.9 kB
-
https://huggingface.co/HWresearch/GNN4Colliders/resolve/main/validation/forward.py
- Command line
-
hf download hf://HWresearch/GNN4Colliders/validation/forward.py
-
curl -L -o forward.py https://huggingface.co/HWresearch/GNN4Colliders/resolve/main/validation/forward.py
2.9 kB
| #!/usr/bin/env python3 | |
| """Run fixed-weight ROOT-GNN inference on a normalized validation artifact.""" | |
| from __future__ import annotations | |
| import argparse | |
| from dataclasses import replace | |
| from pathlib import Path | |
| import dgl | |
| import numpy as np | |
| import torch | |
| from validation.artifacts import load_artifact, save_artifact | |
| def _graph(artifact, index: int): | |
| nodes = torch.from_numpy(artifact.event_nodes(index)).to(torch.float32) | |
| src, dst, edges = artifact.event_edges(index) | |
| graph = dgl.graph( | |
| (torch.from_numpy(src), torch.from_numpy(dst)), num_nodes=nodes.shape[0] | |
| ) | |
| graph.ndata["features"] = nodes | |
| graph.edata["features"] = torch.from_numpy(edges).to(torch.float32) | |
| return 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) | |
| args = parser.parse_args() | |
| artifact = load_artifact(args.artifact) | |
| event_count = ( | |
| artifact.event_count | |
| if args.events is None | |
| else min(args.events, artifact.event_count) | |
| ) | |
| graphs = [_graph(artifact, index) for index in range(event_count)] | |
| batch_graph = dgl.batch(graphs) | |
| first = graphs[0] | |
| empty_globals = torch.empty((1, 0), dtype=torch.float32) | |
| from gnn4colliders.models.root_gnn import ( | |
| EdgeNetwork, | |
| load_legacy_edge_network_state_dict, | |
| ) | |
| model = EdgeNetwork(first, 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(model, payload) | |
| model.eval() | |
| with torch.inference_mode(): | |
| logits = model(batch_graph, None).cpu().numpy() | |
| args.output.mkdir(parents=True, exist_ok=True) | |
| reload_path = args.output / "checkpoint_reload.pt" | |
| torch.save({"model_state_dict": model.state_dict()}, reload_path) | |
| fresh = type(model)(first, empty_globals, 64, 12, 4, 4, dropout=0.0) | |
| fresh.load_state_dict( | |
| torch.load(reload_path, map_location="cpu")["model_state_dict"] | |
| ) | |
| fresh.eval() | |
| with torch.inference_mode(): | |
| reloaded_logits = fresh(batch_graph, None).cpu().numpy() | |
| scores = torch.softmax(torch.from_numpy(logits), dim=1).numpy() | |
| predictions = scores.argmax(axis=1) | |
| updated = replace( | |
| artifact, | |
| logits=logits, | |
| scores=scores, | |
| predictions=predictions, | |
| manifest={ | |
| **artifact.manifest, | |
| "implementation": "rewrite", | |
| "checkpoint": str(args.checkpoint), | |
| "reload_max_abs": float(np.abs(logits - reloaded_logits).max()), | |
| }, | |
| ) | |
| save_artifact(updated, args.output) | |
| return 0 | |
| if __name__ == "__main__": | |
| raise SystemExit(main()) | |