ATLAS Offline Software
Loading...
Searching...
No Matches
MuonML::BucketInferenceToolBase Class Reference

#include <BucketInferenceToolBase.h>

Inheritance diagram for MuonML::BucketInferenceToolBase:
Collaboration diagram for MuonML::BucketInferenceToolBase:

Public Member Functions

 ~BucketInferenceToolBase () override=default
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"}.

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< std::string > m_outputName {this, "OutputName", "logits"}
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 Attributes

ToolHandle< AthOnnx::IOnnxRuntimeSessionTool > m_onnxSessionTool

Detailed Description

BucketInferenceToolBase

Common infra to:

  • read buckets & (optionally) geometry
  • build node features
  • (optionally) build GNN sparse edges (via BucketGraphUtils)
  • wrap tensors and run ONNX sessions

GNN-specific operations are in BucketGraphUtils.* Transformer tools reuse feature building without edges and add a pad mask.

Definition at line 40 of file BucketInferenceToolBase.h.

Constructor & Destructor Documentation

◆ ~BucketInferenceToolBase()

MuonML::BucketInferenceToolBase::~BucketInferenceToolBase ( )
overridedefault

Member Function Documentation

◆ buildFeaturesOnly()

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

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

Definition at line 86 of file BucketInferenceToolBase.cxx.

87 {
88
89 graphData.graph.reset();
90 graphData.srcEdges.clear();
91 graphData.desEdges.clear();
92 graphData.edgeIndexPacked.clear();
93 graphData.featureLeaves.clear();
94 graphData.spacePointsInBucket.clear();
95 graphData.graph = std::make_unique<InferenceGraph>();
96 graphData.graph->dataTensor.reserve(1); // features input; outputs are reserved in runNamedInference()
97
98 const MuonR4::SpacePointContainer* buckets{nullptr};
99 ATH_CHECK(SG::get(buckets, m_readKey, ctx));
100
101 const ActsTrk::GeometryContext* gctx = nullptr;
102 ATH_CHECK(SG::get(gctx, m_geoCtxKey, ctx));
103
104 std::vector<BucketGraphUtils::NodeAux> nodes;
105 BucketGraphUtils::buildNodesAndFeatures(*buckets, *gctx, nodes,
106 graphData.featureLeaves,
107 graphData.spacePointsInBucket); // now int64_t-compatible
108
109 const int64_t numNodes = static_cast<int64_t>(nodes.size());
110 ATH_MSG_DEBUG("Total buckets: " << buckets->size()
111 << " -> nodes (size>0): " << numNodes
112 << " | features.size()=" << graphData.featureLeaves.size());
113
114 if (numNodes == 0) {
115 ATH_MSG_WARNING("No valid buckets found (all have size 0.0). Skipping inference.");
116 return StatusCode::SUCCESS;
117 }
118
119 const int64_t nFeatPerNode = static_cast<int64_t>(kBucketFeatureCount);
120 if (numNodes * nFeatPerNode != static_cast<int64_t>(graphData.featureLeaves.size())) {
121 ATH_MSG_ERROR( "Feature size mismatch: expected " << (numNodes * nFeatPerNode)
122 << " got " << graphData.featureLeaves.size());
123 return StatusCode::FAILURE;
124 }
125
126 Ort::MemoryInfo memInfo = Ort::MemoryInfo::CreateCpu(OrtDeviceAllocator, OrtMemTypeCPU);
127 std::vector<int64_t> featShape{numNodes, nFeatPerNode};
128 graphData.graph->dataTensor.emplace_back(
129 Ort::Value::CreateTensor<float>(memInfo,
130 graphData.featureLeaves.data(),
131 graphData.featureLeaves.size(),
132 featShape.data(),
133 featShape.size()));
134 return StatusCode::SUCCESS;
135}
#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()

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

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

Definition at line 199 of file BucketInferenceToolBase.cxx.

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

◆ buildTransformerInputs()

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

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

Definition at line 137 of file BucketInferenceToolBase.cxx.

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

◆ model()

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

Definition at line 64 of file BucketInferenceToolBase.cxx.

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

◆ parseFeatureNames()

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

Definition at line 31 of file BucketInferenceToolBase.cxx.

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

◆ runInference()

StatusCode BucketInferenceToolBase::runInference ( GraphRawData & graphData) const

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

Definition at line 531 of file BucketInferenceToolBase.cxx.

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

◆ runNamedInference()

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

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

Definition at line 358 of file BucketInferenceToolBase.cxx.

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

◆ setupModel()

StatusCode BucketInferenceToolBase::setupModel ( )
protected

Definition at line 68 of file BucketInferenceToolBase.cxx.

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

◆ trimFeatureToken()

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

Definition at line 24 of file BucketInferenceToolBase.cxx.

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

Member Data Documentation

◆ kBucketFeatureCount

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

Definition at line 52 of file BucketInferenceToolBase.h.

◆ kDefaultNodeFeatureNames

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

Definition at line 55 of file BucketInferenceToolBase.h.

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

◆ kEdgeFeatureCount

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

Definition at line 54 of file BucketInferenceToolBase.h.

◆ kNodeFeatureCount

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

Definition at line 53 of file BucketInferenceToolBase.h.

◆ m_cudaDeviceId

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

Definition at line 101 of file BucketInferenceToolBase.h.

101{0};

◆ m_debugDumpFirstNEdges

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

Definition at line 93 of file BucketInferenceToolBase.h.

93{this, "DebugDumpFirstNEdges", 12};

◆ m_debugDumpFirstNNodes

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

Definition at line 92 of file BucketInferenceToolBase.h.

92{this, "DebugDumpFirstNNodes", 5};

◆ m_geoCtxKey

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

Definition at line 79 of file BucketInferenceToolBase.h.

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

◆ m_isCuda

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

Definition at line 100 of file BucketInferenceToolBase.h.

100{false};

◆ m_maxAbsDz

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

Definition at line 89 of file BucketInferenceToolBase.h.

89{this, "MaxAbsDz", 15000.0};

◆ m_maxChamberDelta

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

Definition at line 86 of file BucketInferenceToolBase.h.

86{this, "MaxChamberDelta", 13};

◆ m_maxDistXY

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

Definition at line 88 of file BucketInferenceToolBase.h.

88{this, "MaxDistXY", 6800.0};

◆ m_maxSectorDelta

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

Definition at line 87 of file BucketInferenceToolBase.h.

87{this, "MaxSectorDelta", 1};

◆ m_minLayers

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

Definition at line 85 of file BucketInferenceToolBase.h.

85{this, "MinLayersValid", 3};

◆ m_onnxSessionTool

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

Definition at line 104 of file BucketInferenceToolBase.h.

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

◆ m_outputName

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

Definition at line 82 of file BucketInferenceToolBase.h.

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

◆ m_readKey

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

Definition at line 78 of file BucketInferenceToolBase.h.

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

◆ m_sanitizeNonFinitePredictions

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

Definition at line 95 of file BucketInferenceToolBase.h.

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

◆ m_validateEdges

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

Definition at line 94 of file BucketInferenceToolBase.h.

94{this, "ValidateEdges", true};

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