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