Download benchmarks/benchmark_preprocessing.py from HWresearch/GNN4Colliders: direct link, hf CLI and curl.
- Browser
- Download file 1.92 kB
-
https://huggingface.co/HWresearch/GNN4Colliders/resolve/main/benchmarks/benchmark_preprocessing.py
- Command line
-
hf download hf://HWresearch/GNN4Colliders/benchmarks/benchmark_preprocessing.py
-
curl -L -o benchmark_preprocessing.py https://huggingface.co/HWresearch/GNN4Colliders/resolve/main/benchmarks/benchmark_preprocessing.py
1.92 kB
| """Measure shared feature and graph preprocessing on fixed synthetic events.""" | |
| from __future__ import annotations | |
| import torch | |
| from gnn4colliders.features import build_node_features | |
| from gnn4colliders.graphs import build_edge_features, fully_connected_edges | |
| try: | |
| from ._common import common_parser, measure, metadata, report | |
| except ImportError: | |
| from _common import common_parser, measure, metadata, report | |
| def main() -> None: | |
| parser = common_parser(__doc__) | |
| parser.add_argument("--nodes", type=int, default=32) | |
| args = parser.parse_args() | |
| torch.manual_seed(args.seed) | |
| event = { | |
| "pt": torch.arange(args.nodes, dtype=torch.float32) + 1, | |
| "eta": torch.linspace(-2, 2, args.nodes), | |
| "phi": torch.linspace(-3.0, 3.0, args.nodes), | |
| } | |
| branches = [["pt"], ["eta"], ["phi"], "CALC_E", [1.0], [0.0], "NODE_TYPE"] | |
| object_types = ["vector"] | |
| scales = [1.0] * 7 | |
| device = torch.device(args.device) | |
| feature_result = measure( | |
| lambda: build_node_features(event, branches, object_types, scales), | |
| iterations=args.iterations, | |
| warmup=args.warmup, | |
| device=device, | |
| ) | |
| nodes = build_node_features(event, branches, object_types, scales)[0] | |
| src, dst = fully_connected_edges(args.nodes) | |
| graph_result = measure( | |
| lambda: build_edge_features(nodes, src, dst, eta_index=1, phi_index=2), | |
| iterations=args.iterations, | |
| warmup=args.warmup, | |
| device=device, | |
| ) | |
| base = {**metadata(device), "nodes": args.nodes, "iterations": args.iterations} | |
| report( | |
| "feature_construction", | |
| { | |
| **base, | |
| **feature_result, | |
| "events_per_second": 1000 / feature_result["mean_ms"], | |
| }, | |
| ) | |
| report( | |
| "edge_features", | |
| {**base, **graph_result, "graphs_per_second": 1000 / graph_result["mean_ms"]}, | |
| ) | |
| if __name__ == "__main__": | |
| main() | |