Download validation/batching.py from HWresearch/GNN4Colliders: direct link, hf CLI and curl.
- Browser
- Download file 2.08 kB
-
https://huggingface.co/HWresearch/GNN4Colliders/resolve/main/validation/batching.py
- Command line
-
hf download hf://HWresearch/GNN4Colliders/validation/batching.py
-
curl -L -o batching.py https://huggingface.co/HWresearch/GNN4Colliders/resolve/main/validation/batching.py
2.08 kB
| #!/usr/bin/env python3 | |
| """Dump deterministic batch membership and aggregate sizes.""" | |
| from __future__ import annotations | |
| import argparse | |
| import json | |
| from pathlib import Path | |
| 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("--output", type=Path, required=True) | |
| parser.add_argument("--batch-size", type=int, default=8) | |
| args = parser.parse_args() | |
| artifact = load_artifact(args.artifact) | |
| from gnn4colliders.data import ( | |
| EventMetadata, | |
| GraphDataLoader, | |
| GraphDataset, | |
| GraphSample, | |
| ) | |
| samples = [] | |
| for index in range(artifact.event_count): | |
| samples.append( | |
| GraphSample( | |
| _graph(artifact, index), | |
| torch.tensor(artifact.labels[index]), | |
| None, | |
| EventMetadata( | |
| int(artifact.folds[index]), | |
| float(artifact.weights[index]), | |
| str(artifact.sample_id[index]), | |
| {"index": index}, | |
| ), | |
| ) | |
| ) | |
| loader = GraphDataLoader(GraphDataset(samples), args.batch_size, shuffle=False) | |
| batches = [] | |
| for batch in loader: | |
| graph = batch.graph | |
| labels = batch.labels | |
| ids = list(batch.metadata.sample_id) | |
| weights = batch.metadata.weight.tolist() | |
| batches.append( | |
| { | |
| "sample_id": ids, | |
| "num_graphs": len(ids), | |
| "nodes": int(graph.num_nodes()), | |
| "edges": int(graph.num_edges()), | |
| "labels": labels.reshape(-1).tolist(), | |
| "weights": weights, | |
| } | |
| ) | |
| args.output.parent.mkdir(parents=True, exist_ok=True) | |
| args.output.write_text( | |
| json.dumps({"batch_size": args.batch_size, "batches": batches}, indent=2) + "\n" | |
| ) | |
| return 0 | |
| if __name__ == "__main__": | |
| raise SystemExit(main()) | |