ATLAS Offline Software
Loading...
Searching...
No Matches
MuonML::SegmentEdgeClassifierTool Class Referencefinal

Runs a segment-level GNN on reconstructed muon segments to classify segment-pair edges as "good" or "background". More...

#include <SegmentEdgeClassifierTool.h>

Inheritance diagram for MuonML::SegmentEdgeClassifierTool:
Collaboration diagram for MuonML::SegmentEdgeClassifierTool:

Public Member Functions

StatusCode initialize () override
 Retrieve the ONNX model and resolve node feature ordering from metadata.
StatusCode finalize () override
 Log a pre-ONNX candidate-edge pruning.
StatusCode runGraphInference (const EventContext &ctx, GraphRawData &graphData) const override
 Not supported by this tool; returns FAILURE.
StatusCode buildGraph (const EventContext &ctx, const xAOD::MuonSegmentContainer &segments, SegmentEdgeGraph &graph) const override
 Build a GNN graph from segments, computing node and edge features and storing the graph structure in graph.
StatusCode classifyEdges (const EventContext &ctx, const SegmentEdgeGraph &graph, std::vector< SegmentEdgeScore > &scoresscores) const override
 Run ONNX inference on graph and populate scores with logit and probability for each edge; called after buildGraph().
bool enableTruthDiagnostics () const override
 Whether this tool was configured to fill SegmentEdgeGraph's truth diagnostics.
StatusCode buildGraph (const EventContext &ctx, GraphRawData &graphData) const
 GNN-style graph builder (features + edges). Kept for tools that want it.
StatusCode runInference (GraphRawData &graphData) const
 Default ONNX run for GNN case: inputs {"features","edge_index"} -> outputs {"logits"}.
 DeclareInterfaceID (ISegmentEdgeClassifierTool, 1, 0)

Protected Member Functions

StatusCode setupModel ()
Ort::Session & model () const
StatusCode buildFeaturesOnly (const EventContext &ctx, GraphRawData &graphData) const
 Build only features (N,6); attaches one tensor in graph.dataTensor[0].
StatusCode buildTransformerInputs (const EventContext &ctx, GraphRawData &graphData) const
 Build Transformer inputs: features [1,S,6] and pad_mask [1,S] (False = valid), as tensors 0 and 1.
StatusCode runNamedInference (GraphRawData &graphData, const std::vector< const char * > &inputNames, const std::vector< const char * > &outputNames) const
 Generic named inference, for tools with different I/O conventions.

Static Protected Member Functions

static std::string trimFeatureToken (std::string s)
static std::vector< std::string > parseFeatureNames (const std::string &raw)

Protected Attributes

SG::ReadHandleKey< MuonR4::SpacePointContainer > m_readKey {this, "ReadSpacePoints", "MuonSpacePoints"}
ActsTrk::GeoContextReadKey_t m_geoCtxKey {this, "AlignmentKey", "ActsAlignment", "cond handle key"}
Gaudi::Property< int > m_minLayers {this, "MinLayersValid", 3}
Gaudi::Property< int > m_maxChamberDelta {this, "MaxChamberDelta", 13}
Gaudi::Property< int > m_maxSectorDelta {this, "MaxSectorDelta", 1}
Gaudi::Property< double > m_maxDistXY {this, "MaxDistXY", 6800.0}
Gaudi::Property< double > m_maxAbsDz {this, "MaxAbsDz", 15000.0}
Gaudi::Property< unsigned int > m_debugDumpFirstNNodes {this, "DebugDumpFirstNNodes", 5}
Gaudi::Property< unsigned int > m_debugDumpFirstNEdges {this, "DebugDumpFirstNEdges", 12}
Gaudi::Property< bool > m_validateEdges {this, "ValidateEdges", true}
Gaudi::Property< bool > m_sanitizeNonFinitePredictions
bool m_isCuda {false}
int m_cudaDeviceId {0}

Static Protected Attributes

static constexpr std::size_t kBucketFeatureCount = 6
static constexpr std::size_t kNodeFeatureCount = 10
static constexpr std::size_t kEdgeFeatureCount = 7
static constexpr std::array< std::string_view, kNodeFeatureCount > kDefaultNodeFeatureNames

Private Member Functions

StatusCode dumpDebugEvent (const EventContext &ctx, const SegmentEdgeGraph &graph, const std::vector< SegmentEdgeScore > &scoresscores) const
void fillTruthDiagnostics (const xAOD::MuonSegmentContainer &segments, const std::unordered_set< const xAOD::MuonSegment * > &bucketRetained, SegmentEdgeGraph &graph) const
 MC-only diagnostics: record, for every input segment and every pair of input segments sharing a truth particle, why it did or did not reach ONNX.

Private Attributes

Gaudi::Property< float > m_maxDeltaThetaDeg {this, "MaxDeltaThetaDeg", 35.f}
Gaudi::Property< int > m_maxDeltaSector {this, "MaxDeltaSector", 1}
Gaudi::Property< int > m_sectorModulo
Gaudi::Property< unsigned int > m_maxSegmentsPerBucket
Gaudi::Property< unsigned int > m_maxEdgesPerNodeBeforeInference
Gaudi::Property< unsigned int > m_maxEdgesPerTargetChamberBeforeInference
Gaudi::Property< bool > m_dropSameChamberEdgesBeforeInference
Gaudi::Property< bool > m_dropIsolatedNodesBeforeInference
Gaudi::Property< std::string > m_inputNodeName {this, "InputNodeName", "x"}
Gaudi::Property< std::string > m_inputEdgeIndexName {this, "InputEdgeIndexName", "edge_index"}
Gaudi::Property< std::string > m_inputEdgeAttrName {this, "InputEdgeAttrName", "edge_attr"}
Gaudi::Property< std::string > m_outputName {this, "OutputName", "logits"}
Gaudi::Property< std::string > m_debugDumpFile {this, "DebugDumpFile", ""}
Gaudi::Property< unsigned int > m_debugDumpMaxEvents {this, "DebugDumpMaxEvents", 0}
Gaudi::Property< bool > m_enableTruthDiagnostics
SG::ReadDecorHandleKey< xAOD::MuonSegmentContainer > m_truthLinkKey
float m_cosMin {0.f}
std::vector< std::string > m_nodeFeatureNames {}
 Node feature order expected by the model metadata (resolved at initialize).
std::vector< SegmentNodeFeatureId > m_nodeFeatureIds {}
std::mutex m_debugDumpMutex
std::atomic< unsigned int > m_debugDumpEvents {0}
std::atomic< std::size_t > m_sumInputSegments {0}
 Job-summed pre-ONNX pruning counters (see buildGraph()).
std::atomic< std::size_t > m_sumCandidatePairs {0}
std::atomic< std::size_t > m_sumRetainedPairs {0}
std::atomic< std::size_t > m_sumNodesBeforeIsolatedDrop {0}
std::atomic< std::size_t > m_sumNodesAfterIsolatedDrop {0}
ToolHandle< AthOnnx::IOnnxRuntimeSessionTool > m_onnxSessionTool

Detailed Description

Runs a segment-level GNN on reconstructed muon segments to classify segment-pair edges as "good" or "background".

The tool reads a xAOD::MuonSegmentContainer and builds a graph where:

  • Nodes are muon segments, each with 10 features:
    • Position and direction (6 floats)
    • Chamber index, layer count, sector, and segment multiplicity (4 floats)
  • Edges connect all segment pairs within an angular threshold (cos(angle) >= cos(MaxDeltaThetaDeg)) and sector distance, with 7 features:
    • Spatial displacement (3 floats: dx, dy, dz)
    • Distance magnitude (1 float)
    • Angle (dot product, 1 float)
    • Chamber and sector match flags (2 flags)

The tool then runs an ONNX model (typically a GIN or GCN variant) to produce a logit or probability for each edge, enabling downstream algorithms to filter low-quality segment associations and improve reconstruction efficiency.

Key difference from GraphBucketFilterTool: operates at segment (edge) level rather than bucket (node) level, and the interface uses discrete graph structures (SegmentEdgeGraph) rather than tensors for input/output.

Note: runGraphInference() is not supported by this tool; use SegmentEdgeInferenceAlg and the ISegmentEdgeClassifierTool methods instead.

Definition at line 64 of file SegmentEdgeClassifierTool.h.

Member Function Documentation

◆ buildFeaturesOnly()

StatusCode BucketInferenceToolBase::buildFeaturesOnly ( const EventContext & ctx,
GraphRawData & graphData ) const
protectedinherited

Build only features (N,6); attaches one tensor in graph.dataTensor[0].

Definition at line 87 of file BucketInferenceToolBase.cxx.

