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());
32 std::vector<std::string> out;
34 if (s.empty())
return out;
37 if (!s.empty() && s.front() ==
'[') {
43 if (!token.empty()) out.push_back(token);
49 if (inQuote) token.push_back(c);
51 if (!out.empty())
return out;
55 std::istringstream
ss(s);
57 while (std::getline(
ss, tok,
',')) {
59 if (!tok.empty()) out.push_back(tok);
78 <<
". I/O binding will be used.");
83 return StatusCode::SUCCESS;
89 graphData.
graph.reset();
95 graphData.
graph = std::make_unique<InferenceGraph>();
96 graphData.
graph->dataTensor.reserve(1);
104 std::vector<BucketGraphUtils::NodeAux> nodes;
109 const int64_t numNodes =
static_cast<int64_t
>(nodes.size());
111 <<
" -> nodes (size>0): " << numNodes
115 ATH_MSG_WARNING(
"No valid buckets found (all have size 0.0). Skipping inference.");
116 return StatusCode::SUCCESS;
120 if (numNodes * nFeatPerNode !=
static_cast<int64_t
>(graphData.
featureLeaves.size())) {
121 ATH_MSG_ERROR(
"Feature size mismatch: expected " << (numNodes * nFeatPerNode)
123 return StatusCode::FAILURE;
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,
134 return StatusCode::SUCCESS;
147 ATH_MSG_WARNING(
"No valid features for transformer input. Skipping inference.");
148 return StatusCode::SUCCESS;
151 if (msgLvl(MSG::DEBUG)) {
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) {
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]);
169 graphData.
graph.reset();
170 graphData.
graph = std::make_unique<InferenceGraph>();
171 graphData.
graph->dataTensor.reserve(2);
173 Ort::MemoryInfo memInfo = Ort::MemoryInfo::CreateCpu(OrtDeviceAllocator, OrtMemTypeCPU);
178 graphData.
graph->dataTensor.emplace_back(
179 Ort::Value::CreateTensor<float>(memInfo,
186 Ort::AllocatorWithDefaultOptions allocator;
187 std::vector<int64_t> mShape{1, S};
188 Ort::Value padVal = Ort::Value::CreateTensor(allocator,
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));
196 return StatusCode::SUCCESS;
202 graphData.
graph.reset();
208 graphData.
graph = std::make_unique<InferenceGraph>();
209 graphData.
graph->dataTensor.reserve(2);
217 std::vector<BucketGraphUtils::NodeAux> nodes;
223 const int64_t numNodes =
static_cast<int64_t
>(nodes.size());
225 <<
" -> nodes (size>0): " << numNodes
229 ATH_MSG_WARNING(
"No valid buckets found (all have size 0.0). Skipping graph building.");
230 return StatusCode::SUCCESS;
234 if (numNodes * nFeatPerNode !=
static_cast<int64_t
>(graphData.
featureLeaves.size())) {
235 ATH_MSG_ERROR(
"Feature size mismatch: expected " << (numNodes * nFeatPerNode)
237 return StatusCode::FAILURE;
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,
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);
270 ATH_MSG_DEBUG(
"Drop invalid edge " << k <<
": (" << u <<
"->" << v
271 <<
"), valid node range [0," << (numNodes-1) <<
"]");
282 const size_t E = graphData.
srcEdges.size();
284 if (msgLvl(MSG::DEBUG)) {
288 for (
size_t k = 0; k < dumpE; ++k) {
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]++;
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]);
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;
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];
320 if (foundAny) connections <<
", ";
323 }
else if (v == nodeIdx) {
324 if (foundAny) connections <<
", ";
330 if (!foundAny) connections <<
"none";
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,
354 ATH_MSG_DEBUG(
"Built sparse bucket graph: N=" << numNodes <<
", E=" << Efinal);
355 return StatusCode::SUCCESS;
360 const std::vector<const char*>& inputNames,
361 const std::vector<const char*>& outputNames)
const
363 if (!graphData.
graph) {
365 return StatusCode::FAILURE;
367 if (graphData.
graph->dataTensor.empty()) {
369 return StatusCode::FAILURE;
374 graphData.
graph->dataTensor.reserve(inputNames.size() + outputNames.size());
375 if (graphData.
graph->dataTensor.size() < inputNames.size()) {
377 <<
" tensors but inference expects " << inputNames.size() <<
" inputs.");
378 return StatusCode::FAILURE;
381 if (msgLvl(MSG::DEBUG)) {
385 if (!graphData.
graph->dataTensor.empty()) {
386 const auto& featureTensor = graphData.
graph->dataTensor[0];
387 auto featShape = featureTensor.GetTensorTypeAndShapeInfo().GetShape();
389 << (featShape.size()>1 ? (
"," + std::to_string(featShape[1])) :
"")
390 << (featShape.size()>2 ? (
"," + std::to_string(featShape[2])) :
"") <<
"]");
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);
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));
403 std::vector<std::string> featNames;
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();
420 if (featNames.empty()) {
421 ATH_MSG_DEBUG(
"No usable feature-name metadata key found in model; using generic fN labels.");
424 auto featLabel = [&](
size_t f) -> std::string {
425 if (f < featNames.size())
return featNames[f];
426 return "f" + std::to_string(f);
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 <<
",";
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 <<
",";
453 Ort::RunOptions run_options;
454 run_options.SetRunLogSeverityLevel(ORT_LOGGING_LEVEL_ERROR);
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]);
463 Ort::MemoryInfo cpuOut = Ort::MemoryInfo::CreateCpu(OrtArenaAllocator, OrtMemTypeDefault);
464 for (
const char* outName : outputNames) {
465 binding.BindOutput(outName, cpuOut);
468 model().Run(run_options, binding);
469 binding.SynchronizeOutputs();
471 std::vector<Ort::Value> outputs = binding.GetOutputValues();
472 if (outputs.empty()) {
474 return StatusCode::FAILURE;
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);
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.");
491 for (
auto& v : outputs) {
492 graphData.
graph->dataTensor.emplace_back(std::move(v));
494 return StatusCode::SUCCESS;
498 std::vector<Ort::Value> outputs =
499 model().Run(run_options,
501 graphData.
graph->dataTensor.data(),
506 if (outputs.empty()) {
508 return StatusCode::FAILURE;
511 float* outData = outputs[0].GetTensorMutableData<
float>();
512 const size_t outSize = outputs[0].GetTensorTypeAndShapeInfo().GetElementCount();
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.");
525 for (
auto& v : outputs) {
526 graphData.
graph->dataTensor.emplace_back(std::move(v));
528 return StatusCode::SUCCESS;
532 std::vector<const char*> inputNames = {
"features",
"edge_index"};
533 std::vector<const char*> outputNames = {
m_outputName.value().c_str()};
#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,...)
#define ATH_MSG_INFO(x,...)
Handle class for reading from StoreGate.
size_type size() const noexcept
Returns the number of elements in the collection.
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.
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)
SessionBackend sessionBackend(const SessionToolHandle &sessionTool)
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.
Helper struct to ship the Graph from the space point buckets to ONNX.
FeatureVec_t featureLeaves
Vector containing all features.
EdgeCounterVec_t edgeIndexPacked
Packed edge index buffer (kept alive for ONNX tensors that reference it) This stores [srcEdges,...
std::unique_ptr< InferenceGraph > graph
Pointer to the graph to be parsed to ONNX.
EdgeCounterVec_t srcEdges
Vector encoding the source index of the.
EdgeCounterVec_t desEdges
Vect.
NodeConnectVec_t spacePointsInBucket
Vector keeping track of how many space points are in each parsed bucket.