ATLAS Offline Software
Loading...
Searching...
No Matches
GnnPipelineTool.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#include "GnnPipelineTool.h"
6
7#ifdef ACTS_GNN_WITH_MODULEMAP
8
9#include <algorithm>
10#include <numbers>
11#include <numeric>
12
13#include "Acts/Utilities/MathHelpers.hpp"
14#include "ActsPlugins/Gnn/CudaTrackBuilding.hpp"
15#include "ActsPlugins/Gnn/GnnPipeline.hpp"
16#include "ActsPlugins/Gnn/ModuleMapCuda.hpp"
17#include "ActsPlugins/Gnn/OnnxEdgeClassifier.hpp"
18#include "ActsPlugins/Gnn/TensorRTEdgeClassifier.hpp"
19#include "ActsPlugins/Gnn/TorchEdgeClassifier.hpp"
20#include "AthenaKernel/Chrono.h"
21#include "detail/GnnFeatures.h"
22
23#endif
24
25namespace ActsTrk {
26
28 const std::string& name,
29 const IInterface* parent)
30 : base_class(type, name, parent) {}
31
33
34#ifdef ACTS_GNN_WITH_MODULEMAP
35
36StatusCode GnnPipelineTool::initialize() {
37 m_logger = makeActsAthenaLogger(this, "ActsGnn");
38
39 if (m_numFeatures.value() != 4 && m_numFeatures.value() != 12) {
40 ATH_MSG_ERROR("numFeatures must be 4 or 12, got " << m_numFeatures.value());
41 return StatusCode::FAILURE;
42 }
43
44 ATH_CHECK(m_chronoSvc.retrieve());
45 ATH_CHECK(detStore()->retrieve(m_pixelIdHelper, "PixelID"));
46 ATH_CHECK(detStore()->retrieve(m_stripIdHelper, "SCT_ID"));
47
48 ActsPlugins::ModuleMapCuda::Config gcCfg;
49 gcCfg.rScale = 1000.f;
50 gcCfg.zScale = 1000.f;
51 gcCfg.phiScale = std::numbers::pi_v<float>;
52 gcCfg.moduleMapPath = m_moduleMapPath.value();
53 gcCfg.gpuBlocks = 512;
54 std::shared_ptr<ActsPlugins::ModuleMapCuda> gc =
55 std::make_shared<ActsPlugins::ModuleMapCuda>(
56 gcCfg, m_logger->cloneWithSuffix("ModuleMap"));
57
58 std::shared_ptr<ActsPlugins::EdgeClassificationBase> gnn;
59 if (m_gnnPath.value().find(".onnx") != std::string::npos) {
60#ifdef ACTS_GNN_ONNX_BACKEND
61 ActsPlugins::OnnxEdgeClassifier::Config gnnCfg;
62 gnnCfg.modelPath = m_gnnPath.value();
63 gnnCfg.cut = m_edgeCut.value();
64 gnn = std::make_shared<ActsPlugins::OnnxEdgeClassifier>(
65 gnnCfg, m_logger->cloneWithSuffix("GNN"));
66#else
67 ATH_MSG_ERROR("GNN .onnx selected but build lacks ONNX backend");
68 return StatusCode::FAILURE;
69#endif
70 } else if (m_gnnPath.value().find(".pt") != std::string::npos) {
71#ifdef ACTS_GNN_TORCH_BACKEND
72 ActsPlugins::TorchEdgeClassifier::Config gnnCfg;
73 gnnCfg.modelPath = m_gnnPath.value();
74 gnnCfg.cut = m_edgeCut.value();
75 gnnCfg.useEdgeFeatures = true;
76 gnn = std::make_shared<ActsPlugins::TorchEdgeClassifier>(
77 gnnCfg, m_logger->cloneWithSuffix("GNN"));
78#else
79 ATH_MSG_ERROR("GNN .pt selected but build lacks libtorch backend");
80 return StatusCode::FAILURE;
81#endif
82 } else if (m_gnnPath.value().find(".engine") != std::string::npos) {
83#ifdef ACTS_GNN_WITH_TENSORRT
84 ActsPlugins::TensorRTEdgeClassifier::Config gnnCfg;
85 gnnCfg.cut = m_edgeCut.value();
86 gnnCfg.modelPath = m_gnnPath.value();
87 gnnCfg.numExecutionContexts = m_numTrtContexts.value();
88 gnn = std::make_shared<ActsPlugins::TensorRTEdgeClassifier>(
89 gnnCfg, m_logger->cloneWithSuffix("GNN"));
90#else
91 ATH_MSG_ERROR("GNN .engine selected but build lacks TensorRT backend");
92 return StatusCode::FAILURE;
93#endif
94 } else {
95 ATH_MSG_ERROR("Unknown GNN model extension: " << m_gnnPath.value());
96 return StatusCode::FAILURE;
97 }
98
99 ActsPlugins::CudaTrackBuilding::Config tbCfg;
100 std::shared_ptr<ActsPlugins::TrackBuildingBase> tb;
101 tbCfg.doJunctionRemoval = true;
102 tb = std::make_shared<ActsPlugins::CudaTrackBuilding>(
103 tbCfg, m_logger->cloneWithSuffix("GraphSeg"));
104
105 m_gnnPipeline = std::make_unique<ActsPlugins::GnnPipeline>(
106 gc, std::vector{gnn}, tb, m_logger->cloneWithSuffix("Pipeline"));
107
108 // Limit the total number of instances on the GPU to avoid out of memory
109 m_gpuInstanceCount.emplace(m_maxGpuInstances.value());
110
111 ACTS_INFO("Use phi overlap spacepoints: " << std::boolalpha
112 << m_usePhiOverlapSps.value());
113 return StatusCode::SUCCESS;
114}
115
117 const std::vector<const xAOD::SpacePointContainer*>& spacePointCollections,
118 ActsTrk::SeedContainer& seeds) const {
119 std::vector<float> features;
120 std::vector<std::uint64_t> moduleIds;
121 std::vector<int> ids;
122 std::vector<const xAOD::SpacePoint*> allSPPtrs;
123 ATH_CHECK(buildFeatures(spacePointCollections, features, moduleIds, ids,
124 allSPPtrs, m_numFeatures.value()));
125
126 std::optional<Athena::Chrono> timer;
127 timer.emplace("GNN inference", m_chronoSvc.get());
128
129 m_gpuInstanceCount->acquire();
130 auto candidates =
131 m_gnnPipeline->run(features, moduleIds, ids,
132 ActsPlugins::Device::Cuda(m_cudaDeviceIndex.value()));
133 m_gpuInstanceCount->release();
134
135 ACTS_DEBUG("Have " << candidates.size() << " candidates after GNN");
136
137 // Rough estimate of the number of spacepoints per seed
138 seeds.reserve(candidates.size(), 1 + m_minCandidateMeasurements.value() * 2);
139
140 // Remove candidates that have too few measurements or no pixel hit
141 auto candidateSelector = [&](const std::vector<int>& c) {
142 bool tooFewMeasurements =
143 std::accumulate(c.begin(), c.end(), 0ul, [&](auto sum, auto spi) {
144 return sum + allSPPtrs.at(spi)->measurements().size();
145 }) < m_minCandidateMeasurements.value();
146 bool noPixelHits = !std::ranges::any_of(c, [&](auto spi) {
147 return allSPPtrs.at(spi)->measurements().size() == 1;
148 });
149 return tooFewMeasurements || noPixelHits;
150 };
151
152 for (const auto& candidate : candidates) {
153 if (candidateSelector(candidate)) {
154 continue;
155 }
156 std::vector<const xAOD::SpacePoint*> seedSPs;
157 seedSPs.reserve(candidate.size());
158 for (int spi : candidate) {
159 seedSPs.push_back(allSPPtrs.at(spi));
160 }
161 // Order the space points from the innermost to the outermost one, as
162 // expected by the downstream parameter estimation and track finding
163 std::ranges::sort(seedSPs, [](const xAOD::SpacePoint* a,
164 const xAOD::SpacePoint* b) {
165 return Acts::fastHypot(a->x(), a->y(), a->z()) <
166 Acts::fastHypot(b->x(), b->y(), b->z());
167 });
168 constexpr float quality = 0.f; // quality is not computed in the GNN pipeline
169 constexpr float vertexZ = 0.f; // vertexZ is not computed in the GNN pipeline
170 seeds.push_back(ActsTrk::SpacePointRange(seedSPs.data(), seedSPs.size()),
171 quality, vertexZ);
172 }
173
174 ACTS_DEBUG("Candidates left with >= " << m_minCandidateMeasurements.value()
175 << " measurements: " << seeds.size());
176 return StatusCode::SUCCESS;
177}
178
179#else // ACTS_GNN_WITH_MODULEMAP
180
182 m_logger = makeActsAthenaLogger(this, "ActsGnn");
183 ACTS_ERROR("Cannot initialize ActsGnn without the Acts GNN plugin");
184 return StatusCode::FAILURE;
185}
186
188 const std::vector<const xAOD::SpacePointContainer*>& /*spacePointCollections*/,
189 ActsTrk::SeedContainer& /*seeds*/) const {
190 return StatusCode::FAILURE;
191}
192
193#endif // ACTS_GNN_WITH_MODULEMAP
194
195} // namespace ActsTrk
#define ATH_CHECK
Evaluate an expression and check for errors.
#define ATH_MSG_ERROR(x,...)
Exception-safe IChronoSvc caller.
static Double_t a
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
Gaudi::Property< std::string > m_gnnPath
StatusCode buildSeed(const std::vector< const xAOD::SpacePointContainer * > &spacePointCollections, ActsTrk::SeedContainer &seeds) const override
Gaudi::Property< unsigned int > m_numFeatures
const SCT_ID * m_stripIdHelper
Gaudi::Property< bool > m_usePhiOverlapSps
StatusCode initialize() override
Gaudi::Property< std::string > m_moduleMapPath
Gaudi::Property< unsigned int > m_maxGpuInstances
Gaudi::Property< unsigned int > m_minCandidateMeasurements
ServiceHandle< IChronoStatSvc > m_chronoSvc
std::unique_ptr< const Acts::Logger > m_logger
GnnPipelineTool(const std::string &type, const std::string &name, const IInterface *parent)
Gaudi::Property< unsigned int > m_numTrtContexts
Gaudi::Property< int > m_cudaDeviceIndex
Gaudi::Property< double > m_edgeCut
const PixelID * m_pixelIdHelper
StatusCode buildFeatures(const std::vector< const xAOD::SpacePointContainer * > &spacePointCollections, std::vector< float > &features, std::vector< std::uint64_t > &moduleIds, std::vector< int > &ids, std::vector< const xAOD::SpacePoint * > &allSPPtrs, std::size_t nFeatures=12) const
Definition GnnFeatures.h:48
The AlignStoreProviderAlg loads the rigid alignment corrections and pipes them through the readout ge...
::StatusCode StatusCode
StatusCode definition for legacy code.
timer(name, disabled=False)
Seed push_back(SpacePointRange spacePoints, float quality, float vertexZ)
void reserve(std::size_t size, float averageSpacePoints=3) noexcept
std::size_t size() const noexcept