88 {
89
90 graphData.graph.reset();
91 graphData.srcEdges.clear();
92 graphData.desEdges.clear();
93 graphData.edgeIndexPacked.clear();
94 graphData.featureLeaves.clear();
95 graphData.spacePointsInBucket.clear();
96 graphData.graph = std::make_unique<InferenceGraph>();
97 graphData.graph->dataTensor.reserve(1); // features input; outputs are reserved in runNamedInference()
98
99 const MuonR4::SpacePointContainer* buckets{nullptr};
100 ATH_CHECK(SG::get(buckets, m_readKey, ctx));
101
102 const ActsTrk::GeometryContext* gctx = nullptr;
103 ATH_CHECK(SG::get(gctx, m_geoCtxKey, ctx));
104
105 std::vector<BucketGraphUtils::NodeAux> nodes;
106 BucketGraphUtils::buildNodesAndFeatures(*buckets, *gctx, nodes,
107 graphData.featureLeaves,
108 graphData.spacePointsInBucket); // now int64_t-compatible
109
110 const int64_t numNodes = static_cast<int64_t>(nodes.size());
111 ATH_MSG_DEBUG("Total buckets: " << buckets->size()
112 << " -> nodes (size>0): " << numNodes
113 << " | features.size()=" << graphData.featureLeaves.size());
114
115 if (numNodes == 0) {
116 ATH_MSG_WARNING("No valid buckets found (all have size 0.0). Skipping inference.");
117 return StatusCode::SUCCESS;
118 }
119
120 const int64_t nFeatPerNode = static_cast<int64_t>(kBucketFeatureCount);
121 if (numNodes * nFeatPerNode != static_cast<int64_t>(graphData.featureLeaves.size())) {
122 ATH_MSG_ERROR( "Feature size mismatch: expected " << (numNodes * nFeatPerNode)
123 << " got " << graphData.featureLeaves.size());
124 return StatusCode::FAILURE;
125 }
126
127 Ort::MemoryInfo memInfo = Ort::MemoryInfo::CreateCpu(OrtDeviceAllocator, OrtMemTypeCPU);
128 std::vector<int64_t> featShape{numNodes, nFeatPerNode};
129 graphData.graph->dataTensor.emplace_back(
130 Ort::Value::CreateTensor<float>(memInfo,
131 graphData.featureLeaves.data(),
132 graphData.featureLeaves.size(),
133 featShape.data(),
134 featShape.size()));
135 return StatusCode::SUCCESS;
136}
#define ATH_CHECK
Evaluate an expression and check for errors.
#define ATH_MSG_DEBUG(x,...)
#define ATH_MSG_ERROR(x,...)
#define ATH_MSG_WARNING(x,...)
size_type size() const noexcept
Returns the number of elements in the collection.
ActsTrk::GeoContextReadKey_t m_geoCtxKey
static constexpr std::size_t kBucketFeatureCount
SG::ReadHandleKey< MuonR4::SpacePointContainer > m_readKey
void buildNodesAndFeatures(const MuonR4::SpacePointContainer &buckets, const ActsTrk::GeometryContext &gctx, std::vector< NodeAux > &nodes, std::vector< float > &featuresLeaves, std::vector< int64_t > &spInBucket)
Build nodes + flat features (N,6) and number of SPs per kept bucket.
DataVector< SpacePointBucket > SpacePointContainer
Abrivation of the space point container type.
const T * get(const ReadCondHandleKey< T > &key, const EventContext &ctx)
Convenience function to retrieve an object given a ReadCondHandleKey.
FeatureVec_t featureLeaves
Vector containing all features.
Definition GraphData.h:30
EdgeCounterVec_t edgeIndexPacked
Packed edge index buffer (kept alive for ONNX tensors that reference it) This stores [srcEdges,...
Definition GraphData.h:42
std::unique_ptr< InferenceGraph > graph
Pointer to the graph to be parsed to ONNX.
Definition GraphData.h:46
EdgeCounterVec_t srcEdges
Vector encoding the source index of the.
Definition GraphData.h:32
EdgeCounterVec_t desEdges
Vect.
Definition GraphData.h:34
NodeConnectVec_t spacePointsInBucket
Vector keeping track of how many space points are in each parsed bucket.
Definition GraphData.h:36

◆ buildGraph() [1/2]

StatusCode BucketInferenceToolBase::buildGraph ( const EventContext & ctx,
GraphRawData & graphData ) const
inherited

GNN-style graph builder (features + edges). Kept for tools that want it.

Definition at line 200 of file BucketInferenceToolBase.cxx.

201 {
202
203 graphData.graph.reset();
204 graphData.srcEdges.clear();
205 graphData.desEdges.clear();
206 graphData.featureLeaves.clear();
207 graphData.spacePointsInBucket.clear();
208 graphData.edgeIndexPacked.clear();
209 graphData.graph = std::make_unique<InferenceGraph>();
210 graphData.graph->dataTensor.reserve(2); // features and edge_index inputs; outputs are reserved in runNamedInference()
211
212 const MuonR4::SpacePointContainer* buckets{nullptr};
213 ATH_CHECK(SG::get(buckets, m_readKey, ctx));
214
215 const ActsTrk::GeometryContext* gctx = nullptr;
216 ATH_CHECK(SG::get(gctx, m_geoCtxKey, ctx));
217
218 std::vector<BucketGraphUtils::NodeAux> nodes;
219
220 BucketGraphUtils::buildNodesAndFeatures(*buckets, *gctx, nodes,
221 graphData.featureLeaves,
222 graphData.spacePointsInBucket);
223
224 const int64_t numNodes = static_cast<int64_t>(nodes.size());
225 ATH_MSG_DEBUG("Total buckets: " << buckets->size()
226 << " -> nodes (size>0): " << numNodes
227 << " | features.size()=" << graphData.featureLeaves.size());
228
229 if (numNodes == 0) {
230 ATH_MSG_WARNING("No valid buckets found (all have size 0.0). Skipping graph building.");
231 return StatusCode::SUCCESS;
232 }
233
234 const int64_t nFeatPerNode = static_cast<int64_t>(kBucketFeatureCount);
235 if (numNodes * nFeatPerNode != static_cast<int64_t>(graphData.featureLeaves.size())) {
236 ATH_MSG_ERROR("Feature size mismatch: expected " << (numNodes * nFeatPerNode)
237 << " got " << graphData.featureLeaves.size());
238 return StatusCode::FAILURE;
239 }
240
241 Ort::MemoryInfo memInfo = Ort::MemoryInfo::CreateCpu(OrtDeviceAllocator, OrtMemTypeCPU);
242 std::vector<int64_t> featShape{numNodes, nFeatPerNode};
243 graphData.graph->dataTensor.emplace_back(
244 Ort::Value::CreateTensor<float>(memInfo,
245 graphData.featureLeaves.data(),
246 graphData.featureLeaves.size(),
247 featShape.data(),
248 featShape.size()));
249
256 graphData.srcEdges, graphData.desEdges);
257 if (m_validateEdges) {
258 size_t bad = 0;
259 size_t write = 0;
260 for (size_t k = 0; k < graphData.srcEdges.size(); ++k) {
261 const int64_t u = graphData.srcEdges[k];
262 const int64_t v = graphData.desEdges[k];
263 const bool okU = (u >= 0 && u < numNodes);
264 const bool okV = (v >= 0 && v < numNodes);
265 if (okU && okV) {
266 graphData.srcEdges[write] = u;
267 graphData.desEdges[write] = v;
268 ++write;
269 } else {
270 ++bad;
271 ATH_MSG_DEBUG( "Drop invalid edge " << k << ": (" << u << "->" << v
272 << "), valid node range [0," << (numNodes-1) << "]");
273 }
274 }
275 if (bad) {
276 ATH_MSG_WARNING( "Removed " << bad << " invalid edges out of "
277 << graphData.srcEdges.size());
278 graphData.srcEdges.resize(write);
279 graphData.desEdges.resize(write);
280 }
281 }
282
283 const size_t E = graphData.srcEdges.size();
284
285 if (msgLvl(MSG::DEBUG)) {
286 // DEBUG: Count connections per node
287 ATH_MSG_DEBUG("Edges built: " << E);
288 const size_t dumpE = std::min<std::size_t>(m_debugDumpFirstNEdges.value(), E);
289 for (size_t k = 0; k < dumpE; ++k) {
290 ATH_MSG_DEBUG("EDGE[" << k << "]: "
291 << graphData.srcEdges[k] << " -> "
292 << graphData.desEdges[k]);
293 }
294
295 std::vector<int> nodeConnections(numNodes, 0);
296 for (size_t k = 0; k < graphData.srcEdges.size(); ++k) {
297 const int64_t u = graphData.srcEdges[k];
298 const int64_t v = graphData.desEdges[k];
299 if (u >= 0 && u < numNodes) nodeConnections[u]++;
300 if (v >= 0 && v < numNodes) nodeConnections[v]++;
301 }
302
303 ATH_MSG_DEBUG("=== DEBUGGING: Node Connections (first 10 nodes) ===");
304 const int64_t debugNodeCount = std::min(numNodes, static_cast<int64_t>(10));
305 for (int64_t i = 0; i < debugNodeCount; ++i) {
306 ATH_MSG_DEBUG("Node[" << i << "] connections: " << nodeConnections[i]);
307 }
308 ATH_MSG_DEBUG("=== END DEBUG NODE CONNECTIONS ===");
309
310 ATH_MSG_DEBUG("=== DEBUGGING: Detailed Edge Connections (first 10 nodes) ===");
311 for (int64_t nodeIdx = 0; nodeIdx < debugNodeCount; ++nodeIdx) {
312 std::stringstream connections;
313 connections << "Node[" << nodeIdx << "] connected to: ";
314 bool foundAny = false;
315
316 for (size_t k = 0; k < graphData.srcEdges.size(); ++k) {
317 const int64_t u = graphData.srcEdges[k];
318 const int64_t v = graphData.desEdges[k];
319
320 if (u == nodeIdx) {
321 if (foundAny) connections << ", ";
322 connections << v;
323 foundAny = true;
324 } else if (v == nodeIdx) {
325 if (foundAny) connections << ", ";
326 connections << u;
327 foundAny = true;
328 }
329 }
330
331 if (!foundAny) connections << "none";
332 ATH_MSG_DEBUG(connections.str());
333 }
334 ATH_MSG_DEBUG("=== END DEBUG DETAILED CONNECTIONS ===");
335 }
336
337 nodes = {};
338
339 graphData.edgeIndexPacked.clear();
340 const size_t Efinal = BucketGraphUtils::packEdgeIndex(graphData.srcEdges,
341 graphData.desEdges,
342 graphData.edgeIndexPacked);
343
344 graphData.srcEdges.clear();
345 graphData.desEdges.clear();
346
347 std::vector<int64_t> edgeShape{2, static_cast<int64_t>(Efinal)};
348 graphData.graph->dataTensor.emplace_back(
349 Ort::Value::CreateTensor<int64_t>(memInfo,
350 graphData.edgeIndexPacked.data(),
351 graphData.edgeIndexPacked.size(),
352 edgeShape.data(),
353 edgeShape.size()));
354
355 ATH_MSG_DEBUG("Built sparse bucket graph: N=" << numNodes << ", E=" << Efinal);
356 return StatusCode::SUCCESS;
357}
Gaudi::Property< unsigned int > m_debugDumpFirstNEdges
Gaudi::Property< double > m_maxDistXY
void buildSparseEdges(const std::vector< NodeAux > &nodes, int minLayers, int maxChamberDelta, int maxSectorDelta, double maxDistXY, double maxAbsDz, std::vector< int64_t > &srcEdges, std::vector< int64_t > &dstEdges)
size_t packEdgeIndex(const std::vector< int64_t > &srcEdges, const std::vector< int64_t > &dstEdges, std::vector< int64_t > &edgeIndexPacked)
@ u
Enums for curvilinear frames.
Definition ParamDefs.h:77

◆ buildGraph() [2/2]

StatusCode MuonML::SegmentEdgeClassifierTool::buildGraph ( const EventContext & ctx,
const xAOD::MuonSegmentContainer & segments,
SegmentEdgeGraph & graph ) const
overridevirtual

Build a GNN graph from segments, computing node and edge features and storing the graph structure in graph.

Implements MuonML::ISegmentEdgeClassifierTool.

Definition at line 234 of file SegmentEdgeClassifierTool.cxx.

