| """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() |
|
|