GNN4Colliders / benchmarks /benchmark_preprocessing.py
ho22joshua's picture
perf: profile and optimize ROOT-GNN execution
9dcd2b7
Raw
History Blame
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()