236 {
237 graph = SegmentEdgeGraph{};
238 graph.segments.reserve(segments.size());
239 graph.nodeFeatures.reserve(segments.size() * kNodeFeatureCount);
240 // Evaluated once per event; short-circuits without touching the message
241 // service unless the property was explicitly enabled.
242 const bool truthDiag = m_enableTruthDiagnostics.value() && msgLvl(MSG::DEBUG);
243
244 /*
245 * Keep the original bucket multiplicity in the node feature even when the
246 * speed configuration retains only the best representatives of a bucket.
247 * This preserves the model's occupancy input while removing duplicate node
248 * and edge work before tensor construction.
249 */
250 std::map<SegmentBucketKey, std::vector<const xAOD::MuonSegment*>>
251 segmentsByBucket;
252 for (const xAOD::MuonSegment* segment : segments) {
253 segmentsByBucket[segmentBucketKey(*segment)].push_back(segment);
254 }
255
256 const InferenceUtils::SegmentQualityOrder betterSegment{};
257
258 std::unordered_set<const xAOD::MuonSegment*> retainedSegments;
259 retainedSegments.reserve(segments.size());
260 for (auto& [_, bucketSegments] : segmentsByBucket) {
261 std::ranges::sort(bucketSegments, betterSegment);
262 const std::size_t nKeep = m_maxSegmentsPerBucket.value() == 0
263 ? bucketSegments.size()
264 : std::min<std::size_t>(
265 bucketSegments.size(),
266 m_maxSegmentsPerBucket.value());
267 retainedSegments.insert(bucketSegments.begin(),
268 bucketSegments.begin() + nKeep);
269 }
270
271 std::vector<Amg::Vector3D> pos;
272 std::vector<Amg::Vector3D> dir;
273 std::vector<BucketSegmentFeatures> bucket;
274 pos.reserve(retainedSegments.size());
275 dir.reserve(retainedSegments.size());
276 bucket.reserve(retainedSegments.size());
277
278 for (const xAOD::MuonSegment* segment : segments) {
279 if (!retainedSegments.contains(segment)) continue;
280
281 const Amg::Vector3D position = segment->position();
282 const Amg::Vector3D direction = segment->direction();
283 const SegmentBucketKey key = segmentBucketKey(*segment);
284 const auto bucketIt = segmentsByBucket.find(key);
285 const int multiplicity =
286 bucketIt == segmentsByBucket.end()
287 ? 1
288 : static_cast<int>(bucketIt->second.size());
289
290 const int chamberIndex = static_cast<int>(segment->chamberIndex());
291 const int layers = layersInBucket(*MuonR4::detailedSegment(*segment)->parent()->parentBucket());
292 const int sector = segment->sector();
293
294 graph.segments.push_back(segment);
295 pos.emplace_back(position / Gaudi::Units::m);
296 dir.emplace_back(direction);
297 bucket.emplace_back(BucketSegmentFeatures{
298 chamberIndex, layers, sector, multiplicity});
299 for (const SegmentNodeFeatureId featureId : m_nodeFeatureIds) {
300 graph.nodeFeatures.push_back(
301 nodeFeatureValue(featureId, pos.back(), dir.back(), bucket.back()));
302 }
303 }
304 graph.nNodes = graph.segments.size();
305
306 if (pos.size() != graph.nNodes || dir.size() != graph.nNodes ||
307 bucket.size() != graph.nNodes) {
308 ATH_MSG_ERROR("Inconsistent vector sizes during graph building: nodes="
309 << graph.nNodes << ", pos=" << pos.size()
310 << ", dir=" << dir.size() << ", bucket=" << bucket.size());
311 return StatusCode::FAILURE;
312 }
313
314 if (graph.nNodes < 2) {
315 graph.nEdges = 0;
316 if (truthDiag) fillTruthDiagnostics(segments, retainedSegments, graph);
317 return StatusCode::SUCCESS;
318 }
319
320 const auto wrapRegularSector = [&](int sector) {
321 // MuonSegment::sector() is a regular sector number, not an ExpandedSector
322 // coordinate. Wrap it to [0, modulo); <= 0 disables wrapping.
323 if (m_sectorModulo.value() > 0) {
324 sector %= m_sectorModulo.value();
325 if (sector < 0) sector += m_sectorModulo.value();
326 }
327 return sector;
328 };
329
330 // The lookup key must use the same wrapping as the target sectors below:
331 // ATLAS sectors are 1-based (1..16), so a raw key of 16 can never match a
332 // wrapped target of 0, which silently dropped every edge into sector 16.
333 // The per-pair sectorDistance check below enforces the true circular
334 // distance on the raw sector numbers.
335 std::unordered_map<int, std::vector<std::size_t>> nodesBySector;
336 nodesBySector.reserve(graph.nNodes);
337 for (std::size_t node = 0; node < graph.nNodes; ++node) {
338 nodesBySector[wrapRegularSector(bucket[node].sector)].push_back(node);
339 }
340
341 std::unordered_map<int, std::vector<int>> targetSectorsBySourceSector;
342 targetSectorsBySourceSector.reserve(nodesBySector.size());
343 std::size_t sectorLocalEdgeUpperBound = 0;
344 for (const auto& [sourceSector, sourceNodes] : nodesBySector) {
345 std::vector<int> targetSectors;
346 targetSectors.reserve(2 * m_maxDeltaSector.value() + 1);
347 for (int delta = -m_maxDeltaSector.value();
348 delta <= m_maxDeltaSector.value(); ++delta) {
349 targetSectors.push_back(wrapRegularSector(sourceSector + delta));
350 }
351 for (const int targetSector : targetSectors) {
352 const auto found = nodesBySector.find(targetSector);
353 if (found == nodesBySector.end()) continue;
354 sectorLocalEdgeUpperBound += sourceNodes.size() * found->second.size();
355 if (targetSector == sourceSector) {
356 sectorLocalEdgeUpperBound -= sourceNodes.size();
357 }
358 }
359 targetSectorsBySourceSector.emplace(sourceSector,
360 std::move(targetSectors));
361 }
362
363 /*
364 * The model receives a directed graph, but the geometric candidate relation
365 * is undirected. Build each pair once, then emit both directions. With a
366 * non-zero input cap, each endpoint nominates its best candidates and the
367 * union is made bidirectional before inference; this preserves the message
368 * passing symmetry expected by the GNN.
369 */
370 struct UndirectedEdge {
371 std::size_t first{0};
372 std::size_t second{0};
373 float dx{0.f};
374 float dy{0.f};
375 float dz{0.f};
376 float distance{0.f};
377 float cosAngle{0.f};
378 };
379 const auto betterEdge = [](const UndirectedEdge& first,
380 const UndirectedEdge& second) {
381 const int cosOrder = InferenceUtils::compareFloatDescending(
382 first.cosAngle, second.cosAngle);
383 if (cosOrder != 0) {
384 return cosOrder < 0;
385 }
386
387 const int distanceOrder =
388 InferenceUtils::compareFloat(first.distance, second.distance);
389 if (distanceOrder != 0) {
390 return distanceOrder < 0;
391 }
392 if (first.first != second.first) return first.first < second.first;
393 return first.second < second.second;
394 };
395 const auto edgeKey = [](const UndirectedEdge& edge) {
396 return (static_cast<std::uint64_t>(edge.first) << 32) |
397 static_cast<std::uint64_t>(edge.second);
398 };
399
400 const unsigned int maxEdgesPerNode =
402 const unsigned int maxEdgesPerTargetChamber =
404 const bool usePreInferenceSelection =
405 maxEdgesPerNode != 0 || maxEdgesPerTargetChamber != 0;
406 std::vector<std::vector<UndirectedEdge>> bestEdgesByNode;
407 if (usePreInferenceSelection) {
408 bestEdgesByNode.resize(graph.nNodes);
409 const unsigned int reservePerNode =
410 maxEdgesPerNode != 0 ? maxEdgesPerNode : maxEdgesPerTargetChamber;
411 for (std::vector<UndirectedEdge>& edges : bestEdgesByNode) {
412 edges.reserve(reservePerNode);
413 }
414 } else {
415 graph.edgeIndex.reserve(2 * sectorLocalEdgeUpperBound);
416 graph.edgeFeatures.reserve(kEdgeFeatureCount * sectorLocalEdgeUpperBound);
417 }
418 const auto appendDirectedPair = [&](const UndirectedEdge& edge) {
419 graph.edgeIndex.push_back(static_cast<int64_t>(edge.first));
420 graph.edgeIndex.push_back(static_cast<int64_t>(edge.second));
421 graph.edgeFeatures.insert(
422 graph.edgeFeatures.end(),
423 {edge.dx, edge.dy, edge.dz, edge.distance, edge.cosAngle,
424 float(bucket[edge.first].chamberIndex ==
425 bucket[edge.second].chamberIndex),
426 float(bucket[edge.first].sector == bucket[edge.second].sector)});
427
428 graph.edgeIndex.push_back(static_cast<int64_t>(edge.second));
429 graph.edgeIndex.push_back(static_cast<int64_t>(edge.first));
430 graph.edgeFeatures.insert(
431 graph.edgeFeatures.end(),
432 {-edge.dx, -edge.dy, -edge.dz, edge.distance, edge.cosAngle,
433 float(bucket[edge.first].chamberIndex ==
434 bucket[edge.second].chamberIndex),
435 float(bucket[edge.first].sector == bucket[edge.second].sector)});
436 };
437
438 const auto retainForNode = [&](std::size_t node,
439 const UndirectedEdge& candidate) {
440 std::vector<UndirectedEdge>& retained = bestEdgesByNode[node];
441 const std::size_t other = candidate.first == node ? candidate.second
442 : candidate.first;
443 const int targetChamber = bucket[other].chamberIndex;
444
445 if (maxEdgesPerTargetChamber != 0) {
446 unsigned int sameChamberCount = 0;
447 auto worstSameChamber = retained.end();
448 for (auto it = retained.begin(); it != retained.end(); ++it) {
449 const std::size_t retainedOther =
450 it->first == node ? it->second : it->first;
451 if (bucket[retainedOther].chamberIndex != targetChamber) continue;
452 ++sameChamberCount;
453 if (worstSameChamber == retained.end() ||
454 betterEdge(*worstSameChamber, *it)) {
455 worstSameChamber = it;
456 }
457 }
458 if (sameChamberCount >= maxEdgesPerTargetChamber) {
459 if (!betterEdge(candidate, *worstSameChamber)) return;
460 *worstSameChamber = candidate;
461 } else {
462 retained.push_back(candidate);
463 }
464 } else {
465 retained.push_back(candidate);
466 }
467
468 if (maxEdgesPerNode != 0 && retained.size() > maxEdgesPerNode) {
469 auto worst = retained.begin();
470 for (auto it = std::next(retained.begin()); it != retained.end(); ++it) {
471 if (betterEdge(*worst, *it)) worst = it;
472 }
473 retained.erase(worst);
474 }
475 };
476
477 std::size_t candidatePairs = 0;
478 for (std::size_t first = 0; first < graph.nNodes; ++first) {
479 const auto sectorsIt =
480 targetSectorsBySourceSector.find(wrapRegularSector(bucket[first].sector));
481 if (sectorsIt == targetSectorsBySourceSector.end()) continue;
482 for (const int sector : sectorsIt->second) {
483 const auto targetIt = nodesBySector.find(sector);
484 if (targetIt == nodesBySector.end()) continue;
485
486 for (const std::size_t second : targetIt->second) {
487 // Every valid pair will be visited from the lower-index endpoint.
488 if (second <= first) continue;
489 if (sectorDistance(bucket[first].sector, bucket[second].sector,
490 m_sectorModulo.value()) >
491 m_maxDeltaSector.value()) {
492 continue;
493 }
495 bucket[first].chamberIndex == bucket[second].chamberIndex) {
496 continue;
497 }
498 const float cosAngle = static_cast<float>(dir[first].dot(dir[second]));
499 if (cosAngle < m_cosMin) continue;
500
501 const Amg::Vector3D delta = pos[second] - pos[first];
502 const UndirectedEdge candidate{
503 first,
504 second,
505 static_cast<float>(delta.x()),
506 static_cast<float>(delta.y()),
507 static_cast<float>(delta.z()),
508 static_cast<float>(delta.mag()),
509 cosAngle};
510 ++candidatePairs;
511
512 if (!usePreInferenceSelection) {
513 appendDirectedPair(candidate);
514 } else {
515 retainForNode(first, candidate);
516 retainForNode(second, candidate);
517 }
518 }
519 }
520 }
521 std::size_t retainedPairs = candidatePairs;
522 if (usePreInferenceSelection) {
523 const unsigned int selectedReservePerNode =
524 maxEdgesPerNode != 0 ? maxEdgesPerNode : maxEdgesPerTargetChamber;
525 std::unordered_set<std::uint64_t> selectedKeys;
526 selectedKeys.reserve(graph.nNodes * selectedReservePerNode);
527 std::vector<UndirectedEdge> selectedEdges;
528 selectedEdges.reserve(graph.nNodes * selectedReservePerNode);
529
530 for (const std::vector<UndirectedEdge>& nodeEdges : bestEdgesByNode) {
531 for (const UndirectedEdge& edge : nodeEdges) {
532 if (selectedKeys.insert(edgeKey(edge)).second) {
533 selectedEdges.push_back(edge);
534 }
535 }
536 }
537 std::sort(selectedEdges.begin(), selectedEdges.end(),
538 [](const UndirectedEdge& first,
539 const UndirectedEdge& second) {
540 if (first.first != second.first) {
541 return first.first < second.first;
542 }
543 return first.second < second.second;
544 });
545
546 retainedPairs = selectedEdges.size();
547 graph.edgeIndex.reserve(4 * retainedPairs);
548 graph.edgeFeatures.reserve(2 * kEdgeFeatureCount * retainedPairs);
549 for (const UndirectedEdge& edge : selectedEdges) {
550 appendDirectedPair(edge);
551 }
552 }
553 graph.nEdges = graph.edgeIndex.size() / 2;
554 const std::size_t nodesBeforeIsolatedNodeDrop = graph.nNodes;
555 if (m_dropIsolatedNodesBeforeInference.value() && graph.nEdges != 0) {
556 std::vector<unsigned char> active(graph.nNodes, 0);
557 for (const int64_t index : graph.edgeIndex) {
558 active[static_cast<std::size_t>(index)] = 1;
559 }
560 const std::size_t activeNodes =
561 std::count(active.begin(), active.end(), static_cast<unsigned char>(1));
562 if (activeNodes != graph.nNodes) {
563 std::vector<std::size_t> oldToNew(graph.nNodes, graph.nNodes);
564 std::vector<const xAOD::MuonSegment*> compactedSegments;
565 std::vector<float> compactedNodeFeatures;
566 compactedSegments.reserve(activeNodes);
567 compactedNodeFeatures.reserve(activeNodes * kNodeFeatureCount);
568 for (std::size_t oldNode = 0; oldNode < graph.nNodes; ++oldNode) {
569 if (!active[oldNode]) continue;
570 oldToNew[oldNode] = compactedSegments.size();
571 compactedSegments.push_back(graph.segments[oldNode]);
572 const auto featureBegin = graph.nodeFeatures.begin() +
573 oldNode * kNodeFeatureCount;
574 compactedNodeFeatures.insert(compactedNodeFeatures.end(),
575 featureBegin,
576 featureBegin + kNodeFeatureCount);
577 }
578 for (int64_t& index : graph.edgeIndex) {
579 index = static_cast<int64_t>(oldToNew[static_cast<std::size_t>(index)]);
580 }
581 graph.segments = std::move(compactedSegments);
582 graph.nodeFeatures = std::move(compactedNodeFeatures);
583 graph.nNodes = activeNodes;
584 }
585 }
586 if (truthDiag) fillTruthDiagnostics(segments, retainedSegments, graph);
587 ATH_MSG_DEBUG("buildGraph: input segments=" << segments.size()
588 << ", kept nodes=" << graph.nNodes
589 << ", nodes before isolated-node drop=" << nodesBeforeIsolatedNodeDrop
590 << ", bucket cap=" << m_maxSegmentsPerBucket.value()
591 << ", candidate pairs=" << candidatePairs
592 << ", retained pairs=" << retainedPairs
593 << ", built directed edges=" << graph.nEdges
594 << ", pre-inference node cap=" << m_maxEdgesPerNodeBeforeInference.value()
595 << ", per-target-chamber cap=" << maxEdgesPerTargetChamber
596 << ", drop same chamber=" << m_dropSameChamberEdgesBeforeInference.value()
597 << ", drop isolated nodes=" << m_dropIsolatedNodesBeforeInference.value()
598 << ", sector-local reserve=" << sectorLocalEdgeUpperBound);
599
600 // Job-summed diagnostics
601 if (msgLvl(MSG::DEBUG)) {
602 m_sumInputSegments += segments.size();
603 m_sumCandidatePairs += candidatePairs;
604 m_sumRetainedPairs += retainedPairs;
605 m_sumNodesBeforeIsolatedDrop += nodesBeforeIsolatedNodeDrop;
606 m_sumNodesAfterIsolatedDrop += graph.nNodes;
607 }
608 return StatusCode::SUCCESS;
609}
static constexpr std::size_t kEdgeFeatureCount
static constexpr std::size_t kNodeFeatureCount
std::atomic< std::size_t > m_sumInputSegments
Job-summed pre-ONNX pruning counters (see buildGraph()).
void fillTruthDiagnostics(const xAOD::MuonSegmentContainer &segments, const std::unordered_set< const xAOD::MuonSegment * > &bucketRetained, SegmentEdgeGraph &graph) const
MC-only diagnostics: record, for every input segment and every pair of input segments sharing a truth...
Gaudi::Property< unsigned int > m_maxEdgesPerNodeBeforeInference
Gaudi::Property< unsigned int > m_maxEdgesPerTargetChamberBeforeInference
std::atomic< std::size_t > m_sumNodesBeforeIsolatedDrop
std::atomic< std::size_t > m_sumRetainedPairs
Gaudi::Property< unsigned int > m_maxSegmentsPerBucket
Gaudi::Property< bool > m_dropSameChamberEdgesBeforeInference
std::vector< SegmentNodeFeatureId > m_nodeFeatureIds
std::atomic< std::size_t > m_sumCandidatePairs
Gaudi::Property< bool > m_dropIsolatedNodesBeforeInference
std::atomic< std::size_t > m_sumNodesAfterIsolatedDrop
const SpacePointBucket * parentBucket() const
Returns the bucket out of which the seed was formed.
const SegmentSeed * parent() const
Returns the seed out of which the segment was built.
float distance(const Amg::Vector3D &p1, const Amg::Vector3D &p2)
calculates the distance between two point in 3D space
Eigen::Matrix< double, 3, 1 > Vector3D
bool first
Definition DeMoScan.py:534
str index
Definition DeMoScan.py:362
int compareFloat(float first, float second)
Three-way float comparison which orders NaN after all numeric values.
int compareFloatDescending(float first, float second)
Three-way descending comparison which also orders NaN last.
SegmentNodeFeatureId
Identifier for each node feature in segment-based GNNs.
Definition MuonMLEvent.h:28
const Segment * detailedSegment(const xAOD::MuonSegment &seg)
Helper function to navigate from the xAOD::MuonSegment to the MuonR4::Segment.
const Amg::Vector3D & direction() const
Method to retrieve the direction at the Intersection.
const Amg::Vector3D & position() const
Method to retrieve the position of the Intersection.
@ active
Definition Layer.h:47
void sort(typename DataModel_detail::iterator< DVL > beg, typename DataModel_detail::iterator< DVL > end)
Specialization of sort for DataVector/List.
MuonSegment_v1 MuonSegment
Reference the current persistent version:

