13 flags, name="DumpObjects", outfile="Dump_GNN4Itk.root", **kwargs):
15 create algorithm which dumps GNN training information to ROOT file
17 acc = ComponentAccumulator()
21 Output=[f
"{name} DATAFILE='{outfile}', OPT='RECREATE'"]
25 kwargs.setdefault(
"NtupleFileName", flags.Tracking.GNN.DumpObjects.NtupleFileName)
26 kwargs.setdefault(
"NtupleTreeName", flags.Tracking.GNN.DumpObjects.NtupleTreeName)
27 kwargs.setdefault(
"rootFile",
True)
29 acc.addEventAlgo(CompFactory.InDet.DumpObjects(name, **kwargs))
33 """Sets up a GNNTrackFinderTool tool and returns it."""
34 acc = ComponentAccumulator()
37 kwargs.setdefault(
"embeddingDim", flags.Tracking.GNN.TrackFinder.embeddingDim)
38 kwargs.setdefault(
"rVal", flags.Tracking.GNN.TrackFinder.rVal)
39 kwargs.setdefault(
"knnVal", flags.Tracking.GNN.TrackFinder.knnVal)
40 kwargs.setdefault(
"filterCut", flags.Tracking.GNN.TrackFinder.filterCut)
41 kwargs.setdefault(
"inputMLModelDir", flags.Tracking.GNN.TrackFinder.inputMLModelDir)
42 kwargs.setdefault(
"ccCut", flags.Tracking.GNN.TrackFinder.ccCut)
43 kwargs.setdefault(
"walkMin", flags.Tracking.GNN.TrackFinder.walkMin)
44 kwargs.setdefault(
"walkMax", flags.Tracking.GNN.TrackFinder.walkMax)
45 kwargs.setdefault(
"EmbeddingFeatureNames", flags.Tracking.GNN.TrackFinder.EmbeddingFeatureNames)
46 kwargs.setdefault(
"EmbeddingFeatureScales", flags.Tracking.GNN.TrackFinder.EmbeddingFeatureScales)
47 kwargs.setdefault(
"FilterFeatureNames", flags.Tracking.GNN.TrackFinder.FilterFeatureNames)
48 kwargs.setdefault(
"FilterFeatureScales", flags.Tracking.GNN.TrackFinder.FilterFeatureScales)
49 kwargs.setdefault(
"GNNFeatureNames", flags.Tracking.GNN.TrackFinder.GNNFeatureNames)
50 kwargs.setdefault(
"GNNFeatureScales", flags.Tracking.GNN.TrackFinder.GNNFeatureScales)
52 from AthOnnxComps.OnnxRuntimeInferenceConfig
import OnnxRuntimeInferenceToolCfg
53 ort_exe_provider = flags.Tracking.GNN.TrackFinder.ORTExeProvider
54 kwargs.setdefault(
"Embedding", acc.popToolsAndMerge(
55 OnnxRuntimeInferenceToolCfg(flags, str(Path(kwargs[
"inputMLModelDir"]) /
"embedding.onnx"),
56 ort_exe_provider, name=
"Embedding")
58 kwargs.setdefault(
"Filtering", acc.popToolsAndMerge(
59 OnnxRuntimeInferenceToolCfg(flags, str(Path(kwargs[
"inputMLModelDir"]) /
"filtering.onnx"),
60 ort_exe_provider, name=
"Filtering")
62 kwargs.setdefault(
"GNN", acc.popToolsAndMerge(
63 OnnxRuntimeInferenceToolCfg(flags, str(Path(kwargs[
"inputMLModelDir"]) /
"gnn.onnx"),
64 ort_exe_provider, name=
"GNN")
67 acc.setPrivateTools(CompFactory.InDet.SiGNNTrackFinderTool(name, **kwargs))
72 """Sets up a GNNTrackFinderTritonTool tool and returns it."""
73 from AthTritonComps.TritonToolConfig
import TritonToolCfg
75 acc = ComponentAccumulator()
77 kwargs.setdefault(
"TritonTool", acc.popToolsAndMerge(
78 TritonToolCfg(flags, model_name=flags.Tracking.GNN.Triton.model,
79 url=flags.Tracking.GNN.Triton.url,
80 port=flags.Tracking.GNN.Triton.port,
84 kwargs.setdefault(
"FeatureNames", flags.Tracking.GNN.spacepointFeatures)
85 acc.setPrivateTools(CompFactory.InDet.GNNTrackFinderTritonTool(name, **kwargs))
90 """Sets up an ActsGnnModuleMapFinderTool and returns it."""
91 acc = ComponentAccumulator()
93 kwargs.setdefault(
"moduleMapPath", flags.Tracking.GNN.ActsPipeline.moduleMapPath)
94 kwargs.setdefault(
"gnnPath", flags.Tracking.GNN.ActsPipeline.gnnPath)
95 kwargs.setdefault(
"edgeCut", flags.Tracking.GNN.ActsPipeline.edgeCut)
96 kwargs.setdefault(
"numTrtContexts", flags.Tracking.GNN.ActsPipeline.numTrtContexts)
97 kwargs.setdefault(
"minCandidateMeasurements", flags.Tracking.GNN.ActsPipeline.minCandidateMeasurements)
98 kwargs.setdefault(
"useEdgeLayerConnector", flags.Tracking.GNN.ActsPipeline.useEdgeLayerConnector)
99 kwargs.setdefault(
"elcMaxHitsPerTrack", flags.Tracking.GNN.ActsPipeline.elcMaxHitsPerTrack)
102 acc.setPrivateTools(CompFactory.InDet.ActsGnnModuleMapFinderTool(name, **kwargs))
144 """Sets up a GNNTrackMaker algorithm and returns it."""
146 acc = ComponentAccumulator()
150 kwargs.setdefault(
"SeedFitterTool", SeedFitterTool)
152 from TrkConfig.CommonTrackFitterConfig
import ITkTrackFitterCfg
153 InDetTrackFitter = acc.popToolsAndMerge(ITkTrackFitterCfg(flags))
154 kwargs.setdefault(
"TrackFitter", InDetTrackFitter)
156 if "TrackSummaryTool" not in kwargs:
157 from TrkConfig.TrkTrackSummaryToolConfig
import ITkTrackSummaryToolCfg
160 "TrackSummaryTool", acc.popToolsAndMerge(ITkTrackSummaryToolCfg(flags))
163 if flags.Tracking.GNN.ToolType == GNNTrackFinderToolType.TrackFinder:
165 kwargs.setdefault(
"GNNTrackFinderTool", InDetGNNTrackFinderTool)
166 kwargs.setdefault(
"GNNTrackReaderTool",
None)
167 elif flags.Tracking.GNN.ToolType == GNNTrackFinderToolType.TrackReader:
169 kwargs.setdefault(
"GNNTrackReaderTool", InDetGNNTrackReader)
170 kwargs.setdefault(
"GNNTrackFinderTool",
None)
171 elif flags.Tracking.GNN.ToolType == GNNTrackFinderToolType.Triton:
173 kwargs.setdefault(
"GNNTrackReaderTool",
None)
174 kwargs.setdefault(
"GNNTrackFinderTool", InDetGNNTrackFinderTool)
175 elif flags.Tracking.GNN.ToolType == GNNTrackFinderToolType.ActsPipeline:
177 kwargs.setdefault(
"GNNTrackReaderTool",
None)
178 kwargs.setdefault(
"GNNTrackFinderTool", InDetGNNTrackFinderTool)
180 raise RuntimeError(
"GNNTrackFinder or GNNTrackReader must be enabled!")
182 kwargs.setdefault(
"areInputClusters", flags.Tracking.GNN.useClusterTracks)
183 kwargs.setdefault(
"doRecoTrackCuts", flags.Tracking.GNN.doRecoTrackCuts)
184 kwargs.setdefault(
"saveEdgeScore", flags.Tracking.GNN.ActsPipeline.saveEdgeScore)
187 if "InDetEtaDependentCutSvc" not in kwargs:
188 from InDetConfig.InDetEtaDependentCutsConfig
import ITkEtaDependentCutsSvcCfg
189 acc.merge(ITkEtaDependentCutsSvcCfg(flags))
190 kwargs.setdefault(
"InDetEtaDependentCutsSvc", acc.getService(
"ITkEtaDependentCutsSvc"+flags.Tracking.ActiveConfig.extension))
192 kwargs.setdefault(
"minClusters", flags.Tracking.GNN.minClusters)
193 kwargs.setdefault(
"pTmin", flags.Tracking.GNN.pTmin)
194 kwargs.setdefault(
"etamax", flags.Tracking.GNN.etamax)
195 kwargs.setdefault(
"minPixelClusters", flags.Tracking.GNN.minPixelClusters)
196 kwargs.setdefault(
"minStripClusters", flags.Tracking.GNN.minStripClusters)
198 acc.addEventAlgo(CompFactory.InDet.SiSPGNNTrackMaker(name, **kwargs))
202 """Sets up a GNN for seeding algorithm and returns it."""
203 acc = ComponentAccumulator()
205 from InDetConfig.SiCombinatorialTrackFinderToolConfig
import SiDetElementBoundaryLinksCondAlg_xk_ITkPixel_Cfg, SiDetElementBoundaryLinksCondAlg_xk_ITkStrip_Cfg
206 acc.merge(SiDetElementBoundaryLinksCondAlg_xk_ITkPixel_Cfg(flags))
207 acc.merge(SiDetElementBoundaryLinksCondAlg_xk_ITkStrip_Cfg(flags))
210 from MagFieldServices.MagFieldServicesConfig
import (
211 AtlasFieldCacheCondAlgCfg)
212 acc.merge(AtlasFieldCacheCondAlgCfg(flags))
214 from TrkConfig.TrkRIO_OnTrackCreatorConfig
import ITkRotCreatorCfg
215 ITkRotCreator = acc.popToolsAndMerge(ITkRotCreatorCfg(
216 flags, name=
"ITkRotCreator"+flags.Tracking.ActiveConfig.extension))
217 acc.addPublicTool(ITkRotCreator)
218 kwargs.setdefault(
"RIOonTrackTool", ITkRotCreator)
220 from TrkConfig.TrkExRungeKuttaPropagatorConfig
import (
221 RungeKuttaPropagatorCfg)
222 ITkPatternPropagator = acc.popToolsAndMerge(
223 RungeKuttaPropagatorCfg(flags, name=
"ITkPatternPropagator"))
224 acc.addPublicTool(ITkPatternPropagator)
225 kwargs.setdefault(
"PropagatorTool", ITkPatternPropagator)
227 from TrkConfig.TrkMeasurementUpdatorConfig
import KalmanUpdator_xkCfg
228 ITkPatternUpdator = acc.popToolsAndMerge(
229 KalmanUpdator_xkCfg(flags, name=
"ITkPatternUpdator"))
230 acc.addPublicTool(ITkPatternUpdator)
231 kwargs.setdefault(
"UpdatorTool", ITkPatternUpdator)
233 from InDetConfig.InDetBoundaryCheckToolConfig
import ITkBoundaryCheckToolCfg
234 kwargs.setdefault(
"BoundaryCheckTool", acc.popToolsAndMerge(
235 ITkBoundaryCheckToolCfg(flags)))
237 from PixelConditionsTools.ITkPixelConditionsSummaryConfig
import (
238 ITkPixelConditionsSummaryCfg)
239 kwargs.setdefault(
"PixelSummaryTool", acc.popToolsAndMerge(
240 ITkPixelConditionsSummaryCfg(flags)))
242 from SCT_ConditionsTools.ITkStripConditionsToolsConfig
import (
243 ITkStripConditionsSummaryToolCfg)
244 kwargs.setdefault(
"StripSummaryTool", acc.popToolsAndMerge(
245 ITkStripConditionsSummaryToolCfg(flags)))
247 if flags.Tracking.GNN.useTrackFinder:
249 kwargs.setdefault(
"GNNTrackReaderTool",
None)
250 elif flags.Tracking.GNN.useTrackReader:
252 kwargs.setdefault(
"GNNTrackFinderTool",
None)
254 raise RuntimeError(
"GNNTrackFinder or GNNTrackReader must be enabled!")
256 kwargs.setdefault(
"SeedFitterTool", acc.popToolsAndMerge(
SeedFitterToolCfg(flags)))
258 from TrkConfig.CommonTrackFitterConfig
import ITkTrackFitterCfg
259 kwargs.setdefault(
"TrackFitter", acc.popToolsAndMerge(ITkTrackFitterCfg(flags)))
261 from InDetConfig.SiDetElementsRoadToolConfig
import ITkSiDetElementsRoadMaker_xkCfg
262 kwargs.setdefault(
"RoadTool", acc.popToolsAndMerge(ITkSiDetElementsRoadMaker_xkCfg(flags)))
266 kwargs.setdefault(
"nClustersMin", flags.Tracking.ActiveConfig.minClusters[0])
267 kwargs.setdefault(
"nWeightedClustersMin", flags.Tracking.ActiveConfig.nWeightedClustersMin[0])
268 kwargs.setdefault(
"nHolesMax", flags.Tracking.ActiveConfig.nHolesMax[0])
269 kwargs.setdefault(
"nHolesGapMax", flags.Tracking.ActiveConfig.nHolesGapMax[0])
271 kwargs.setdefault(
"pTmin", flags.Tracking.ActiveConfig.minPT[0])
272 kwargs.setdefault(
"pTminBrem", flags.Tracking.ActiveConfig.minPTBrem[0])
273 kwargs.setdefault(
"Xi2max", flags.Tracking.ActiveConfig.Xi2max[0])
274 kwargs.setdefault(
"Xi2maxNoAdd", flags.Tracking.ActiveConfig.Xi2maxNoAdd[0])
275 kwargs.setdefault(
"Xi2maxMultiTracks", flags.Tracking.ActiveConfig.Xi2max[0])
276 kwargs.setdefault(
"doMultiTracksProd",
False)
278 acc.addEventAlgo(CompFactory.InDet.GNNSeedingTrackMaker(name, **kwargs))