ATLAS Offline Software
Loading...
Searching...
No Matches
ActsGnnModuleMapFinderTool.cxx
Go to the documentation of this file.
1/*
2 Copyright (C) 2002-2026 CERN for the benefit of the ATLAS collaboration
3*/
4
5#ifdef ACTS_GNN_WITH_MODULEMAP
6
8
10
11#include "ActsPlugins/Gnn/CudaTrackBuilding.hpp"
12#include "ActsPlugins/Gnn/GnnPipeline.hpp"
13#include "ActsPlugins/Gnn/ModuleMapCuda.hpp"
14#include "ActsPlugins/Gnn/OnnxEdgeClassifier.hpp"
15#include "ActsPlugins/Gnn/TensorRTEdgeClassifier.hpp"
16#include "ActsPlugins/Gnn/TorchEdgeClassifier.hpp"
17#include "ActsPlugins/Gnn/EdgeLayerConnector.hpp"
18
22#include "ActsGnnHookTool.h"
23
24#include <algorithm>
25#include <numeric>
26
27
30
31 m_logger = makeActsAthenaLogger(this, "ActsGnnModuleMapFinderTool");
32
33 // Build ACTS GNN pipeline components
34
35 // 1. Graph constructor (ModuleMapCuda)
36 ActsPlugins::ModuleMapCuda::Config gcCfg;
37 gcCfg.rScale = kScaleR;
38 gcCfg.zScale = kScaleZ;
39 gcCfg.phiScale = kScalePhi;
40 gcCfg.moduleMapPath = m_moduleMapPath.value();
41 gcCfg.gpuBlocks = 512;
42 auto gc = std::make_shared<ActsPlugins::ModuleMapCuda>(
43 gcCfg, m_logger->cloneWithSuffix("ModuleMap"));
44
45 // 2. Edge classifier (ONNX / Torch / TensorRT)
46 // ONNX and Torch are not thread-safe: a mutex is emplaced to serialise run().
47 // TensorRT manages concurrency internally via execution contexts: no guard needed.
48 std::shared_ptr<ActsPlugins::EdgeClassificationBase> gnn;
49 if (m_gnnPath.value().find(".onnx") != std::string::npos) {
50#ifdef ACTS_GNN_ONNX_BACKEND
51 ActsPlugins::OnnxEdgeClassifier::Config gnnCfg;
52 gnnCfg.modelPath = m_gnnPath.value();
53 gnnCfg.cut = m_edgeCut.value();
54 gnn = std::make_shared<ActsPlugins::OnnxEdgeClassifier>(
55 gnnCfg, m_logger->cloneWithSuffix("GNN"));
56 m_runMutex.emplace();
57#else
58 ATH_MSG_FATAL("Not compiled with ONNX, cannot interpret *.onnx files");
59 return StatusCode::FAILURE;
60#endif
61 } else if (m_gnnPath.value().find(".pt") != std::string::npos) {
62#ifdef ACTS_GNN_TORCH_BACKEND
63 ActsPlugins::TorchEdgeClassifier::Config gnnCfg;
64 gnnCfg.modelPath = m_gnnPath.value();
65 gnnCfg.cut = m_edgeCut.value();
66 gnn = std::make_shared<ActsPlugins::TorchEdgeClassifier>(
67 gnnCfg, m_logger->cloneWithSuffix("GNN"));
68 m_runMutex.emplace();
69#else
70 ATH_MSG_FATAL("Not compiled with Torch, cannot interpret *.pt files");
71 return StatusCode::FAILURE;
72#endif
73 } else if (m_gnnPath.value().find(".engine") != std::string::npos) {
74#ifdef ACTS_GNN_WITH_TENSORRT
75 ActsPlugins::TensorRTEdgeClassifier::Config gnnCfg;
76 gnnCfg.modelPath = m_gnnPath.value();
77 gnnCfg.cut = m_edgeCut.value();
78 gnnCfg.numExecutionContexts = m_numTrtContexts.value();
79 gnn = std::make_shared<ActsPlugins::TensorRTEdgeClassifier>(
80 gnnCfg, m_logger->cloneWithSuffix("GNN"));
81#else
82 ATH_MSG_FATAL("Not compiled with TensorRT, cannot interpret *.engine files");
83 return StatusCode::FAILURE;
84#endif
85 } else {
86 ATH_MSG_FATAL("Unknown extension for GNN model: " << m_gnnPath.value());
87 return StatusCode::FAILURE;
88 }
89
90 // 3. Track builder
91 std::shared_ptr<ActsPlugins::TrackBuildingBase> tb;
92 ATH_MSG_INFO("Configure CC&JunctionRemoval as graph segmentation algorithm");
94 ActsPlugins::EdgeLayerConnector::Config tbCfg;
95 tbCfg.maxHitsPerTrack = m_elcMaxHitsPerTrack;
96 tbCfg.blockSize = 512;
97 tbCfg.weightsCut = m_edgeCut;
98 tb = std::make_shared<ActsPlugins::EdgeLayerConnector>(
99 tbCfg, m_logger->cloneWithSuffix("ELC"));
100 } else {
101 ActsPlugins::CudaTrackBuilding::Config tbCfg;
102 tbCfg.doJunctionRemoval = true;
103 tb = std::make_shared<ActsPlugins::CudaTrackBuilding>(
104 tbCfg, m_logger->cloneWithSuffix("CC&JR"));
105 }
106
107 // 4. Assemble pipeline
108 m_gnnPipeline = std::make_unique<ActsPlugins::GnnPipeline>(
109 gc, std::vector{std::move(gnn)}, tb, m_logger->cloneWithSuffix("Pipeline"));
110
111 return StatusCode::SUCCESS;
112}
113
115 const std::vector<const Trk::SpacePoint*>& spacepoints,
116 std::vector<std::vector<uint32_t>>& tracks,
117 std::unordered_map<int, std::unordered_map<int, float>>* edgeMap) const {
118
119 const std::size_t nSP = spacepoints.size();
120
121 ATH_MSG_DEBUG("Processing " << nSP << " spacepoints with " << NUM_FEATURES << " features");
122
123 // Sort spacepoint indices by module ID (required by module map graph construction)
124 std::vector<std::size_t> sortIdx(nSP);
125 std::iota(sortIdx.begin(), sortIdx.end(), 0);
126 std::ranges::sort(sortIdx, std::less{}, [&](std::size_t i) {
127 return spacepoints[i]->clusterList().first->detectorElement()->identify().get_compact();
128 });
129
130 // Build features, module IDs, and IDs directly in sorted order
131 std::vector<float> features(NUM_FEATURES * nSP);
132 std::vector<std::uint64_t> moduleIds(nSP);
133 std::vector<int> ids(nSP);
134
135 for (std::size_t k = 0; k < nSP; ++k) {
136 const std::size_t origIdx = sortIdx[k];
137 auto featureMap = m_spacepointFeatureTool->getFeatures(spacepoints[origIdx]);
138 // Use detector element identifier, not cluster identifier, to get the module ID
139 moduleIds[k] = spacepoints[origIdx]->clusterList().first->detectorElement()->identify().get_compact();
140 ids[k] = static_cast<int>(k);
141 for (std::size_t j = 0; j < NUM_FEATURES; ++j) {
142 features[k * NUM_FEATURES + j] = featureMap[FEATURE_NAMES[j]] / FEATURE_SCALES[j];
143 }
144 }
145
146 // Run GNN pipeline (mutex present for ONNX/Torch, absent for TRT)
147 auto candidates = [&] {
148 std::unique_lock<std::mutex> lock;
149 if (m_runMutex) lock = std::unique_lock<std::mutex>(*m_runMutex);
150
151 if (edgeMap != nullptr) {
152 ScoredGraphHook hook;
153 auto result = m_gnnPipeline->run(features, moduleIds, ids, ActsPlugins::Device::Cuda(0), hook);
154
155 // Retrieve edgeScores and edgeIndex from hook
156 const std::vector<float>& edgeScores = hook.getEdgeScores();
157 const std::vector<std::int64_t>& edgeIndex = hook.getEdgeIndex();
158 const std::size_t nEdges = edgeScores.size();
159
160 // Create a map to acces edge score (sorted indices back to original spacepoint indices)
161 for (std::size_t i = 0; i < nEdges; ++i) {
162 std::int64_t src = edgeIndex[i];
163 std::int64_t dst = edgeIndex[nEdges + i];
164 (*edgeMap)[sortIdx[src]][sortIdx[dst]] = edgeScores[i];
165 }
166 return result;
167 }
168
169 return m_gnnPipeline->run(features, moduleIds, ids, ActsPlugins::Device::Cuda(0));;
170 }();
171
172 ATH_MSG_DEBUG("GNN pipeline returned " << candidates.size() << " candidates");
173
174 // Filter by minimum measurements and convert indices back to original ordering
175 tracks.clear();
176 tracks.reserve(candidates.size());
177
178 for (const auto& candidate : candidates) {
179 if (candidate.size() < m_minCandidateMeasurements.value()) {
180 continue;
181 }
182
183 // Map sorted indices back to original spacepoint indices
184 std::vector<uint32_t> track;
185 track.reserve(candidate.size());
186 for (int sortedIdx : candidate) {
187 track.push_back(static_cast<uint32_t>(sortIdx[sortedIdx]));
188 }
189 tracks.push_back(std::move(track));
190 }
191
192 ATH_MSG_DEBUG("Returning " << tracks.size() << " track candidates after filtering (>= "
193 << m_minCandidateMeasurements.value() << " measurements)");
194
195 return StatusCode::SUCCESS;
196}
197
198MsgStream& InDet::ActsGnnModuleMapFinderTool::dump(MsgStream& out) const {
199 out << "\n";
200 out << "|---------------------------------------------------------------------|\n" ;
201 out << "| ActsGnnModuleMapFinderTool |\n" ;
202 out << "|---------------------------------------------------------------------|\n" ;
203 return out;
204}
205
206std::ostream& InDet::ActsGnnModuleMapFinderTool::dump(std::ostream& out) const {
207 return out;
208}
209
210#endif // ACTS_GNN_WITH_MODULEMAP
#define ATH_CHECK
Evaluate an expression and check for errors.
#define ATH_MSG_DEBUG(x,...)
#define ATH_MSG_INFO(x,...)
#define ATH_MSG_FATAL(x,...)
virtual void lock()=0
Interface to allow an object to lock itself when made const in SG.
std::unique_ptr< const Acts::Logger > makeActsAthenaLogger(IMessageSvc *svc, const std::string &name, int level, std::optional< std::string > parent_name)
Definition Logger.cxx:64
virtual StatusCode initialize() override
virtual MsgStream & dump(MsgStream &out) const override
static constexpr std::array< float, NUM_FEATURES > FEATURE_SCALES
static constexpr std::array< const char *, NUM_FEATURES > FEATURE_NAMES
std::unique_ptr< const Acts::Logger > m_logger
ToolHandle< ISpacepointFeatureTool > m_spacepointFeatureTool
static constexpr std::size_t NUM_FEATURES
virtual StatusCode getTracks(const std::vector< const Trk::SpacePoint * > &spacepoints, std::vector< std::vector< uint32_t > > &tracks, std::unordered_map< int, std::unordered_map< int, float > > *edgeMap=nullptr) const override
std::unique_ptr< ActsPlugins::GnnPipeline > m_gnnPipeline
::StatusCode StatusCode
StatusCode definition for legacy code.
float j(const xAOD::IParticle &, const xAOD::TrackMeasurementValidation &hit, const Eigen::Matrix3d &jab_inv)
setEventNumber uint32_t