◆ buildTransformerInputs()

StatusCode BucketInferenceToolBase::buildTransformerInputs ( const EventContext & ctx,
GraphRawData & graphData ) const
protectedinherited

Build Transformer inputs: features [1,S,6] and pad_mask [1,S] (False = valid), as tensors 0 and 1.

Definition at line 138 of file BucketInferenceToolBase.cxx.

139 {
140 // Start from (N,6)
141 ATH_CHECK(buildFeaturesOnly(ctx, graphData));
142
143 // Copy features flat buffer for lifetime management
144 std::vector<float> featuresFlat = graphData.featureLeaves;
145 const int64_t S = static_cast<int64_t>(featuresFlat.size() / kBucketFeatureCount);
146
147 if (S == 0) {
148 ATH_MSG_WARNING("No valid features for transformer input. Skipping inference.");
149 return StatusCode::SUCCESS;
150 }
151
152 if (msgLvl(MSG::DEBUG)) {
153 // DEBUG: Print transformer input features for first 10 nodes
154 ATH_MSG_DEBUG("=== DEBUGGING: Transformer input features for first 10 nodes ===");
155 const int64_t debugNodes = std::min(S, static_cast<int64_t>(10));
156 for (int64_t nodeIdx = 0; nodeIdx < debugNodes; ++nodeIdx) {
157 const int64_t baseIdx = nodeIdx * static_cast<int64_t>(kBucketFeatureCount);
158 ATH_MSG_DEBUG("TransformerNode[" << nodeIdx << "]: "
159 << "x=" << featuresFlat[baseIdx + 0] << ", "
160 << "y=" << featuresFlat[baseIdx + 1] << ", "
161 << "z=" << featuresFlat[baseIdx + 2] << ", "
162 << "layers=" << featuresFlat[baseIdx + 3] << ", "
163 << "nSp=" << featuresFlat[baseIdx + 4] << ", "
164 << "bucketSize=" << featuresFlat[baseIdx + 5]);
165 }
166 ATH_MSG_DEBUG("=== END DEBUG TRANSFORMER FEATURES ===");
167 }
168
169 // Rebuild graph with exactly 2 inputs: features [1,S,6], pad_mask [1,S]
170 graphData.graph.reset();
171 graphData.graph = std::make_unique<InferenceGraph>();
172 graphData.graph->dataTensor.reserve(2); // features and pad_mask inputs; outputs are reserved in runNamedInference()
173
174 Ort::MemoryInfo memInfo = Ort::MemoryInfo::CreateCpu(OrtDeviceAllocator, OrtMemTypeCPU);
175
176 // features: [1,S,6] (backed by graphData.featureLeaves to keep alive)
177 std::vector<int64_t> fShape{1, S, static_cast<int64_t>(kBucketFeatureCount)};
178 graphData.featureLeaves.swap(featuresFlat);
179 graphData.graph->dataTensor.emplace_back(
180 Ort::Value::CreateTensor<float>(memInfo,
181 graphData.featureLeaves.data(),
182 graphData.featureLeaves.size(),
183 fShape.data(),
184 fShape.size()));
185
186 // pad_mask: [1,S] (bool). Create ORT-owned tensor and fill with False (=valid).
187 Ort::AllocatorWithDefaultOptions allocator;
188 std::vector<int64_t> mShape{1, S};
189 Ort::Value padVal = Ort::Value::CreateTensor(allocator,
190 mShape.data(),
191 mShape.size(),
192 ONNX_TENSOR_ELEMENT_DATA_TYPE_BOOL);
193 bool* maskPtr = padVal.GetTensorMutableData<bool>();
194 for (int64_t i = 0; i < S; ++i) maskPtr[i] = false;
195 graphData.graph->dataTensor.emplace_back(std::move(padVal));
196
197 return StatusCode::SUCCESS;
198}
StatusCode buildFeaturesOnly(const EventContext &ctx, GraphRawData &graphData) const
Build only features (N,6); attaches one tensor in graph.dataTensor[0].

◆ classifyEdges()

StatusCode MuonML::SegmentEdgeClassifierTool::classifyEdges ( const EventContext & ctx,
const SegmentEdgeGraph & graph,
std::vector< SegmentEdgeScore > & scores ) const
overridevirtual

Run ONNX inference on graph and populate scores with logit and probability for each edge; called after buildGraph().

Implements MuonML::ISegmentEdgeClassifierTool.

Definition at line 710 of file SegmentEdgeClassifierTool.cxx.

