ATLAS Offline Software
Loading...
Searching...
No Matches
muonEdgeRecoChain.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# End-to-end test: bucket filter -> segment edge inference -> ML-assisted track seeding.
4
5def main(args):
6 from MuonGeoModelTestR4.testGeoModel import setupGeoR4TestCfg
7 from MuonConfig.MuonConfigUtils import executeTest, setupHistSvcCfg
8 from AthenaConfiguration.AllConfigFlags import initConfigFlags
9 flags = initConfigFlags()
10 flags.PerfMon.doFullMonMT = not args.noPerfMon
11 flags.PerfMon.OutputJSON = "perfmonmt_MuonR4Reco.json"
12 flags.Trigger.Muon.useNewRegionSelector = False
13
14 run_bucket_filter = args.enableBucketFilter and not args.skip_onnx
15 run_edge_classifier = args.enableEdgeClassifier and not args.skip_onnx
16 run_ml_seeder = args.useMlSeeder and not args.skip_onnx
17 filter_segment_container = (
18 args.filterSegmentsWithoutMlConnections and run_edge_classifier
19 )
20 filtered_segment_key = "MuonSegmentsFromR4MlConnected"
21
22 if args.skip_onnx and (args.enableBucketFilter or args.enableEdgeClassifier):
23 print("INFO: --skip-onnx requested. Disabling bucket filter and edge classifier inference stages.")
24 if args.skip_onnx and args.useMlSeeder:
25 print("INFO: --skip-onnx requested. Switching to the standard seeder for the non-ONNX baseline.")
26 if args.filterSegmentsWithoutMlConnections and not run_edge_classifier:
27 print("WARNING: --filterSegmentsWithoutMlConnections requires the edge "
28 "classifier and will be ignored.")
29
30 if args.athenaDebug:
31 flags.Exec.DebugMessageComponents = [
32 "GraphInferenceAlg",
33 "GraphInferenceAlg.GraphBucketFilterTool",
34 "GraphInferenceAlg.GraphBucketFilterTool.OnnxRuntimeSessionToolCPU",
35 "GraphInferenceAlg.GraphBucketFilterTool.OnnxRuntimeSessionToolCUDA",
36 "SegmentEdgeInferenceAlg",
37 "SegmentEdgeInferenceAlg.SegmentEdgeClassifierTool",
38 "SegmentEdgeInferenceAlg.SegmentEdgeClassifierTool.OnnxRuntimeSessionToolCPU",
39 "SegmentEdgeInferenceAlg.SegmentEdgeClassifierTool.OnnxRuntimeSessionToolCUDA",
40 "MSTrackFinderAlg.MlMsTrackSeeder",
41 ]
42
43 from AthOnnxComps.OnnxRuntimeFlags import OnnxRuntimeType
44 if run_bucket_filter or run_edge_classifier:
45 flags.AthOnnx.ExecutionProvider = (
46 OnnxRuntimeType.CPU if args.use_cpu else OnnxRuntimeType.CUDA
47 )
48 else:
49 flags.AthOnnx.ExecutionProvider = OnnxRuntimeType.CPU
50
51 flags, cfg = setupGeoR4TestCfg(args, flags)
52
53 if not args.skipTrackTester:
54 cfg.merge(setupHistSvcCfg(flags, outFile=args.outRootFile,
55 outStream="MuonTrackTester"))
56
57 output_level = 1 if args.athenaDebug else 3
58
59 if run_bucket_filter:
60 from MuonInference.InferenceConfig import GraphBucketFilterToolCfg, GraphInferenceAlgCfg
61 bucketTool = cfg.popToolsAndMerge(
62 GraphBucketFilterToolCfg(
63 flags,
64 ModelPath=args.bucket_model_path,
65 ScoreThreshold=args.score_threshold,
66 OutputName=args.output_name,
67 SingleOutputMode=args.single_output_mode,
68 OutputLevel=output_level,
69 )
70 )
71 cfg.merge(GraphInferenceAlgCfg(flags, InferenceTools=[bucketTool]))
72
73 from MuonConfig.ReconstructionConfigR4 import MuonReconstructionConfig
74 cfg.merge(MuonReconstructionConfig(flags))
75 if run_bucket_filter:
76 cfg.getEventAlgo("MuonEtaHoughTransformAlg").SpacePointContainer = "FilteredMlBuckets"
77
78 if run_edge_classifier:
79 from MuonInference.InferenceConfig import SegmentEdgeInferenceAlgCfg
80 edge_classifier_kwargs = {
81 "ModelPath": args.edgeModel,
82 "ReadSpacePoints": (
83 "FilteredMlBuckets" if run_bucket_filter else "MuonSpacePoints"
84 ),
85 # These cuts run before ONNX. MaxEdgesPerSegment below acts only
86 # after all model scores have already been computed.
87 "MaxSegmentsPerBucket": args.maxSegmentsPerBucket,
88 "MaxEdgesPerNodeBeforeInference": args.maxEdgesBeforeInference,
89 "MaxEdgesPerTargetChamberBeforeInference": (
90 args.maxEdgesPerTargetChamber
91 ),
92 "DropSameChamberEdgesBeforeInference": (
93 not args.keepSameChamberEdgesBeforeInference
94 ),
95 "DropIsolatedNodesBeforeInference": (
96 not args.keepIsolatedNodesBeforeInference
97 ),
98 "EnableTruthDiagnostics": args.truthDiagnostics,
99 }
100 if args.maxDeltaThetaDeg is not None:
101 edge_classifier_kwargs["MaxDeltaThetaDeg"] = args.maxDeltaThetaDeg
102 edge_inference_kwargs = {
103 "EdgeClassifierTool": edge_classifier_kwargs,
104 "PairGateDecoration": "MuonSegmentsFromR4.mlTrackComponent",
105 "PairGateThreshold": args.edgeThreshold,
106 "MaxEdgesPerNode": args.maxEdgesPerSegment,
107 "UseDegreeCappedComponents": args.useDegreeCappedMlComponents,
108 "RequireMutualTopKEdges": not args.allowOneSidedMlEdges,
109 "RecoverOrphanNodes": not args.disableOrphanRecovery,
110 "SeedAnchorsPerComponent": args.seedAnchorsPerComponent,
111 "AnchorInnermostLayer": args.anchorInnermostLayer,
112 "MinSegmentsPerComponent": args.minSegmentsPerComponent,
113 "KeepBestSegmentPerChamber": not args.keepAllSegmentsPerChamber,
114 "OutputLevel": output_level,
115 }
116 if filter_segment_container:
117 edge_inference_kwargs["FilteredSegmentKey"] = filtered_segment_key
118 if args.truthDiagnostics and not flags.Input.isMC:
119 print("WARNING: --truthDiagnostics requested on non-MC input.")
120 cfg.merge(SegmentEdgeInferenceAlgCfg(flags, **edge_inference_kwargs))
121
122 if run_ml_seeder and not run_edge_classifier:
123 print("WARNING: ML seeder enabled while edge classifier is disabled."
124 " The decoration 'mlTrackComponent' may be missing.")
125
126 ms_track_finder = cfg.getEventAlgo("MSTrackFinderAlg")
127 ms_track_finder.OutputLevel = output_level
128 if args.athenaDebug:
129 ms_track_finder.FittingTool.OutputLevel = 3
130 if run_ml_seeder:
131 from MuonTrackFindingAlgs.TrackFindingConfig import MsTrackSeedingToolCfg
132 from AthenaConfiguration.ComponentFactory import CompFactory
133 baseline_seeder = cfg.popToolsAndMerge(MsTrackSeedingToolCfg(flags))
134 ms_track_finder.SeedingTool = CompFactory.MuonR4.MlMsTrackSeeder(
135 "MlMsTrackSeeder",
136 BaselineSeeder=baseline_seeder,
137 SegmentContainer="MuonSegmentsFromR4",
138 CandidateDecoration="mlTrackComponent",
139 MinSegmentsPerCandidate=args.minSegmentsPerComponent,
140 MaxSegmentsPerCandidate=args.maxSegmentsPerSeedCandidate,
141 MinCosConsistency=args.minSeedCosConsistency,
142 )
143 elif filter_segment_container:
144 ms_track_finder.SeedingTool.SegmentContainer = filtered_segment_key
145
146 if not args.skipTrackTester:
147 from MuonTrackFindingTest.MsTrackFindingTester import MsTrackTesterCfg
148 cfg.merge(MsTrackTesterCfg(flags, scheduleLegacy=False, outFile=args.outRootFile))
149
150 if args.enableRecoChainTester:
151 from MuonTrackFindingAlgs.TrackFindingConfig import MuonActsToTrkConvCfg
152 cfg.merge(MuonActsToTrkConvCfg(flags,
153 ACTSTracksLocation="MsTracks",
154 TracksLocation="MsTracksConv"))
155
156 from xAODTrackingCnv.xAODTrackingCnvConfig import MuonStandaloneTrackParticleCnvAlgCfg
157 cfg.merge(MuonStandaloneTrackParticleCnvAlgCfg(flags,
158 name="MuonXAODParticleConvR4",
159 TrackContainerName="MsTracksConv",
160 xAODTrackParticlesFromTracksContainerName="MuonSpectrometerTrackParticlesR4"))
161
162 if flags.Input.isMC:
163 from MuonTruthAlgsR4.MuonTruthAlgsConfig import RecoSegmentTruthAssocCfg, TrackToTruthPartAssocCfg
164 cfg.merge(RecoSegmentTruthAssocCfg(flags,
165 name="MuonSegmentsFromR4TruthMatching",
166 SegmentKey="MuonSegmentsFromR4"))
167 cfg.merge(TrackToTruthPartAssocCfg(flags,
168 name="TrackToTruthMuonSpectrometerTrackParticlesR4",
169 TrackCollection="MuonSpectrometerTrackParticlesR4"))
170
171 from MuonPatternRecognitionTest.PatternTestConfig import MuonRecoChainTesterCfg
172 cfg.merge(MuonRecoChainTesterCfg(flags,
173 LegacySegmentKey="MuonSegmentsFromR4",
174 SegmentFromR4HoughKey="",
175 R4SegmentKey="MuonSegmentsFromR4",
176 LegacyTrackKey="MuonSpectrometerTrackParticlesR4",
177 TrackKeyHoughR4="",
178 TrackKeyR4="MuonSpectrometerTrackParticlesR4"))
179
180 cfg.printConfig(withDetails=True, summariseProps=True)
181 executeTest(cfg)
182
183if __name__ == "__main__":
184 from MuonGeoModelTestR4.testGeoModel import SetupArgParser, MuonPhaseIITestDefaults
185 parser = SetupArgParser()
186 parser.set_defaults(nEvents=-1)
187 parser.set_defaults(inputFile=MuonPhaseIITestDefaults.HITS_PG_R3)
188 parser.set_defaults(outRootFile="EdgeRecoChain.root")
189 from MuonInference.InferenceConfig import (
190 DEFAULT_BUCKET_MODEL_PATH,
191 DEFAULT_BUCKET_SCORE_THRESHOLD,
192 DEFAULT_BUCKET_SINGLE_OUTPUT_MODE,
193 )
194 parser.add_argument("--bucketModel", "--bucket-model-path", dest="bucket_model_path", default=DEFAULT_BUCKET_MODEL_PATH)
195 parser.add_argument("--bucketThreshold", "--score-threshold", dest="score_threshold", type=float, default=DEFAULT_BUCKET_SCORE_THRESHOLD)
196 parser.add_argument("--output-name", default="logits", dest="output_name",
197 help="Bucket filter ONNX output tensor name")
198 score_mode = parser.add_mutually_exclusive_group()
199 score_mode.add_argument("--single-output-mode", choices=("logit", "prob"), default=DEFAULT_BUCKET_SINGLE_OUTPUT_MODE, dest="single_output_mode",
200 help="Scalar ONNX-output interpretation. 'logit' applies sigmoid before thresholding.")
201 score_mode.add_argument("--is-logit", action="store_const", const="logit", dest="single_output_mode",
202 help="Alias for --single-output-mode logit.")
203 score_mode.add_argument("--is-prob", action="store_const", const="prob", dest="single_output_mode",
204 help="Alias for --single-output-mode prob.")
205 parser.add_argument("--edgeModel")
206 parser.add_argument("--maxDeltaThetaDeg", type=float, default=None,
207 help="Override the edge-building opening-angle gate (deg); 180 disables it")
208 parser.add_argument("--athenaDebug", action="store_true",
209 help="Enable Athena DEBUG verbosity for inference and seeding components")
210 parser.add_argument("--truthDiagnostics", action="store_true", default=False,
211 help="MC-only: DEBUG-log a truth-vs-background breakdown of the segment-edge selection.")
212 parser.add_argument("--noPerfMon", default=False, action="store_true",
213 help="Disable performance monitoring")
214 parser.add_argument("--edgeThreshold", type=float, default=0.975,
215 help="Minimum high-confidence edge probability used to form ML track components")
216 parser.add_argument("--maxEdgesPerSegment", type=int, default=2,
217 help="Keep this many highest-score neighbours per segment in the ML path graph (default: 2)")
218 parser.add_argument("--useDegreeCappedMlComponents", action="store_true", default=False,
219 help="Use a global greedy degree cap instead of mutual top-K path extraction")
220 parser.add_argument("--allowOneSidedMlEdges", "--allowBranchingMlComponents",
221 dest="allowOneSidedMlEdges", action="store_true", default=False,
222 help="Keep an edge selected by only one endpoint; use only for validation/recovery")
223 parser.add_argument("--disableOrphanRecovery", action="store_true", default=False,
224 help="Disable bounded one-sided recovery for nodes with no mutual top-K ML edge")
225 parser.add_argument("--seedAnchorsPerComponent", type=int, default=0,
226 help="Launch this many ranked ML anchors per component; zero keeps every retained segment (default: 0)")
227 parser.add_argument("--anchorInnermostLayer", action="store_true", default=False,
228 help="Restrict seed anchors to inner segment(s)")
229 parser.add_argument("--minSegmentsPerComponent", type=int, default=2,
230 help="Require this many retained chambers in an ML component before seeding (default: 2)")
231 parser.add_argument("--maxSegmentsPerSeedCandidate", type=int, default=0,
232 help="Discard ML components with more than this many segments before seeding "
233 "(MlMsTrackSeeder.MaxSegmentsPerCandidate); 0 disables the cap (default: 0)")
234 parser.add_argument("--minSeedCosConsistency", type=float, default=0.819,
235 help="Minimum cos(angle) between a segment's direction and its seed anchor's "
236 "direction for the segment to be kept in an ML seed.")
237 parser.add_argument("--maxSegmentsPerBucket", type=int, default=0,
238 help="Keep at most this many best duplicate segments in each "
239 "(sector,chamber,eta) bucket before ONNX; 0 keeps all.")
240 parser.add_argument("--maxEdgesBeforeInference", type=int, default=6,
241 help="Each node nominates this many geometrically best "
242 "undirected edges before ONNX; 0 keeps all")
243 parser.add_argument("--maxEdgesPerTargetChamber", type=int, default=1,
244 help="Keep at most this many geometrical neighbours from a "
245 "single target chamber for each node before ONNX; 0 keeps all")
246 parser.add_argument("--keepSameChamberEdgesBeforeInference",
247 action="store_true", default=False,
248 help="Keep same-chamber edges in the ONNX input graph. "
249 "Disabled by default because direct ML seeding keeps "
250 "only one segment per chamber.")
251 parser.add_argument("--keepIsolatedNodesBeforeInference",
252 action="store_true", default=False,
253 help="Keep nodes with no retained pre-ONNX edge. Disabled by "
254 "default because isolated nodes cannot contribute to edge scores.")
255 chamber_representatives = parser.add_mutually_exclusive_group()
256 chamber_representatives.add_argument("--keepAllSegmentsPerChamber", dest="keepAllSegmentsPerChamber",
257 action="store_true", default=True, help="Keep all ML component segments from a chamber")
258 chamber_representatives.add_argument(
259 "--keepBestSegmentPerChamber", dest="keepAllSegmentsPerChamber", action="store_false",
260 help="Keep only the highest-ranked segment per chamber")
261 parser.add_argument("--enableRecoChainTester", action="store_true", default=False,
262 help="Enable MuonRecoChainTester (can crash for some custom chains)")
263 parser.add_argument("--skipTrackTester", action="store_true", default=False,
264 help="Do not write the MsTrackValidTest validation tree")
265 parser.add_argument("--enableBucketFilter", dest="enableBucketFilter", action="store_true", default=True,
266 help="Enable ML bucket filtering stage")
267 parser.add_argument("--disableBucketFilter", dest="enableBucketFilter", action="store_false",
268 help="Disable ML bucket filtering stage")
269 parser.add_argument("--enableEdgeClassifier", dest="enableEdgeClassifier", action="store_true", default=True,
270 help="Enable segment-edge classifier stage")
271 parser.add_argument("--disableEdgeClassifier", dest="enableEdgeClassifier", action="store_false",
272 help="Disable segment-edge classifier stage")
273 parser.add_argument("--filterSegmentsWithoutMlConnections",
274 "--filter-segments-without-ml-connections",
275 action="store_true", default=False,
276 help="Pass MSTrackFinderAlg a VIEW containing only "
277 "segments incident to a selected ML edge")
278 parser.add_argument("--useMlSeeder", dest="useMlSeeder", action="store_true", default=True,
279 help="Use new ML-assisted seeder (default)")
280 parser.add_argument("--useStandardSeeder", dest="useMlSeeder", action="store_false",
281 help="Use the standard seeder")
282 parser.add_argument("--use-cpu", action="store_true", default=False,
283 help="Force CPU for ONNX inference")
284 parser.add_argument("--skip-onnx", action="store_true", default=False,
285 help="Skip all ONNX inference stages (bucket filter + edge classifier)")
286
287 args = parser.parse_args()
288
289 if not args.skip_onnx and args.enableEdgeClassifier and not args.edgeModel:
290 parser.error("--edgeModel is required when edge classifier is enabled")
291 main(args)
void print(char *figname, TCanvas *c1)
int main()
Definition hello.cxx:18