10#include "Acts/Utilities/Helpers.hpp"
12#include "GaudiKernel/SystemOfUnits.h"
13#include <nlohmann/json.hpp>
23#include <unordered_map>
24#include <unordered_set>
27using SegmentGroupKey = std::tuple<int, int, int>;
38 std::vector<unsigned int> uniqueLayers;
39 uniqueLayers.reserve(bucket.size());
40 for (
const MuonR4::SpacePointBucket::value_type&
sp : bucket) {
41 const unsigned int layNum =
sorter.sectorLayerNum(*
sp);
42 if (!Acts::rangeContainsValue(uniqueLayers, layNum)) {
43 uniqueLayers.push_back(layNum);
46 return static_cast<int>(uniqueLayers.size());
54inline int sectorDistance(
int a,
int b,
int mod) {
55 int d = std::abs(
a - b);
56 return mod > 0 ? std::min(d, mod - d) :
d;
59std::optional<MuonML::SegmentNodeFeatureId> nodeFeatureIdFromName(
const std::string& name) {
61 if (name ==
"segmentPositionX_m")
return FeatureId::SegmentPositionX;
62 if (name ==
"segmentPositionY_m")
return FeatureId::SegmentPositionY;
63 if (name ==
"segmentPositionZ_m")
return FeatureId::SegmentPositionZ;
64 if (name ==
"segmentDirectionX")
return FeatureId::SegmentDirectionX;
65 if (name ==
"segmentDirectionY")
return FeatureId::SegmentDirectionY;
66 if (name ==
"segmentDirectionZ")
return FeatureId::SegmentDirectionZ;
67 if (name ==
"bucket_chamberIndex")
return FeatureId::BucketChamberIndex;
68 if (name ==
"bucket_layers")
return FeatureId::BucketLayers;
69 if (name ==
"bucket_sector")
return FeatureId::BucketSector;
70 if (name ==
"bucket_segments")
return FeatureId::BucketSegments;
80 case FeatureId::SegmentPositionX:
return static_cast<float>(
pos.x());
81 case FeatureId::SegmentPositionY:
return static_cast<float>(
pos.y());
82 case FeatureId::SegmentPositionZ:
return static_cast<float>(
pos.z());
83 case FeatureId::SegmentDirectionX:
return static_cast<float>(
dir.x());
84 case FeatureId::SegmentDirectionY:
return static_cast<float>(
dir.y());
85 case FeatureId::SegmentDirectionZ:
return static_cast<float>(
dir.z());
86 case FeatureId::BucketChamberIndex:
return static_cast<float>(bucket.
chamberIndex);
87 case FeatureId::BucketLayers:
return static_cast<float>(bucket.
layers);
88 case FeatureId::BucketSector:
return static_cast<float>(bucket.
sector);
89 case FeatureId::BucketSegments:
return static_cast<float>(bucket.
nSegments);
102 Ort::AllocatorWithDefaultOptions allocator;
103 Ort::ModelMetadata
meta =
model().GetModelMetadata();
104 auto keys =
meta.GetCustomMetadataMapKeysAllocated(allocator);
105 std::vector<std::string> keyList;
106 keyList.reserve(keys.size());
107 for (
const auto& k : keys) keyList.emplace_back(k.get());
109 constexpr std::array<std::string_view, 4> candidates{
110 "x_feature_names",
"node_feature_names",
"feature_names",
"input_feature_names"};
112 std::vector<std::string> names;
113 for (std::string_view key : candidates) {
114 const std::string keyStr{key};
115 if (std::find(keyList.begin(), keyList.end(), keyStr) == keyList.end())
continue;
117 if (!names.empty()) {
126 " (tried x_feature_names/node_feature_names/feature_names/input_feature_names)."
127 " Falling back to default training order.");
130 ATH_MSG_ERROR(
"Model metadata key '" << usedKey <<
"' has " << names.size()
132 return StatusCode::FAILURE;
134 for (
const std::string& n : names) {
135 if (!nodeFeatureIdFromName(n).has_value()) {
136 ATH_MSG_ERROR(
"Unsupported node feature name in model metadata ('" << usedKey
137 <<
"'): '" << n <<
"'."
138 " Add mapping in SegmentEdgeClassifierTool::nodeFeatureValue().");
139 return StatusCode::FAILURE;
143 ATH_MSG_DEBUG(
"Using node feature names from model metadata key '" << usedKey <<
"'.");
148 const auto id = nodeFeatureIdFromName(n);
149 if (!
id.has_value()) {
150 ATH_MSG_ERROR(
"Internal feature-id resolution failed for node feature name '" << n <<
"'.");
151 return StatusCode::FAILURE;
156 std::ostringstream order;
157 order <<
"Node feature order:";
168 return StatusCode::FAILURE;
173 return StatusCode::FAILURE;
179 std::ofstream out{
m_debugDumpFile.value(), std::ios::out | std::ios::trunc};
181 ATH_MSG_ERROR(
"Could not create segment-edge debug dump file: "
183 return StatusCode::FAILURE;
186 nlohmann::ordered_json metadata;
187 metadata[
"record_type"] =
"metadata";
188 metadata[
"format_version"] = 1;
189 metadata[
"tool"] =
"SegmentEdgeClassifierTool";
195 metadata[
"edge_attr_feature_names"] = {
196 "deltaPositionX_m",
"deltaPositionY_m",
"deltaPositionZ_m",
197 "distance_m",
"cos_opening_angle",
"same_chamber",
"same_sector"};
198 metadata[
"edge_index_layout"] =
"row_major_2_by_E";
199 metadata[
"edge_order"] =
"directed src_to_dst; row 0 then row 1";
204 out << metadata.dump() <<
'\n';
208 <<
" (DebugDumpMaxEvents="
212 return StatusCode::SUCCESS;
216 ATH_MSG_ERROR(
"runGraphInference is not supported by SegmentEdgeClassifierTool. Use SegmentEdgeInferenceAlg + ISegmentEdgeClassifierTool methods.");
217 return StatusCode::FAILURE;
226 std::vector<Amg::Vector3D> pos, dir;
227 std::vector<BucketSegmentFeatures> bucket;
230 std::map<SegmentGroupKey, int> segmentMultiplicity{};
233 ++segmentMultiplicity[segmentGroupKey(*seg)];
241 const int chamberIdx =
static_cast<int>(seg->
chamberIndex());
243 const int sec = seg->
sector();
244 const auto multIt = segmentMultiplicity.find(segmentGroupKey(*seg));
245 const int nSeg = (multIt != segmentMultiplicity.end()) ? multIt->second : 1;
248 pos.emplace_back(p.x() / Gaudi::Units::m,
249 p.y() / Gaudi::Units::m,
250 p.z() / Gaudi::Units::m);
251 dir.emplace_back(d.x(), d.y(), d.z());
254 graph.
nodeFeatures.push_back(nodeFeatureValue(featureId, pos.back(), dir.back(), bucket.back()));
260 if (pos.size() != graph.
nNodes || dir.size() != graph.
nNodes || bucket.size() != graph.
nNodes) {
262 <<
", pos=" << pos.size() <<
", dir=" << dir.size() <<
", bucket=" << bucket.size());
263 return StatusCode::FAILURE;
268 return StatusCode::SUCCESS;
271 auto normalizeSector = [&](
int s) {
285 std::unordered_map<int, std::vector<std::size_t>> nodesBySector;
286 nodesBySector.reserve(graph.
nNodes);
287 for (std::size_t i = 0; i < graph.
nNodes; ++i) {
288 nodesBySector[normalizeSector(bucket[i].sector)].push_back(i);
291 const std::size_t maxEdges = graph.
nNodes * (graph.
nNodes - 1);
295 for (std::size_t i = 0; i < graph.
nNodes; ++i) {
296 std::unordered_set<int> targetSectors;
299 targetSectors.insert(normalizeSector(bucket[i].sector + delta));
302 for (
const int sec : targetSectors) {
303 auto it = nodesBySector.find(sec);
304 if (it == nodesBySector.end())
continue;
305 for (
const std::size_t j : it->second) {
306 if (i == j)
continue;
308 const float cosang =
static_cast<float>(dir[i].dot(dir[j]));
311 graph.
edgeIndex.push_back(
static_cast<int64_t
>(i));
312 graph.
edgeIndex.push_back(
static_cast<int64_t
>(j));
315 const float dx =
static_cast<float>(delta.x());
316 const float dy =
static_cast<float>(delta.y());
317 const float dz =
static_cast<float>(delta.z());
318 const float dist =
static_cast<float>(delta.mag());
319 graph.
edgeFeatures.insert(graph.
edgeFeatures.end(), {dx,dy,dz,dist,cosang, float(bucket[i].chamberIndex==bucket[j].chamberIndex), float(bucket[i].sector==bucket[j].sector)});
325 <<
", kept nodes=" << graph.
nNodes
326 <<
", built edges=" << graph.
nEdges);
327 return StatusCode::SUCCESS;
332 std::vector<SegmentEdgeScore>& scores)
const {
334 if (!graph.
nNodes)
return StatusCode::SUCCESS;
337 return StatusCode::SUCCESS;
343 return StatusCode::FAILURE;
347 <<
"; expected " << (2 * graph.
nEdges));
348 return StatusCode::FAILURE;
353 return StatusCode::FAILURE;
357 raw.
graph = std::make_unique<InferenceGraph>();
362 for (std::size_t e = 0; e < graph.
nEdges; ++e) {
369 Ort::MemoryInfo memInfo = Ort::MemoryInfo::CreateCpu(OrtDeviceAllocator, OrtMemTypeCPU);
371 const std::vector<int64_t> nodeShape{
static_cast<int64_t
>(graph.
nNodes),
static_cast<int64_t
>(
kNodeFeatureCount)};
372 raw.
graph->dataTensor.emplace_back(
373 Ort::Value::CreateTensor<float>(memInfo,
379 const std::vector<int64_t> edgeIndexShape{2,
static_cast<int64_t
>(graph.
nEdges)};
380 raw.
graph->dataTensor.emplace_back(
381 Ort::Value::CreateTensor<int64_t>(memInfo,
384 edgeIndexShape.data(),
385 edgeIndexShape.size()));
390 const std::vector<int64_t> edgeAttrShape{
static_cast<int64_t
>(graph.
nEdges),
static_cast<int64_t
>(
kEdgeFeatureCount)};
391 raw.
graph->dataTensor.emplace_back(
392 Ort::Value::CreateTensor<float>(memInfo,
395 edgeAttrShape.data(),
396 edgeAttrShape.size()));
398 const std::vector<const char*> inputNames{
402 const std::vector<const char*> outputNames{
m_outputName.value().c_str()};
403 ATH_MSG_DEBUG(
"classifyEdges: ONNX inputs shapes x=[" << nodeShape[0] <<
"," << nodeShape[1]
404 <<
"], edge_index=[" << edgeIndexShape[0] <<
"," << edgeIndexShape[1]
405 <<
"], edge_attr=[" << edgeAttrShape[0] <<
"," << edgeAttrShape[1] <<
"]");
408 if (raw.
graph->dataTensor.size() <= inputNames.size()) {
409 ATH_MSG_ERROR(
"Missing ONNX output tensor for segment edge inference");
410 return StatusCode::FAILURE;
413 const Ort::Value& outTensor = raw.
graph->dataTensor[inputNames.size()];
414 const auto outInfo = outTensor.GetTensorTypeAndShapeInfo();
415 const std::vector<int64_t> outShape = outInfo.GetShape();
416 const size_t outSize = outInfo.GetElementCount();
417 if (!outShape.empty()) {
418 ATH_MSG_DEBUG(
"classifyEdges: ONNX output rank=" << outShape.size()
419 <<
", first dim=" << outShape.front()
420 <<
", elements=" << outSize);
422 ATH_MSG_DEBUG(
"classifyEdges: ONNX scalar output, elements=" << outSize);
424 if (outSize < graph.
nEdges) {
425 ATH_MSG_ERROR(
"ONNX logits tensor has " << outSize <<
" entries for " << graph.
nEdges <<
" edges");
426 return StatusCode::FAILURE;
429 const float* logits = outTensor.GetTensorData<
float>();
430 scores.reserve(graph.
nEdges);
431 for (std::size_t e=0; e<graph.
nEdges; ++e) {
432 const float l = logits[e];
433 scores.push_back({std::size_t(graph.
edgeIndex[2 * e]),
440 return StatusCode::SUCCESS;
444 const EventContext& ctx,
446 const std::vector<SegmentEdgeScore>& scores)
const {
453 return StatusCode::SUCCESS;
459 scores.size() != graph.
nEdges) {
460 ATH_MSG_ERROR(
"Cannot write segment-edge debug dump: inconsistent graph/output sizes"
461 <<
" nodes=" << graph.
nNodes
463 <<
" edges=" << graph.
nEdges
464 <<
" edgeIndex=" << graph.
edgeIndex.size()
466 <<
" scores=" << scores.size());
467 return StatusCode::FAILURE;
470 nlohmann::json
x = nlohmann::json::array();
471 x.get_ref<nlohmann::json::array_t&>().reserve(graph.
nodeFeatures.size());
473 x.push_back(std::isfinite(value) ? nlohmann::json(value)
474 : nlohmann::json(
nullptr));
477 nlohmann::json edgeIndex = nlohmann::json::array();
478 edgeIndex.get_ref<nlohmann::json::array_t&>().reserve(graph.
nEdges * 2);
480 for (std::size_t edge = 0; edge < graph.
nEdges; ++edge) {
481 edgeIndex.push_back(graph.
edgeIndex[2 * edge]);
483 for (std::size_t edge = 0; edge < graph.
nEdges; ++edge) {
484 edgeIndex.push_back(graph.
edgeIndex[2 * edge + 1]);
487 nlohmann::json edgeAttr = nlohmann::json::array();
488 edgeAttr.get_ref<nlohmann::json::array_t&>().reserve(graph.
edgeFeatures.size());
490 edgeAttr.push_back(std::isfinite(value) ? nlohmann::json(value)
491 : nlohmann::json(
nullptr));
494 nlohmann::json logits = nlohmann::json::array();
495 nlohmann::json probabilities = nlohmann::json::array();
496 nlohmann::json edgeSrc = nlohmann::json::array();
497 nlohmann::json edgeDst = nlohmann::json::array();
498 logits.get_ref<nlohmann::json::array_t&>().reserve(scores.size());
499 probabilities.get_ref<nlohmann::json::array_t&>().reserve(scores.size());
500 edgeSrc.get_ref<nlohmann::json::array_t&>().reserve(scores.size());
501 edgeDst.get_ref<nlohmann::json::array_t&>().reserve(scores.size());
503 edgeSrc.push_back(score.src);
504 edgeDst.push_back(score.dst);
505 logits.push_back(std::isfinite(score.logit) ? nlohmann::json(score.logit)
506 : nlohmann::json(
nullptr));
507 probabilities.push_back(std::isfinite(score.probability)
508 ? nlohmann::json(score.probability)
509 : nlohmann::json(
nullptr));
512 std::ofstream out{
m_debugDumpFile.value(), std::ios::out | std::ios::app};
514 ATH_MSG_ERROR(
"Could not append to segment-edge debug dump file: "
516 return StatusCode::FAILURE;
519 const unsigned int dumpIndex =
521 nlohmann::ordered_json event;
522 event[
"record_type"] =
"event";
523 event[
"format_version"] = 1;
524 event[
"dump_index"] = dumpIndex;
525 event[
"run_number"] = ctx.eventID().run_number();
526 event[
"lumi_block"] = ctx.eventID().lumi_block();
527 event[
"event_number"] = ctx.eventID().event_number();
528 event[
"slot"] = ctx.slot();
529 event[
"n_nodes"] = graph.
nNodes;
530 event[
"n_edges"] = graph.
nEdges;
532 event[
"edge_index_shape"] = {2, graph.
nEdges};
534 event[
"logits_shape"] = {graph.
nEdges};
535 event[
"x"] = std::move(
x);
536 event[
"edge_index"] = std::move(edgeIndex);
537 event[
"edge_attr"] = std::move(edgeAttr);
538 event[
"edge_src"] = std::move(edgeSrc);
539 event[
"edge_dst"] = std::move(edgeDst);
540 event[
"logits"] = std::move(logits);
541 event[
"probabilities"] = std::move(probabilities);
542 out <<
event.dump() <<
'\n';
547 return StatusCode::SUCCESS;
#define ATH_CHECK
Evaluate an expression and check for errors.
#define ATH_MSG_WARNING(x)
virtual void lock()=0
Interface to allow an object to lock itself when made const in SG.
Define macros for attributes used to control the static checker.
#define ATLAS_THREAD_SAFE
size_type size() const noexcept
Returns the number of elements in the collection.
: The muon space point bucket represents a collection of points that will bre processed together in t...
The SpacePointPerLayerSorter sort two given space points by their layer Identifier.
Amg::Vector3D direction() const
Returns the direction as Amg::Vector.
::Muon::MuonStationIndex::ChIndex chamberIndex() const
Returns the chamber index.
Amg::Vector3D position() const
Returns the position as Amg::Vector.
int etaIndex() const
Returns the eta index, which corresponds to stationEta in the offline identifiers (and the ).
Eigen::Matrix< double, 3, 1 > Vector3D
SegmentNodeFeatureId
Identifier for each node feature in segment-based GNNs.
const Segment * detailedSegment(const xAOD::MuonSegment &seg)
Helper function to navigate from the xAOD::MuonSegment to the MuonR4::Segment.
MuonSegmentContainer_v1 MuonSegmentContainer
Definition of the current "MuonSegment container version".
MuonSegment_v1 MuonSegment
Reference the current persistent version:
Segment features derived from or stored in bucket metadata.
int chamberIndex
Muon chamber index of the segment.
int sector
Sector number (typically 0–15).
int layers
Total number of active layers in the segment.
int nSegments
Count of segments in the same chamber/sector/eta group.
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.
std::vector< float > edgeFeatures
packed [E,7]: dpos(3), dist, cos, same_chamber, same_sector
std::vector< int64_t > edgeIndex
packed edge pairs [src0,dst0,src1,dst1,...]
std::vector< const xAOD::MuonSegment_v1 * > segments
std::vector< float > nodeFeatures
packed [N,10]: pos_m(3), dir_u(3), bucket(4)