712 {
713 scores.clear();
714 if (!graph.nNodes) return StatusCode::SUCCESS;
715 if (!graph.nEdges) {
716 ATH_CHECK(dumpDebugEvent(ctx, graph, scores));
717 return StatusCode::SUCCESS;
718 }
719
720 if (graph.nodeFeatures.size() != graph.nNodes * kNodeFeatureCount) {
721 ATH_MSG_ERROR("Unexpected node feature size " << graph.nodeFeatures.size()
722 << "; expected " << (graph.nNodes * kNodeFeatureCount));
723 return StatusCode::FAILURE;
724 }
725 if (graph.edgeIndex.size() != 2 * graph.nEdges) {
726 ATH_MSG_ERROR("Unexpected edge index size " << graph.edgeIndex.size()
727 << "; expected " << (2 * graph.nEdges));
728 return StatusCode::FAILURE;
729 }
730 if (graph.edgeFeatures.size() != graph.nEdges * kEdgeFeatureCount) {
731 ATH_MSG_ERROR("Unexpected edge feature size " << graph.edgeFeatures.size()
732 << "; expected " << (graph.nEdges * kEdgeFeatureCount));
733 return StatusCode::FAILURE;
734 }
735
736 GraphRawData raw{};
737 raw.graph = std::make_unique<InferenceGraph>();
738 raw.edgeIndexPacked.resize(2 * graph.nEdges);
739 for (std::size_t e = 0; e < graph.nEdges; ++e) {
740 raw.edgeIndexPacked[e] = graph.edgeIndex[2 * e];
741 raw.edgeIndexPacked[graph.nEdges + e] = graph.edgeIndex[2 * e + 1];
742 }
743
744 Ort::MemoryInfo memInfo = Ort::MemoryInfo::CreateCpu(OrtDeviceAllocator, OrtMemTypeCPU);
745
746 const std::vector<int64_t> nodeShape{static_cast<int64_t>(graph.nNodes), static_cast<int64_t>(kNodeFeatureCount)};
747 // The graph outlives the synchronous ONNX call below. Use its node
748 // buffer directly instead of allocating and copying featureLeaves per event.
749 ATLAS_THREAD_SAFE float* nodeFeaturesData =
750 const_cast<float*>(graph.nodeFeatures.data());
751 raw.graph->dataTensor.emplace_back(
752 Ort::Value::CreateTensor<float>(memInfo,
753 nodeFeaturesData,
754 graph.nodeFeatures.size(),
755 nodeShape.data(),
756 nodeShape.size()));
757
758 const std::vector<int64_t> edgeIndexShape{2, static_cast<int64_t>(graph.nEdges)};
759 raw.graph->dataTensor.emplace_back(
760 Ort::Value::CreateTensor<int64_t>(memInfo,
761 raw.edgeIndexPacked.data(),
762 raw.edgeIndexPacked.size(),
763 edgeIndexShape.data(),
764 edgeIndexShape.size()));
765
766 // ONNX Runtime's CreateTensor API takes a non-const pointer, but it does not
767 // mutate input buffers during inference. Avoid copying edge_attr every event.
768 ATLAS_THREAD_SAFE float* edgeFeaturesData = const_cast<float*>(graph.edgeFeatures.data());
769 const std::vector<int64_t> edgeAttrShape{static_cast<int64_t>(graph.nEdges), static_cast<int64_t>(kEdgeFeatureCount)};
770 raw.graph->dataTensor.emplace_back(
771 Ort::Value::CreateTensor<float>(memInfo,
772 edgeFeaturesData,
773 graph.edgeFeatures.size(),
774 edgeAttrShape.data(),
775 edgeAttrShape.size()));
776
777 const std::vector<const char*> inputNames{
778 m_inputNodeName.value().c_str(),
779 m_inputEdgeIndexName.value().c_str(),
780 m_inputEdgeAttrName.value().c_str()};
781 const std::vector<const char*> outputNames{m_outputName.value().c_str()};
782 ATH_MSG_DEBUG("classifyEdges: ONNX inputs shapes x=[" << nodeShape[0] << "," << nodeShape[1]
783 << "], edge_index=[" << edgeIndexShape[0] << "," << edgeIndexShape[1]
784 << "], edge_attr=[" << edgeAttrShape[0] << "," << edgeAttrShape[1] << "]");
785 ATH_CHECK(runNamedInference(raw, inputNames, outputNames));
786
787 if (raw.graph->dataTensor.size() <= inputNames.size()) {
788 ATH_MSG_ERROR("Missing ONNX output tensor for segment edge inference");
789 return StatusCode::FAILURE;
790 }
791
792 const Ort::Value& outTensor = raw.graph->dataTensor[inputNames.size()];
793 const auto outInfo = outTensor.GetTensorTypeAndShapeInfo();
794 const std::vector<int64_t> outShape = outInfo.GetShape();
795 const size_t outSize = outInfo.GetElementCount();
796 if (!outShape.empty()) {
797 ATH_MSG_DEBUG("classifyEdges: ONNX output rank=" << outShape.size()
798 << ", first dim=" << outShape.front()
799 << ", elements=" << outSize);
800 } else {
801 ATH_MSG_DEBUG("classifyEdges: ONNX scalar output, elements=" << outSize);
802 }
803 if (outSize < graph.nEdges) {
804 ATH_MSG_ERROR("ONNX logits tensor has " << outSize << " entries for " << graph.nEdges << " edges");
805 return StatusCode::FAILURE;
806 }
807
808 const float* logits = outTensor.GetTensorData<float>();
809 scores.reserve(graph.nEdges);
810 for (std::size_t e=0; e<graph.nEdges; ++e) {
811 const float l = logits[e];
812 scores.push_back({std::size_t(graph.edgeIndex[2 * e]),
813 std::size_t(graph.edgeIndex[2 * e + 1]),
814 l,
816 }
817
818 ATH_CHECK(dumpDebugEvent(ctx, graph, scores));
819 return StatusCode::SUCCESS;
820}
std::vector< std::vector< float > > scores
#define ATLAS_THREAD_SAFE
StatusCode runNamedInference(GraphRawData &graphData, const std::vector< const char * > &inputNames, const std::vector< const char * > &outputNames) const
Generic named inference, for tools with different I/O conventions.
Gaudi::Property< std::string > m_outputName
Gaudi::Property< std::string > m_inputEdgeAttrName
Gaudi::Property< std::string > m_inputEdgeIndexName
Gaudi::Property< std::string > m_inputNodeName
StatusCode dumpDebugEvent(const EventContext &ctx, const SegmentEdgeGraph &graph, const std::vector< SegmentEdgeScore > &scores) const
l
Printing final latex table to .tex output file.

◆ DeclareInterfaceID()

MuonML::ISegmentEdgeClassifierTool::DeclareInterfaceID ( ISegmentEdgeClassifierTool ,
1 ,
0  )
inherited

◆ dumpDebugEvent()

StatusCode MuonML::SegmentEdgeClassifierTool::dumpDebugEvent ( const EventContext & ctx,
const SegmentEdgeGraph & graph,
const std::vector< SegmentEdgeScore > & scores ) const
private

Definition at line 822 of file SegmentEdgeClassifierTool.cxx.

825 {
826 if (m_debugDumpFile.value().empty()) return StatusCode::SUCCESS;
827
828 std::lock_guard<std::mutex> lock{m_debugDumpMutex};
829 if (m_debugDumpMaxEvents.value() != 0 &&
830 m_debugDumpEvents.load(std::memory_order_relaxed) >=
831 m_debugDumpMaxEvents.value()) {
832 return StatusCode::SUCCESS;
833 }
834
835 if (graph.nodeFeatures.size() != graph.nNodes * kNodeFeatureCount ||
836 graph.edgeIndex.size() != graph.nEdges * 2 ||
837 graph.edgeFeatures.size() != graph.nEdges * kEdgeFeatureCount ||
838 scores.size() != graph.nEdges) {
839 ATH_MSG_ERROR("Cannot write segment-edge debug dump: inconsistent graph/output sizes"
840 << " nodes=" << graph.nNodes
841 << " nodeFeatures=" << graph.nodeFeatures.size()
842 << " edges=" << graph.nEdges
843 << " edgeIndex=" << graph.edgeIndex.size()
844 << " edgeFeatures=" << graph.edgeFeatures.size()
845 << " scores=" << scores.size());
846 return StatusCode::FAILURE;
847 }
848
849 nlohmann::json x = nlohmann::json::array();
850 x.get_ref<nlohmann::json::array_t&>().reserve(graph.nodeFeatures.size());
851 for (const float value : graph.nodeFeatures) {
852 x.push_back(std::isfinite(value) ? nlohmann::json(value)
853 : nlohmann::json(nullptr));
854 }
855
856 nlohmann::json edgeIndex = nlohmann::json::array();
857 edgeIndex.get_ref<nlohmann::json::array_t&>().reserve(graph.nEdges * 2);
858 // This is the actual ONNX [2,E] row-major buffer: all sources then all destinations.
859 for (std::size_t edge = 0; edge < graph.nEdges; ++edge) {
860 edgeIndex.push_back(graph.edgeIndex[2 * edge]);
861 }
862 for (std::size_t edge = 0; edge < graph.nEdges; ++edge) {
863 edgeIndex.push_back(graph.edgeIndex[2 * edge + 1]);
864 }
865
866 nlohmann::json edgeAttr = nlohmann::json::array();
867 edgeAttr.get_ref<nlohmann::json::array_t&>().reserve(graph.edgeFeatures.size());
868 for (const float value : graph.edgeFeatures) {
869 edgeAttr.push_back(std::isfinite(value) ? nlohmann::json(value)
870 : nlohmann::json(nullptr));
871 }
872
873 nlohmann::json logits = nlohmann::json::array();
874 nlohmann::json probabilities = nlohmann::json::array();
875 nlohmann::json edgeSrc = nlohmann::json::array();
876 nlohmann::json edgeDst = nlohmann::json::array();
877 logits.get_ref<nlohmann::json::array_t&>().reserve(scores.size());
878 probabilities.get_ref<nlohmann::json::array_t&>().reserve(scores.size());
879 edgeSrc.get_ref<nlohmann::json::array_t&>().reserve(scores.size());
880 edgeDst.get_ref<nlohmann::json::array_t&>().reserve(scores.size());
881 for (const SegmentEdgeScore& score : scores) {
882 edgeSrc.push_back(score.src);
883 edgeDst.push_back(score.dst);
884 logits.push_back(std::isfinite(score.logit) ? nlohmann::json(score.logit)
885 : nlohmann::json(nullptr));
886 probabilities.push_back(std::isfinite(score.probability)
887 ? nlohmann::json(score.probability)
888 : nlohmann::json(nullptr));
889 }
890
891 std::ofstream out{m_debugDumpFile.value(), std::ios::out | std::ios::app};
892 if (!out) {
893 ATH_MSG_ERROR("Could not append to segment-edge debug dump file: "
894 << m_debugDumpFile.value());
895 return StatusCode::FAILURE;
896 }
897
898 const unsigned int dumpIndex =
899 m_debugDumpEvents.fetch_add(1, std::memory_order_relaxed);
900 nlohmann::ordered_json event;
901 event["record_type"] = "event";
902 event["format_version"] = 1;
903 event["dump_index"] = dumpIndex;
904 event["run_number"] = ctx.eventID().run_number();
905 event["lumi_block"] = ctx.eventID().lumi_block();
906 event["event_number"] = ctx.eventID().event_number();
907 event["slot"] = ctx.slot();
908 event["n_nodes"] = graph.nNodes;
909 event["n_edges"] = graph.nEdges;
910 event["x_shape"] = {graph.nNodes, kNodeFeatureCount};
911 event["edge_index_shape"] = {2, graph.nEdges};
912 event["edge_attr_shape"] = {graph.nEdges, kEdgeFeatureCount};
913 event["logits_shape"] = {graph.nEdges};
914 event["x"] = std::move(x);
915 event["edge_index"] = std::move(edgeIndex);
916 event["edge_attr"] = std::move(edgeAttr);
917 event["edge_src"] = std::move(edgeSrc);
918 event["edge_dst"] = std::move(edgeDst);
919 event["logits"] = std::move(logits);
920 event["probabilities"] = std::move(probabilities);
921 out << event.dump() << '\n';
922
923 ATH_MSG_DEBUG("Wrote segment-edge debug event " << dumpIndex
924 << " to " << m_debugDumpFile.value());
925
926 return StatusCode::SUCCESS;
927}
virtual void lock()=0
Interface to allow an object to lock itself when made const in SG.
#define x
Gaudi::Property< unsigned int > m_debugDumpMaxEvents
std::atomic< unsigned int > m_debugDumpEvents
Gaudi::Property< std::string > m_debugDumpFile
virtual void reserve(size_t sz) override
Change the capacity of all aux data vectors.

◆ enableTruthDiagnostics()

bool MuonML::SegmentEdgeClassifierTool::enableTruthDiagnostics ( ) const
inlineoverridevirtual

Whether this tool was configured to fill SegmentEdgeGraph's truth diagnostics.

Implements MuonML::ISegmentEdgeClassifierTool.

Definition at line 93 of file SegmentEdgeClassifierTool.h.

◆ fillTruthDiagnostics()

void MuonML::SegmentEdgeClassifierTool::fillTruthDiagnostics ( const xAOD::MuonSegmentContainer & segments,
const std::unordered_set< const xAOD::MuonSegment * > & bucketRetained,
SegmentEdgeGraph & graph ) const
private

MC-only diagnostics: record, for every input segment and every pair of input segments sharing a truth particle, why it did or did not reach ONNX.

Only called when EnableTruthDiagnostics and DEBUG output are set.

Definition at line 611 of file SegmentEdgeClassifierTool.cxx.

614 {
615 std::unordered_map<const xAOD::MuonSegment*, std::int32_t> nodeOf;
616 nodeOf.reserve(graph.segments.size());
617 for (std::size_t node = 0; node < graph.segments.size(); ++node) {
618 nodeOf.emplace(graph.segments[node], static_cast<std::int32_t>(node));
619 }
620
621 graph.inputNodeIndex.assign(segments.size(), kDroppedAsIsolated);
622 std::unordered_map<std::int32_t, std::vector<std::uint32_t>> byTruth;
623 std::uint32_t inputIndex = 0;
624 for (const xAOD::MuonSegment* segment : segments) {
625 const auto found = nodeOf.find(segment);
626 if (found != nodeOf.end()) {
627 graph.inputNodeIndex[inputIndex] = found->second;
628 } else if (!bucketRetained.contains(segment)) {
629 graph.inputNodeIndex[inputIndex] = kDroppedByBucketCap;
630 }
631 if (const xAOD::TruthParticle* truthPart =
633 byTruth[static_cast<std::int32_t>(truthPart->index())].push_back(inputIndex);
634 }
635 ++inputIndex;
636 }
637
638 const auto pairKey = [](std::uint64_t first, std::uint64_t second) {
639 return first < second ? (first << 32) | second : (second << 32) | first;
640 };
641 std::unordered_set<std::uint64_t> scoredPairs;
642 scoredPairs.reserve(graph.nEdges);
643 for (std::size_t edge = 0; edge < graph.nEdges; ++edge) {
644 scoredPairs.insert(
645 pairKey(static_cast<std::uint64_t>(graph.edgeIndex[2 * edge]),
646 static_cast<std::uint64_t>(graph.edgeIndex[2 * edge + 1])));
647 }
648
649 // Same order as the pair loop in buildGraph(): bucket cap (node level),
650 // sector window, same-chamber drop, angle window; a pair that clears all of
651 // them but was not scored can only have been evicted by the edge caps.
652 for (const auto& entry : byTruth) {
653 const std::vector<std::uint32_t>& members = entry.second;
654 for (std::size_t x = 0; x < members.size(); ++x) {
655 for (std::size_t y = x + 1; y < members.size(); ++y) {
656 const std::uint32_t i = members[x];
657 const std::uint32_t j = members[y];
658 const xAOD::MuonSegment* first = segments[i];
659 const xAOD::MuonSegment* second = segments[j];
660 PairFate fate{i, j, PairGate::Scored, 0};
661 const int sectorDelta = sectorDistance(first->sector(), second->sector(),
662 m_sectorModulo.value());
663 fate.sectorDelta = static_cast<std::uint8_t>(std::min(sectorDelta, 255));
664 if (!bucketRetained.contains(first) || !bucketRetained.contains(second)) {
665 fate.gate = PairGate::BucketCap;
666 } else if (sectorDelta > m_maxDeltaSector.value()) {
667 fate.gate = PairGate::SectorWindow;
668 } else if (m_dropSameChamberEdgesBeforeInference.value() &&
669 first->chamberIndex() == second->chamberIndex()) {
670 fate.gate = PairGate::SameChamber;
671 } else if (static_cast<float>(first->direction().dot(second->direction())) <
672 m_cosMin) {
673 fate.gate = PairGate::AngleWindow;
674 } else {
675 const std::int32_t firstNode = graph.inputNodeIndex[i];
676 const std::int32_t secondNode = graph.inputNodeIndex[j];
677 const bool scored =
678 firstNode >= 0 && secondNode >= 0 &&
679 scoredPairs.contains(pairKey(static_cast<std::uint64_t>(firstNode),
680 static_cast<std::uint64_t>(secondNode)));
681 if (!scored) fate.gate = PairGate::EdgeCaps;
682 }
683 graph.truthPairFates.push_back(fate);
684 }
685 }
686 }
687}
#define y
@ Scored
the pair was sent to ONNX
@ SameChamber
same-chamber pair dropped by design
@ AngleWindow
cos(opening angle) below MaxDeltaThetaDeg
@ EdgeCaps
candidate evicted by the pre-ONNX per-node / per-target-chamber caps
@ SectorWindow
|delta sector| > MaxDeltaSector
@ BucketCap
an endpoint was dropped by MaxSegmentsPerBucket
constexpr std::int32_t kDroppedAsIsolated
constexpr std::int32_t kDroppedByBucketCap
Values of SegmentEdgeGraph::inputNodeIndex for input segments without a node.
const xAOD::TruthParticle * getTruthMatchedParticle(const xAOD::MuonSegment &segment)
Returns the particle truth-matched to the segment.
float j(const xAOD::IParticle &, const xAOD::TrackMeasurementValidation &hit, const Eigen::Matrix3d &jab_inv)
TruthParticle_v1 TruthParticle
Typedef to implementation.

◆ finalize()

StatusCode MuonML::SegmentEdgeClassifierTool::finalize ( )
override

Log a pre-ONNX candidate-edge pruning.

Definition at line 689 of file SegmentEdgeClassifierTool.cxx.

689 {
691 "SegmentEdgeClassifierTool pre-ONNX pruning summary (job-summed, "
692 "independent of PairGateThreshold): "
693 << "inputSegments=" << m_sumInputSegments
694 << ", candidatePairs=" << m_sumCandidatePairs
695 << " (geometric pairs before MaxEdgesPerNodeBeforeInference/"
696 "MaxEdgesPerTargetChamberBeforeInference caps)"
697 << ", retainedPairs=" << m_sumRetainedPairs
698 << " (pairs actually sent to ONNX; MaxEdgesPerNodeBeforeInference="
700 << ", MaxEdgesPerTargetChamberBeforeInference="
702 << ", nodesBeforeIsolatedDrop=" << m_sumNodesBeforeIsolatedDrop
703 << ", nodesAfterIsolatedDrop=" << m_sumNodesAfterIsolatedDrop
704 << " (DropIsolatedNodesBeforeInference="
706 << "; nodes dropped here never reached ONNX or the pair-gate threshold)");
707 return StatusCode::SUCCESS;
708}

◆ initialize()

StatusCode MuonML::SegmentEdgeClassifierTool::initialize ( )
override

Retrieve the ONNX model and resolve node feature ordering from metadata.

Definition at line 101 of file SegmentEdgeClassifierTool.cxx.

101 {
102 if (m_sectorModulo.value() > 0 &&
103 2ULL * static_cast<unsigned long long>(m_maxDeltaSector.value()) + 1ULL >
104 static_cast<unsigned long long>(m_sectorModulo.value())) {
105 ATH_MSG_ERROR("MaxDeltaSector=" << m_maxDeltaSector.value()
106 << " spans duplicate sectors for SectorModulo="
107 << m_sectorModulo.value());
108 return StatusCode::FAILURE;
109 }
111
112 // Resolve node feature names from model metadata, matching the ONNX exporter.
113 {
114 Ort::AllocatorWithDefaultOptions allocator;
115 Ort::ModelMetadata meta = model().GetModelMetadata();
116 auto keys = meta.GetCustomMetadataMapKeysAllocated(allocator);
117 std::vector<std::string> keyList;
118 keyList.reserve(keys.size());
119 for (const auto& k : keys) keyList.emplace_back(k.get());
120
121 constexpr std::array<std::string_view, 4> candidates{
122 "x_feature_names", "node_feature_names", "feature_names", "input_feature_names"};
123 std::string usedKey;
124 std::vector<std::string> names;
125 for (std::string_view key : candidates) {
126 const std::string keyStr{key};
127 if (std::find(keyList.begin(), keyList.end(), keyStr) == keyList.end()) continue;
128 names = parseFeatureNames(meta.LookupCustomMetadataMapAllocated(keyStr.c_str(), allocator).get());
129 if (!names.empty()) {
130 usedKey = keyStr;
131 break;
132 }
133 }
134
135 if (names.empty()) {
137 ATH_MSG_WARNING("Model metadata has no usable node feature name key"
138 " (tried x_feature_names/node_feature_names/feature_names/input_feature_names)."
139 " Falling back to default training order.");
140 } else {
141 if (names.size() != kNodeFeatureCount) {
142 ATH_MSG_ERROR("Model metadata key '" << usedKey << "' has " << names.size()
143 << " features, expected " << kNodeFeatureCount);
144 return StatusCode::FAILURE;
145 }
146 for (const std::string& n : names) {
147 if (!nodeFeatureIdFromName(n).has_value()) {
148 ATH_MSG_ERROR("Unsupported node feature name in model metadata ('" << usedKey
149 << "'): '" << n << "'."
150 " Add mapping in SegmentEdgeClassifierTool::nodeFeatureValue().");
151 return StatusCode::FAILURE;
152 }
153 }
154 m_nodeFeatureNames = std::move(names);
155 ATH_MSG_DEBUG("Using node feature names from model metadata key '" << usedKey << "'.");
156 }
157
158 m_nodeFeatureIds.reserve(m_nodeFeatureNames.size());
159 for (const std::string& n : m_nodeFeatureNames) {
160 const auto id = nodeFeatureIdFromName(n);
161 if (!id.has_value()) {
162 ATH_MSG_ERROR("Internal feature-id resolution failed for node feature name '" << n << "'.");
163 return StatusCode::FAILURE;
164 }
165 m_nodeFeatureIds.push_back(*id);
166 }
167
168 std::ostringstream order;
169 order << "Node feature order:";
170 for (std::size_t i = 0; i < m_nodeFeatureNames.size(); ++i) {
171 order << " f" << i << "=" << m_nodeFeatureNames[i];
172 if (i + 1 < m_nodeFeatureNames.size()) order << ",";
173 }
174 ATH_MSG_DEBUG(order.str());
175 }
176
178 ATH_MSG_ERROR("Internal node feature setup has " << m_nodeFeatureNames.size()
179 << " entries, expected " << kNodeFeatureCount);
180 return StatusCode::FAILURE;
181 }
182 if (m_nodeFeatureIds.size() != kNodeFeatureCount) {
183 ATH_MSG_ERROR("Internal node feature id setup has " << m_nodeFeatureIds.size()
184 << " entries, expected " << kNodeFeatureCount);
185 return StatusCode::FAILURE;
186 }
187
188 m_cosMin = std::cos(m_maxDeltaThetaDeg.value() * Gaudi::Units::deg);
189
190 if (!m_debugDumpFile.value().empty()) {
191 std::ofstream out{m_debugDumpFile.value(), std::ios::out | std::ios::trunc};
192 if (!out) {
193 ATH_MSG_ERROR("Could not create segment-edge debug dump file: "
194 << m_debugDumpFile.value());
195 return StatusCode::FAILURE;
196 }
197
198 nlohmann::ordered_json metadata;
199 metadata["record_type"] = "metadata";
200 metadata["format_version"] = 1;
201 metadata["tool"] = "SegmentEdgeClassifierTool";
202 metadata["input_names"] = {m_inputNodeName.value(),
203 m_inputEdgeIndexName.value(),
204 m_inputEdgeAttrName.value()};
205 metadata["output_name"] = m_outputName.value();
206 metadata["x_feature_names"] = m_nodeFeatureNames;
207 metadata["edge_attr_feature_names"] = {
208 "deltaPositionX_m", "deltaPositionY_m", "deltaPositionZ_m",
209 "distance_m", "cos_opening_angle", "same_chamber", "same_sector"};
210 metadata["edge_index_layout"] = "row_major_2_by_E";
211 metadata["edge_order"] = "directed src_to_dst; row 0 then row 1";
212 metadata["max_delta_theta_deg"] = m_maxDeltaThetaDeg.value();
213 metadata["max_delta_sector"] = m_maxDeltaSector.value();
214 metadata["sector_modulo"] = m_sectorModulo.value();
215 metadata["debug_dump_max_events"] = m_debugDumpMaxEvents.value();
216 out << metadata.dump() << '\n';
217
218 ATH_MSG_INFO("Writing segment-edge ONNX debug dump to "
219 << m_debugDumpFile.value()
220 << " (DebugDumpMaxEvents="
221 << m_debugDumpMaxEvents.value() << ")");
222 }
223
225
226 return StatusCode::SUCCESS;
227}
#define ATH_MSG_INFO(x,...)
static constexpr std::array< std::string_view, kNodeFeatureCount > kDefaultNodeFeatureNames
static std::vector< std::string > parseFeatureNames(const std::string &raw)
SG::ReadDecorHandleKey< xAOD::MuonSegmentContainer > m_truthLinkKey
std::vector< std::string > m_nodeFeatureNames
Node feature order expected by the model metadata (resolved at initialize).
order
Configure Herwig7.

