ATLAS Offline Software
Loading...
Searching...
No Matches
muonEdgeSegmentInference.py
Go to the documentation of this file.
1#!/usr/bin/env python
2# Copyright (C) 2002-2026 CERN for the benefit of the ATLAS collaboration
3"""Run segment-edge ONNX inference with an optional JSONL parity dump.
4"""
5
6def main(args):
7 from MuonGeoModelTestR4.testGeoModel import setupGeoR4TestCfg
8 from MuonConfig.MuonConfigUtils import executeTest
9 from AthenaConfiguration.AllConfigFlags import initConfigFlags
10 flags = initConfigFlags()
11
12 if args.athenaDebug:
13 flags.Exec.DebugMessageComponents = [
14 "SegmentEdgeInferenceAlg",
15 "SegmentEdgeInferenceAlg.SegmentEdgeClassifierTool",
16 "SegmentEdgeInferenceAlg.SegmentEdgeClassifierTool.OnnxRuntimeSessionToolCPU",
17 "SegmentEdgeInferenceAlg.SegmentEdgeClassifierTool.OnnxRuntimeSessionToolCUDA",
18 "GraphInferenceAlg",
19 "GraphBucketFilterTool",
20 ]
21 print("INFO: Exec.DebugMessageComponents configured:", flags.Exec.DebugMessageComponents)
22
23 from AthOnnxComps.OnnxRuntimeFlags import OnnxRuntimeType
24 if args.use_cpu:
25 flags.AthOnnx.ExecutionProvider = OnnxRuntimeType.CPU
26 else:
27 flags.AthOnnx.ExecutionProvider = OnnxRuntimeType.CUDA
28
29 flags, cfg = setupGeoR4TestCfg(args, flags)
30
31 from MuonConfig.MuonDataPrepConfig import xAODUncalibMeasPrepCfg
32 cfg.merge(xAODUncalibMeasPrepCfg(flags))
33
34 from MuonSpacePointFormation.SpacePointFormationConfig import MuonSpacePointFormationCfg
35 cfg.merge(MuonSpacePointFormationCfg(flags))
36
37 from MuonPatternRecognitionAlgs.MuonPatternRecognitionConfig import MuonPatternRecognitionCfg
38 if args.doMLBucketFilter:
39 from MuonInference.InferenceConfig import GraphBucketFilterToolCfg, GraphInferenceAlgCfg
40 bucket_tool = cfg.popToolsAndMerge(
41 GraphBucketFilterToolCfg(
42 flags,
43 ModelPath=args.bucket_model_path,
44 ScoreThreshold=args.score_threshold,
45 OutputName=args.output_name,
46 SingleOutputMode=args.single_output_mode,
47 )
48 )
49 cfg.merge(GraphInferenceAlgCfg(flags, InferenceTools=[bucket_tool]))
50 cfg.merge(MuonPatternRecognitionCfg(flags))
51 cfg.getEventAlgo("MuonEtaHoughTransformAlg").SpacePointContainer = "FilteredMlBuckets"
52 else:
53 cfg.merge(MuonPatternRecognitionCfg(flags))
54
55 output_level = 1 if args.athenaDebug else 3
56 edge_space_point_key = "FilteredMlBuckets" if args.doMLBucketFilter else "MuonSpacePoints"
57
58 edge_classifier_kwargs = {
59 "ModelPath": args.edgeModel,
60 "ReadSpacePoints": edge_space_point_key,
61 "DebugDumpFile": args.segment_edge_debug_dump_file,
62 "DebugDumpMaxEvents": args.segment_edge_debug_dump_max_events,
63 "MaxDeltaThetaDeg": args.max_delta_theta_deg,
64 "MaxDeltaSector": args.max_delta_sector,
65 "SectorModulo": args.sector_modulo,
66 "OutputLevel": output_level,
67 }
68 from MuonInference.InferenceConfig import SegmentEdgeInferenceAlgCfg
69 cfg.merge(
70 SegmentEdgeInferenceAlgCfg(
71 flags,
72 EdgeClassifierTool=edge_classifier_kwargs,
73 PairGateThreshold=args.edge_threshold,
74 OutputLevel=output_level,
75 )
76 )
77
78 executeTest(cfg)
79
80if __name__ == "__main__":
81 from MuonGeoModelTestR4.testGeoModel import SetupArgParser, MuonPhaseIITestDefaults
82 parser = SetupArgParser()
83 parser.set_defaults(nEvents=-1)
84 parser.set_defaults(inputFile=MuonPhaseIITestDefaults.HITS_PG_R3)
85 parser.add_argument("--edgeModel", "--edge-model", required=True, dest="edgeModel",
86 help="ONNX segment-edge classifier")
87 parser.add_argument("--edge-threshold", type=float, default=0.975,
88 help="Minimum high-confidence edge probability used to form ML track components")
89 parser.add_argument("--max-delta-theta-deg", type=float, default=35.0,
90 help="Graph edge direction window in degrees")
91 parser.add_argument("--max-delta-sector", type=int, default=1,
92 help="Graph edge sector window")
93 parser.add_argument("--sector-modulo", type=int, default=16,
94 help="Sector wrap-around modulo")
95
96 from MuonInference.InferenceConfig import (
97 DEFAULT_BUCKET_MODEL_PATH,
98 DEFAULT_BUCKET_SCORE_THRESHOLD,
99 DEFAULT_BUCKET_SINGLE_OUTPUT_MODE,
100 )
101 parser.add_argument("--doMLBucketFilter", dest="doMLBucketFilter", action="store_true", default=True)
102 parser.add_argument("--noMLBucketFilter", dest="doMLBucketFilter", action="store_false")
103 parser.add_argument("--bucket-model-path", dest="bucket_model_path", default=DEFAULT_BUCKET_MODEL_PATH)
104 parser.add_argument("--score-threshold", type=float, default=DEFAULT_BUCKET_SCORE_THRESHOLD)
105 parser.add_argument("--output-name", default="logits", dest="output_name",
106 help="Bucket filter ONNX output tensor name")
107 score_mode = parser.add_mutually_exclusive_group()
108 score_mode.add_argument("--single-output-mode", choices=("logit", "prob"), default=DEFAULT_BUCKET_SINGLE_OUTPUT_MODE, dest="single_output_mode",
109 help="Scalar ONNX-output interpretation. 'logit' applies sigmoid before thresholding.")
110 score_mode.add_argument("--is-logit", action="store_const", const="logit", dest="single_output_mode",
111 help="Alias for --single-output-mode logit.")
112 score_mode.add_argument("--is-prob", action="store_const", const="prob", dest="single_output_mode",
113 help="Alias for --single-output-mode prob.")
114
115 parser.add_argument("--segment-edge-debug-dump-file", default="",
116 help="Optional JSONL with exact x, edge_index, edge_attr, logits and probabilities")
117 parser.add_argument("--segment-edge-debug-dump-max-events", type=int, default=0,
118 help="Maximum graph events written to the segment-edge JSONL dump; 0 means all")
119 parser.add_argument("--athenaDebug", action="store_true",
120 help="Enable Athena DEBUG verbosity")
121 parser.add_argument("--use-cpu", action="store_true", default=False, help="Force CPU for ONNX inference")
122
123 args = parser.parse_args()
124 main(args)
void print(char *figname, TCanvas *c1)
int main()
Definition hello.cxx:18