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
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
20 filtered_segment_key =
"MuonSegmentsFromR4MlConnected"
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.")
31 flags.Exec.DebugMessageComponents = [
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",
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
49 flags.AthOnnx.ExecutionProvider = OnnxRuntimeType.CPU
51 flags, cfg = setupGeoR4TestCfg(args, flags)
53 if not args.skipTrackTester:
54 cfg.merge(setupHistSvcCfg(flags, outFile=args.outRootFile,
55 outStream=
"MuonTrackTester"))
57 output_level = 1
if args.athenaDebug
else 3
60 from MuonInference.InferenceConfig
import GraphBucketFilterToolCfg, GraphInferenceAlgCfg
61 bucketTool = cfg.popToolsAndMerge(
62 GraphBucketFilterToolCfg(
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,
71 cfg.merge(GraphInferenceAlgCfg(flags, InferenceTools=[bucketTool]))
73 from MuonConfig.ReconstructionConfigR4
import MuonReconstructionConfig
76 cfg.getEventAlgo(
"MuonEtaHoughTransformAlg").SpacePointContainer =
"FilteredMlBuckets"
78 if run_edge_classifier:
79 from MuonInference.InferenceConfig
import SegmentEdgeInferenceAlgCfg
80 edge_classifier_kwargs = {
81 "ModelPath": args.edgeModel,
83 "FilteredMlBuckets" if run_bucket_filter
else "MuonSpacePoints"
87 "MaxSegmentsPerBucket": args.maxSegmentsPerBucket,
88 "MaxEdgesPerNodeBeforeInference": args.maxEdgesBeforeInference,
89 "MaxEdgesPerTargetChamberBeforeInference": (
90 args.maxEdgesPerTargetChamber
92 "DropSameChamberEdgesBeforeInference": (
93 not args.keepSameChamberEdgesBeforeInference
95 "DropIsolatedNodesBeforeInference": (
96 not args.keepIsolatedNodesBeforeInference
98 "EnableTruthDiagnostics": args.truthDiagnostics,
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,
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))
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.")
126 ms_track_finder = cfg.getEventAlgo(
"MSTrackFinderAlg")
127 ms_track_finder.OutputLevel = output_level
129 ms_track_finder.FittingTool.OutputLevel = 3
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(
136 BaselineSeeder=baseline_seeder,
137 SegmentContainer=
"MuonSegmentsFromR4",
138 CandidateDecoration=
"mlTrackComponent",
139 MinSegmentsPerCandidate=args.minSegmentsPerComponent,
140 MaxSegmentsPerCandidate=args.maxSegmentsPerSeedCandidate,
141 MinCosConsistency=args.minSeedCosConsistency,
143 elif filter_segment_container:
144 ms_track_finder.SeedingTool.SegmentContainer = filtered_segment_key
146 if not args.skipTrackTester:
147 from MuonTrackFindingTest.MsTrackFindingTester
import MsTrackTesterCfg
148 cfg.merge(MsTrackTesterCfg(flags, scheduleLegacy=
False, outFile=args.outRootFile))
150 if args.enableRecoChainTester:
151 from MuonTrackFindingAlgs.TrackFindingConfig
import MuonActsToTrkConvCfg
152 cfg.merge(MuonActsToTrkConvCfg(flags,
153 ACTSTracksLocation=
"MsTracks",
154 TracksLocation=
"MsTracksConv"))
156 from xAODTrackingCnv.xAODTrackingCnvConfig
import MuonStandaloneTrackParticleCnvAlgCfg
157 cfg.merge(MuonStandaloneTrackParticleCnvAlgCfg(flags,
158 name=
"MuonXAODParticleConvR4",
159 TrackContainerName=
"MsTracksConv",
160 xAODTrackParticlesFromTracksContainerName=
"MuonSpectrometerTrackParticlesR4"))
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"))
171 from MuonPatternRecognitionTest.PatternTestConfig
import MuonRecoChainTesterCfg
172 cfg.merge(MuonRecoChainTesterCfg(flags,
173 LegacySegmentKey=
"MuonSegmentsFromR4",
174 SegmentFromR4HoughKey=
"",
175 R4SegmentKey=
"MuonSegmentsFromR4",
176 LegacyTrackKey=
"MuonSpectrometerTrackParticlesR4",
178 TrackKeyR4=
"MuonSpectrometerTrackParticlesR4"))
180 cfg.printConfig(withDetails=
True, summariseProps=
True)