◆ model()

Ort::Session & BucketInferenceToolBase::model ( ) const
protectedinherited

Definition at line 65 of file BucketInferenceToolBase.cxx.

65 {
66 return m_onnxSessionTool->session();
67}
ToolHandle< AthOnnx::IOnnxRuntimeSessionTool > m_onnxSessionTool

◆ parseFeatureNames()

std::vector< std::string > BucketInferenceToolBase::parseFeatureNames ( const std::string & raw)
staticprotectedinherited

Definition at line 32 of file BucketInferenceToolBase.cxx.

32 {
33 std::vector<std::string> out;
34 const std::string s = trimFeatureToken(raw);
35 if (s.empty()) return out;
36
37 // Preferred exporter format: JSON list of strings.
38 if (!s.empty() && s.front() == '[') {
39 bool inQuote = false;
40 std::string token;
41 for (char c : s) {
42 if (c == '"') {
43 if (inQuote) {
44 if (!token.empty()) out.push_back(token);
45 token.clear();
46 }
47 inQuote = !inQuote;
48 continue;
49 }
50 if (inQuote) token.push_back(c);
51 }
52 if (!out.empty()) return out;
53 }
54
55 // Backward-compatible format: comma-separated.
56 std::istringstream ss(s);
57 std::string tok;
58 while (std::getline(ss, tok, ',')) {
59 tok = trimFeatureToken(tok);
60 if (!tok.empty()) out.push_back(tok);
61 }
62 return out;
63}
static Double_t ss
static std::string trimFeatureToken(std::string s)

◆ runGraphInference()

StatusCode MuonML::SegmentEdgeClassifierTool::runGraphInference ( const EventContext & ctx,
GraphRawData & graphData ) const
override

Not supported by this tool; returns FAILURE.

Use SegmentEdgeInferenceAlg + buildGraph() + classifyEdges() instead.

Definition at line 229 of file SegmentEdgeClassifierTool.cxx.

229 {
230 ATH_MSG_ERROR("runGraphInference is not supported by SegmentEdgeClassifierTool. Use SegmentEdgeInferenceAlg + ISegmentEdgeClassifierTool methods.");
231 return StatusCode::FAILURE;
232}

◆ runInference()

StatusCode BucketInferenceToolBase::runInference ( GraphRawData & graphData) const
inherited

Default ONNX run for GNN case: inputs {"features","edge_index"} -> outputs {"logits"}.

Definition at line 532 of file BucketInferenceToolBase.cxx.

532 {
533 std::vector<const char*> inputNames = {"features", "edge_index"};
534 std::vector<const char*> outputNames = {m_outputName.value().c_str()};
535 return runNamedInference(graphData, inputNames, outputNames);
536}
Gaudi::Property< std::string > m_outputName

◆ runNamedInference()

