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