File size: 1,924 Bytes
9dcd2b7 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 | """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()
|