StatusCode BucketInferenceToolBase::runNamedInference ( GraphRawData & graphData,
const std::vector< const char * > & inputNames,
const std::vector< const char * > & outputNames ) const
protectedinherited

Generic named inference, for tools with different I/O conventions.

Definition at line 359 of file BucketInferenceToolBase.cxx.

363{
364 if (!graphData.graph) {
365 ATH_MSG_ERROR("Graph data is not built.");
366 return StatusCode::FAILURE;
367 }
368 if (graphData.graph->dataTensor.empty()) {
369 ATH_MSG_ERROR("No input tensors prepared for inference.");
370 return StatusCode::FAILURE;
371 }
372
373 // Reserve the final size here from the actual I/O lists instead
374 // of hard-coding assumptions in the graph builders.
375 graphData.graph->dataTensor.reserve(inputNames.size() + outputNames.size());
376 if (graphData.graph->dataTensor.size() < inputNames.size()) {
377 ATH_MSG_ERROR("Prepared " << graphData.graph->dataTensor.size()
378 << " tensors but inference expects " << inputNames.size() << " inputs.");
379 return StatusCode::FAILURE;
380 }
381
382 if (msgLvl(MSG::DEBUG)) {
383 // DEBUG: Print actual input tensor data for features tensor
384
385 ATH_MSG_DEBUG("=== DEBUGGING: ONNX Input tensor data ===");
386 if (!graphData.graph->dataTensor.empty()) {
387 const auto& featureTensor = graphData.graph->dataTensor[0];
388 auto featShape = featureTensor.GetTensorTypeAndShapeInfo().GetShape();
389 ATH_MSG_DEBUG("Features tensor shape: [" << featShape[0]
390 << (featShape.size()>1 ? ("," + std::to_string(featShape[1])) : "")
391 << (featShape.size()>2 ? ("," + std::to_string(featShape[2])) : "") << "]");
392
393 float* featData = const_cast<Ort::Value&>(featureTensor).GetTensorMutableData<float>();
394 const size_t totalElements = featureTensor.GetTensorTypeAndShapeInfo().GetElementCount();
395 ATH_MSG_DEBUG("Features tensor total elements: " << totalElements);
396
397 // Print up to 10 nodes; stride = nFeat from tensor shape
398 const size_t nFeat = (featShape.size() > 1 && featShape[1] > 0) ? static_cast<size_t>(featShape[1]) : 1;
399 const size_t nNodes = totalElements / nFeat;
400 const size_t debugNodes = std::min(nNodes, static_cast<size_t>(10));
401
402 // Try to read feature names from model custom metadata.
403 // Prefer x_feature_names (current exporter), then fall back to legacy keys.
404 std::vector<std::string> featNames;
405 {
406 Ort::AllocatorWithDefaultOptions allocator;
407 Ort::ModelMetadata meta = model().GetModelMetadata();
408 auto keys = meta.GetCustomMetadataMapKeysAllocated(allocator);
409 std::vector<std::string> keyNames;
410 keyNames.reserve(keys.size());
411 for (const auto& k : keys) keyNames.emplace_back(k.get());
412 const std::array<std::string, 4> candidates{
413 "x_feature_names", "node_feature_names", "feature_names", "input_feature_names"};
414 for (const std::string& key : candidates) {
415 if (std::find(keyNames.begin(), keyNames.end(), key) != keyNames.end()) {
416 std::string val = meta.LookupCustomMetadataMapAllocated(key.c_str(), allocator).get();
417 featNames = parseFeatureNames(val);
418 break;
419 }
420 }
421 if (featNames.empty()) {
422 ATH_MSG_DEBUG("No usable feature-name metadata key found in model; using generic fN labels.");
423 }
424 }
425 auto featLabel = [&](size_t f) -> std::string {
426 if (f < featNames.size()) return featNames[f];
427 return "f" + std::to_string(f);
428 };
429
430 // Print legend
431 {
432 std::ostringstream legend;
433 legend << "Node feature legend (" << nFeat << " features):";
434 for (size_t f = 0; f < nFeat; ++f) {
435 legend << " f" << f << "=" << featLabel(f);
436 if (f + 1 < nFeat) legend << ",";
437 }
438 ATH_MSG_DEBUG(legend.str());
439 }
440
441 for (size_t n = 0; n < debugNodes; ++n) {
442 std::ostringstream row;
443 row << "ONNXNode[" << n << "]:";
444 for (size_t f = 0; f < nFeat; ++f) {
445 row << " f" << f << "=" << featData[n * nFeat + f];
446 if (f + 1 < nFeat) row << ",";
447 }
448 ATH_MSG_DEBUG(row.str());
449 }
450 }
451 ATH_MSG_DEBUG("=== END DEBUG ONNX INPUT ===");
452 }
453
454 Ort::RunOptions run_options;
455 run_options.SetRunLogSeverityLevel(ORT_LOGGING_LEVEL_ERROR);
456
457 if (m_isCuda) {
458 // ---- CUDA path: use IoBinding so tensors stay on device ----
459 Ort::IoBinding binding(model());
460 for (std::size_t i = 0; i < inputNames.size(); ++i) {
461 binding.BindInput(inputNames[i], graphData.graph->dataTensor[i]);
462 }
463 // Bind outputs to CPU so predictions are directly readable after sync.
464 Ort::MemoryInfo cpuOut = Ort::MemoryInfo::CreateCpu(OrtArenaAllocator, OrtMemTypeDefault);
465 for (const char* outName : outputNames) {
466 binding.BindOutput(outName, cpuOut);
467 }
468
469 model().Run(run_options, binding);
470 binding.SynchronizeOutputs();
471
472 std::vector<Ort::Value> outputs = binding.GetOutputValues();
473 if (outputs.empty()) {
474 ATH_MSG_ERROR("IoBinding inference returned empty output.");
475 return StatusCode::FAILURE;
476 }
477
478 float* outData = outputs[0].GetTensorMutableData<float>();
479 const size_t outSize = outputs[0].GetTensorTypeAndShapeInfo().GetElementCount();
480 ATH_MSG_DEBUG("ONNX (IoBinding) raw output elementCount = " << outSize);
481
482 if (m_sanitizeNonFinitePredictions.value()) {
483 std::span<float> preds(outData, outData + outSize);
484 for (size_t i = 0; i < outSize; ++i) {
485 if (!std::isfinite(preds[i])) {
486 ATH_MSG_WARNING("Non-finite prediction detected at " << i << " -> set to -100.");
487 preds[i] = -100.0f;
488 }
489 }
490 }
491
492 for (auto& v : outputs) {
493 graphData.graph->dataTensor.emplace_back(std::move(v));
494 }
495 return StatusCode::SUCCESS;
496 }
497
498 // ---- CPU path ----
499 std::vector<Ort::Value> outputs =
500 model().Run(run_options,
501 inputNames.data(),
502 graphData.graph->dataTensor.data(),
503 inputNames.size(),
504 outputNames.data(),
505 outputNames.size());
506
507 if (outputs.empty()) {
508 ATH_MSG_ERROR("Inference returned empty output.");
509 return StatusCode::FAILURE;
510 }
511
512 float* outData = outputs[0].GetTensorMutableData<float>();
513 const size_t outSize = outputs[0].GetTensorTypeAndShapeInfo().GetElementCount();
514 ATH_MSG_DEBUG("ONNX raw output elementCount = " << outSize);
515
516 if (m_sanitizeNonFinitePredictions.value()) {
517 std::span<float> preds(outData, outData + outSize);
518 for (size_t i = 0; i < outSize; ++i) {
519 if (!std::isfinite(preds[i])) {
520 ATH_MSG_WARNING("Non-finite prediction detected at " << i << " -> set to -100.");
521 preds[i] = -100.0f;
522 }
523 }
524 }
525
526 for (auto& v : outputs) {
527 graphData.graph->dataTensor.emplace_back(std::move(v));
528 }
529 return StatusCode::SUCCESS;
530}
Gaudi::Property< bool > m_sanitizeNonFinitePredictions
row
Appending html table to final .html summary file.

◆ setupModel()

StatusCode BucketInferenceToolBase::setupModel ( )
protectedinherited

Definition at line 69 of file BucketInferenceToolBase.cxx.

69 {
70 ATH_CHECK(m_onnxSessionTool.retrieve());
71 ATH_CHECK(m_readKey.initialize());
72 ATH_CHECK(m_geoCtxKey.initialize());
73
74 const InferenceUtils::SessionBackend backend = InferenceUtils::sessionBackend(m_onnxSessionTool);
75 m_isCuda = backend.isCuda;
76 m_cudaDeviceId = backend.cudaDeviceId;
77 if (m_isCuda) {
78 ATH_MSG_INFO("ONNX session is running on CUDA device " << m_cudaDeviceId
79 << ". I/O binding will be used.");
80 } else {
81 ATH_MSG_INFO("ONNX session is running on CPU.");
82 }
83
84 return StatusCode::SUCCESS;
85}
SessionBackend sessionBackend(const SessionToolHandle &sessionTool)

◆ trimFeatureToken()

std::string BucketInferenceToolBase::trimFeatureToken ( std::string s)
staticprotectedinherited

Definition at line 25 of file BucketInferenceToolBase.cxx.

25 {
26 auto notSpace = [](unsigned char c) { return !std::isspace(c); };
27 s.erase(s.begin(), std::find_if(s.begin(), s.end(), notSpace));
28 s.erase(std::find_if(s.rbegin(), s.rend(), notSpace).base(), s.end());
29 return s;
30}

Member Data Documentation

◆ kBucketFeatureCount

std::size_t MuonML::BucketInferenceToolBase::kBucketFeatureCount = 6
staticconstexprprotectedinherited

Definition at line 53 of file BucketInferenceToolBase.h.

◆ kDefaultNodeFeatureNames

std::array<std::string_view, kNodeFeatureCount> MuonML::BucketInferenceToolBase::kDefaultNodeFeatureNames
staticconstexprprotectedinherited
Initial value:
= {
"segmentPositionX_m", "segmentPositionY_m", "segmentPositionZ_m",
"segmentDirectionX", "segmentDirectionY", "segmentDirectionZ",
"bucket_chamberIndex", "bucket_layers", "bucket_sector", "bucket_segments"}

Definition at line 56 of file BucketInferenceToolBase.h.

56 {
57 "segmentPositionX_m", "segmentPositionY_m", "segmentPositionZ_m",
58 "segmentDirectionX", "segmentDirectionY", "segmentDirectionZ",
59 "bucket_chamberIndex", "bucket_layers", "bucket_sector", "bucket_segments"};

◆ kEdgeFeatureCount

std::size_t MuonML::BucketInferenceToolBase::kEdgeFeatureCount = 7
staticconstexprprotectedinherited

Definition at line 55 of file BucketInferenceToolBase.h.

◆ kNodeFeatureCount

std::size_t MuonML::BucketInferenceToolBase::kNodeFeatureCount = 10
staticconstexprprotectedinherited

Definition at line 54 of file BucketInferenceToolBase.h.

◆ m_cosMin

float MuonML::SegmentEdgeClassifierTool::m_cosMin {0.f}
private

Definition at line 134 of file SegmentEdgeClassifierTool.h.

134{0.f};

◆ m_cudaDeviceId

int MuonML::BucketInferenceToolBase::m_cudaDeviceId {0}
protectedinherited

Definition at line 102 of file BucketInferenceToolBase.h.

102{0};

◆ m_debugDumpEvents

std::atomic<unsigned int> MuonML::SegmentEdgeClassifierTool::m_debugDumpEvents {0}
mutableprivate

Definition at line 140 of file SegmentEdgeClassifierTool.h.

140{0};

◆ m_debugDumpFile

Gaudi::Property<std::string> MuonML::SegmentEdgeClassifierTool::m_debugDumpFile {this, "DebugDumpFile", ""}
private

Definition at line 127 of file SegmentEdgeClassifierTool.h.

127{this, "DebugDumpFile", ""};

◆ m_debugDumpFirstNEdges

Gaudi::Property<unsigned int> MuonML::BucketInferenceToolBase::m_debugDumpFirstNEdges {this, "DebugDumpFirstNEdges", 12}
protectedinherited

Definition at line 94 of file BucketInferenceToolBase.h.

94{this, "DebugDumpFirstNEdges", 12};

◆ m_debugDumpFirstNNodes

Gaudi::Property<unsigned int> MuonML::BucketInferenceToolBase::m_debugDumpFirstNNodes {this, "DebugDumpFirstNNodes", 5}
protectedinherited

