181def DisplacedVertexInferenceAlgCfg(flags, name="DisplacedVertexInferenceAlg", **kwargs):
182 """Configure a runnable event-level DisplacedVertex inference algorithm."""
183 result = ComponentAccumulator()
184 do_calo_tower_build = kwargs.pop("DoCaloTowerBuild", True)
185 do_ml_bucket_filter = kwargs.pop("DoMLBucketFilter", True)
186 bucket_model_path = kwargs.pop("BucketModelPath", None)
187 bucket_threshold = kwargs.pop("BucketThreshold", None)
188 filtered_bucket_key = kwargs.pop("FilteredBucketKey", "FilteredMlBuckets")
189 use_filtered_buckets_for_dv_graph = kwargs.pop("UseFilteredBucketsForDVGraph", False)
190 alg_output_level = kwargs.get("OutputLevel", None)
191 tool_kwargs = {}
192 for key in (
193 "ModelPath",
194 "InputNodeName",
195 "InputEdgeIndexName",
196 "InputEdgeAttrName",
197 "InputNMuonNodesName",
198 "OutputName",
199 "SingleOutputMode",
200 "SegmentKey",
201 "SpacePointKeys",
202 "UseBucketSegmentSelection",
203 "TowerContainerKey",
204 "MinTowerEnergyMeV",
205 "MaxTowerSegmentDR",
206 "CaloRMaxMm",
207 "CaloZMaxMm",
208 "SectorModulo",
209 "RequireEdges",
210 "MaxEdges",
211 "FallbackToAllSegments",
212 "DebugDumpFirstNNodes",
213 "DebugDumpFirstNEdges",
214 "SpacePointKeys",
215 "UseBucketSegmentSelection",
216 "OutputLevel",
217 ):
218 if key in kwargs:
219 tool_kwargs[key] = kwargs.pop(key)
220
221 if isinstance(kwargs.get("InferenceTool"), dict):
222 tool_kwargs.update(kwargs.pop("InferenceTool"))
223
224 tower_key = tool_kwargs.get("TowerContainerKey", "CombinedTower")
225 if do_calo_tower_build and tower_key:
226 result.merge(DisplacedVertexCaloTowerCfg(flags))
227
228 if do_ml_bucket_filter:
229 bucket_filter_kwargs = {
230 "WriteSpacePointKey": filtered_bucket_key,
231 "ReadSpacePoints": "MuonSpacePoints",
232 }
233 if bucket_model_path is not None:
234 bucket_filter_kwargs["ModelPath"] = bucket_model_path
235 if bucket_threshold is not None:
236 bucket_filter_kwargs["ScoreThreshold"] = bucket_threshold
237 bucket_tool = result.popToolsAndMerge(
238 GraphBucketFilterToolCfg(flags, **bucket_filter_kwargs)
239 )
240 result.merge(
241 GraphInferenceAlgCfg(
242 flags,
243 name="DVBucketPrefilterAlg",
244 InferenceTools=[bucket_tool],
245 )
246 )
247 if use_filtered_buckets_for_dv_graph:
248 tool_kwargs.setdefault("SpacePointKeys", [filtered_bucket_key])
249 tool_kwargs.setdefault("UseBucketSegmentSelection", True)
250
251 if "InferenceTool" not in kwargs:
252 kwargs["InferenceTool"] = result.popToolsAndMerge(
253 DisplacedVertexInferenceToolCfg(flags, **tool_kwargs)
254 )
255
256 if alg_output_level is not None:
257 kwargs["OutputLevel"] = alg_output_level
258
259 kwargs.setdefault("ScoreDecoration", "EventInfo.dv_score")
260 kwargs.setdefault("RawOutputDecoration", "EventInfo.dv_rawOutput")
261 kwargs.setdefault("PassDecoration", "EventInfo.dv_pass")
262 kwargs.setdefault("NNodesDecoration", "EventInfo.dv_nNodes")
263 kwargs.setdefault("NEdgesDecoration", "EventInfo.dv_nEdges")
264 alg = CompFactory.MuonML.DVInferenceAlg(name=name, **kwargs)
265 result.addEventAlgo(alg, primary=True)
266 return result