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