Definition at line 93 of file BucketInferenceToolBase.h.

93{this, "DebugDumpFirstNNodes", 5};

◆ m_debugDumpMaxEvents

Gaudi::Property<unsigned int> MuonML::SegmentEdgeClassifierTool::m_debugDumpMaxEvents {this, "DebugDumpMaxEvents", 0}
private

Definition at line 128 of file SegmentEdgeClassifierTool.h.

128{this, "DebugDumpMaxEvents", 0};

◆ m_debugDumpMutex

std::mutex MuonML::SegmentEdgeClassifierTool::m_debugDumpMutex
mutableprivate

Definition at line 139 of file SegmentEdgeClassifierTool.h.

◆ m_dropIsolatedNodesBeforeInference

Gaudi::Property<bool> MuonML::SegmentEdgeClassifierTool::m_dropIsolatedNodesBeforeInference
private
Initial value:
{this, "DropIsolatedNodesBeforeInference", true,
"Remove nodes without a retained pre-ONNX edge before creating ONNX tensors"}

Definition at line 121 of file SegmentEdgeClassifierTool.h.

121 {this, "DropIsolatedNodesBeforeInference", true,
122 "Remove nodes without a retained pre-ONNX edge before creating ONNX tensors"};

◆ m_dropSameChamberEdgesBeforeInference

Gaudi::Property<bool> MuonML::SegmentEdgeClassifierTool::m_dropSameChamberEdgesBeforeInference
private
Initial value:
{this, "DropSameChamberEdgesBeforeInference", true,
"Drop same-chamber segment pairs before ONNX inference"}

Definition at line 119 of file SegmentEdgeClassifierTool.h.

119 {this, "DropSameChamberEdgesBeforeInference", true,
120 "Drop same-chamber segment pairs before ONNX inference"};

◆ m_enableTruthDiagnostics

Gaudi::Property<bool> MuonML::SegmentEdgeClassifierTool::m_enableTruthDiagnostics
private
Initial value:
{
this, "EnableTruthDiagnostics", false,
"MC-only: fill SegmentEdgeGraph's truth diagnostics."}

Definition at line 129 of file SegmentEdgeClassifierTool.h.

129 {
130 this, "EnableTruthDiagnostics", false,
131 "MC-only: fill SegmentEdgeGraph's truth diagnostics."};

◆ m_geoCtxKey

ActsTrk::GeoContextReadKey_t MuonML::BucketInferenceToolBase::m_geoCtxKey {this, "AlignmentKey", "ActsAlignment", "cond handle key"}
protectedinherited

Definition at line 80 of file BucketInferenceToolBase.h.

80{this, "AlignmentKey", "ActsAlignment", "cond handle key"};

◆ m_inputEdgeAttrName

Gaudi::Property<std::string> MuonML::SegmentEdgeClassifierTool::m_inputEdgeAttrName {this, "InputEdgeAttrName", "edge_attr"}
private

Definition at line 125 of file SegmentEdgeClassifierTool.h.

125{this, "InputEdgeAttrName", "edge_attr"};

◆ m_inputEdgeIndexName

Gaudi::Property<std::string> MuonML::SegmentEdgeClassifierTool::m_inputEdgeIndexName {this, "InputEdgeIndexName", "edge_index"}
private

Definition at line 124 of file SegmentEdgeClassifierTool.h.

124{this, "InputEdgeIndexName", "edge_index"};

◆ m_inputNodeName

Gaudi::Property<std::string> MuonML::SegmentEdgeClassifierTool::m_inputNodeName {this, "InputNodeName", "x"}
private

Definition at line 123 of file SegmentEdgeClassifierTool.h.

123{this, "InputNodeName", "x"};

◆ m_isCuda

bool MuonML::BucketInferenceToolBase::m_isCuda {false}
protectedinherited

Definition at line 101 of file BucketInferenceToolBase.h.

101{false};

◆ m_maxAbsDz

Gaudi::Property<double> MuonML::BucketInferenceToolBase::m_maxAbsDz {this, "MaxAbsDz", 15000.0}
protectedinherited

Definition at line 90 of file BucketInferenceToolBase.h.

90{this, "MaxAbsDz", 15000.0};

◆ m_maxChamberDelta

Gaudi::Property<int> MuonML::BucketInferenceToolBase::m_maxChamberDelta {this, "MaxChamberDelta", 13}
protectedinherited

Definition at line 87 of file BucketInferenceToolBase.h.

87{this, "MaxChamberDelta", 13};

◆ m_maxDeltaSector

Gaudi::Property<int> MuonML::SegmentEdgeClassifierTool::m_maxDeltaSector {this, "MaxDeltaSector", 1}
private

Definition at line 109 of file SegmentEdgeClassifierTool.h.

109{this, "MaxDeltaSector", 1};

◆ m_maxDeltaThetaDeg

Gaudi::Property<float> MuonML::SegmentEdgeClassifierTool::m_maxDeltaThetaDeg {this, "MaxDeltaThetaDeg", 35.f}
private

Definition at line 108 of file SegmentEdgeClassifierTool.h.

108{this, "MaxDeltaThetaDeg", 35.f};

◆ m_maxDistXY

Gaudi::Property<double> MuonML::BucketInferenceToolBase::m_maxDistXY {this, "MaxDistXY", 6800.0}
protectedinherited

Definition at line 89 of file BucketInferenceToolBase.h.

89{this, "MaxDistXY", 6800.0};

◆ m_maxEdgesPerNodeBeforeInference

Gaudi::Property<unsigned int> MuonML::SegmentEdgeClassifierTool::m_maxEdgesPerNodeBeforeInference
private
Initial value:
{this, "MaxEdgesPerNodeBeforeInference", 0,
"Keep at most this many geometrical neighbour pairs per node before ONNX inference; 0 keeps all"}

Definition at line 114 of file SegmentEdgeClassifierTool.h.

114 {this, "MaxEdgesPerNodeBeforeInference", 0,
115 "Keep at most this many geometrical neighbour pairs per node before ONNX inference; 0 keeps all"};

◆ m_maxEdgesPerTargetChamberBeforeInference

Gaudi::Property<unsigned int> MuonML::SegmentEdgeClassifierTool::m_maxEdgesPerTargetChamberBeforeInference
private
Initial value:
{
this, "MaxEdgesPerTargetChamberBeforeInference", 0,
"Keep at most this many pre-ONNX neighbours from one target chamber per node; 0 keeps all"}

Definition at line 116 of file SegmentEdgeClassifierTool.h.

116 {
117 this, "MaxEdgesPerTargetChamberBeforeInference", 0,
118 "Keep at most this many pre-ONNX neighbours from one target chamber per node; 0 keeps all"};

◆ m_maxSectorDelta

Gaudi::Property<int> MuonML::BucketInferenceToolBase::m_maxSectorDelta {this, "MaxSectorDelta", 1}
protectedinherited

Definition at line 88 of file BucketInferenceToolBase.h.

88{this, "MaxSectorDelta", 1};

◆ m_maxSegmentsPerBucket

Gaudi::Property<unsigned int> MuonML::SegmentEdgeClassifierTool::m_maxSegmentsPerBucket
private
Initial value:
{this, "MaxSegmentsPerBucket", 0,
"Keep at most this many quality-ranked segments per (sector, chamber, eta) bucket before inference; 0 keeps all"}

Definition at line 112 of file SegmentEdgeClassifierTool.h.

112 {this, "MaxSegmentsPerBucket", 0,
113 "Keep at most this many quality-ranked segments per (sector, chamber, eta) bucket before inference; 0 keeps all"};

◆ m_minLayers

Gaudi::Property<int> MuonML::BucketInferenceToolBase::m_minLayers {this, "MinLayersValid", 3}
protectedinherited

Definition at line 86 of file BucketInferenceToolBase.h.

86{this, "MinLayersValid", 3};

◆ m_nodeFeatureIds

std::vector<SegmentNodeFeatureId> MuonML::SegmentEdgeClassifierTool::m_nodeFeatureIds {}
private

Definition at line 138 of file SegmentEdgeClassifierTool.h.

138{};

◆ m_nodeFeatureNames

std::vector<std::string> MuonML::SegmentEdgeClassifierTool::m_nodeFeatureNames {}
private

Node feature order expected by the model metadata (resolved at initialize).

Definition at line 137 of file SegmentEdgeClassifierTool.h.

137{};

◆ m_onnxSessionTool

ToolHandle<AthOnnx::IOnnxRuntimeSessionTool> MuonML::BucketInferenceToolBase::m_onnxSessionTool
privateinherited
Initial value:
{
this, "ModelSession", ""}

Definition at line 105 of file BucketInferenceToolBase.h.

105 {
106 this, "ModelSession", ""};

◆ m_outputName

Gaudi::Property<std::string> MuonML::SegmentEdgeClassifierTool::m_outputName {this, "OutputName", "logits"}
private

Definition at line 126 of file SegmentEdgeClassifierTool.h.

126{this, "OutputName", "logits"};

◆ m_readKey

SG::ReadHandleKey<MuonR4::SpacePointContainer> MuonML::BucketInferenceToolBase::m_readKey {this, "ReadSpacePoints", "MuonSpacePoints"}
protectedinherited

Definition at line 79 of file BucketInferenceToolBase.h.

79{this, "ReadSpacePoints", "MuonSpacePoints"};

◆ m_sanitizeNonFinitePredictions

Gaudi::Property<bool> MuonML::BucketInferenceToolBase::m_sanitizeNonFinitePredictions
protectedinherited
Initial value:
{
this, "SanitizeNonFinitePredictions", false,
"When true, replace non-finite ONNX outputs with -100 and log a warning."}

Definition at line 96 of file BucketInferenceToolBase.h.

96 {
97 this, "SanitizeNonFinitePredictions", false,
98 "When true, replace non-finite ONNX outputs with -100 and log a warning."};

◆ m_sectorModulo

Gaudi::Property<int> MuonML::SegmentEdgeClassifierTool::m_sectorModulo
private
Initial value:
{this, "SectorModulo", 16,
"Number of muon sectors used when applying wrap-around sector distance"}

Definition at line 110 of file SegmentEdgeClassifierTool.h.

110 {this, "SectorModulo", 16,
111 "Number of muon sectors used when applying wrap-around sector distance"};

◆ m_sumCandidatePairs

std::atomic<std::size_t> MuonML::SegmentEdgeClassifierTool::m_sumCandidatePairs {0}
mutableprivate

Definition at line 144 of file SegmentEdgeClassifierTool.h.

144{0};

◆ m_sumInputSegments

std::atomic<std::size_t> MuonML::SegmentEdgeClassifierTool::m_sumInputSegments {0}
mutableprivate

Job-summed pre-ONNX pruning counters (see buildGraph()).

Definition at line 143 of file SegmentEdgeClassifierTool.h.

143{0};

◆ m_sumNodesAfterIsolatedDrop

std::atomic<std::size_t> MuonML::SegmentEdgeClassifierTool::m_sumNodesAfterIsolatedDrop {0}
mutableprivate

Definition at line 147 of file SegmentEdgeClassifierTool.h.

147{0};

◆ m_sumNodesBeforeIsolatedDrop

std::atomic<std::size_t> MuonML::SegmentEdgeClassifierTool::m_sumNodesBeforeIsolatedDrop {0}
mutableprivate

Definition at line 146 of file SegmentEdgeClassifierTool.h.

146{0};

◆ m_sumRetainedPairs

std::atomic<std::size_t> MuonML::SegmentEdgeClassifierTool::m_sumRetainedPairs {0}
mutableprivate

Definition at line 145 of file SegmentEdgeClassifierTool.h.

145{0};

◆ m_truthLinkKey

SG::ReadDecorHandleKey<xAOD::MuonSegmentContainer> MuonML::SegmentEdgeClassifierTool::m_truthLinkKey
private
Initial value:
{
this, "TruthLinkKey", "MuonSegmentsFromR4.truthParticleLink"}

Definition at line 132 of file SegmentEdgeClassifierTool.h.

132 {
133 this, "TruthLinkKey", "MuonSegmentsFromR4.truthParticleLink"};

◆ m_validateEdges

Gaudi::Property<bool> MuonML::BucketInferenceToolBase::m_validateEdges {this, "ValidateEdges", true}
protectedinherited

Definition at line 95 of file BucketInferenceToolBase.h.

95{this, "ValidateEdges", true};

The documentation for this class was generated from the following files: