7#ifdef ACTS_GNN_WITH_MODULEMAP
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"
28 const std::string& name,
29 const IInterface* parent)
30 : base_class(
type, name, parent) {}
34#ifdef ACTS_GNN_WITH_MODULEMAP
41 return StatusCode::FAILURE;
48 ActsPlugins::ModuleMapCuda::Config gcCfg;
49 gcCfg.rScale = 1000.f;
50 gcCfg.zScale = 1000.f;
51 gcCfg.phiScale = std::numbers::pi_v<float>;
53 gcCfg.gpuBlocks = 512;
54 std::shared_ptr<ActsPlugins::ModuleMapCuda> gc =
55 std::make_shared<ActsPlugins::ModuleMapCuda>(
56 gcCfg,
m_logger->cloneWithSuffix(
"ModuleMap"));
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;
64 gnn = std::make_shared<ActsPlugins::OnnxEdgeClassifier>(
65 gnnCfg,
m_logger->cloneWithSuffix(
"GNN"));
67 ATH_MSG_ERROR(
"GNN .onnx selected but build lacks ONNX backend");
68 return StatusCode::FAILURE;
70 }
else if (
m_gnnPath.value().find(
".pt") != std::string::npos) {
71#ifdef ACTS_GNN_TORCH_BACKEND
72 ActsPlugins::TorchEdgeClassifier::Config gnnCfg;
75 gnnCfg.useEdgeFeatures =
true;
76 gnn = std::make_shared<ActsPlugins::TorchEdgeClassifier>(
77 gnnCfg,
m_logger->cloneWithSuffix(
"GNN"));
79 ATH_MSG_ERROR(
"GNN .pt selected but build lacks libtorch backend");
80 return StatusCode::FAILURE;
82 }
else if (
m_gnnPath.value().find(
".engine") != std::string::npos) {
83#ifdef ACTS_GNN_WITH_TENSORRT
84 ActsPlugins::TensorRTEdgeClassifier::Config gnnCfg;
88 gnn = std::make_shared<ActsPlugins::TensorRTEdgeClassifier>(
89 gnnCfg,
m_logger->cloneWithSuffix(
"GNN"));
91 ATH_MSG_ERROR(
"GNN .engine selected but build lacks TensorRT backend");
92 return StatusCode::FAILURE;
96 return StatusCode::FAILURE;
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"));
105 m_gnnPipeline = std::make_unique<ActsPlugins::GnnPipeline>(
106 gc, std::vector{gnn},
tb,
m_logger->cloneWithSuffix(
"Pipeline"));
111 ACTS_INFO(
"Use phi overlap spacepoints: " << std::boolalpha
113 return StatusCode::SUCCESS;
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;
126 std::optional<Athena::Chrono>
timer;
129 m_gpuInstanceCount->acquire();
131 m_gnnPipeline->run(features, moduleIds, ids,
133 m_gpuInstanceCount->release();
135 ACTS_DEBUG(
"Have " <<
candidates.size() <<
" candidates after GNN");
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();
146 bool noPixelHits = !std::ranges::any_of(c, [&](
auto spi) {
147 return allSPPtrs.at(spi)->measurements().size() == 1;
149 return tooFewMeasurements || noPixelHits;
152 for (
const auto& candidate : candidates) {
153 if (candidateSelector(candidate)) {
156 std::vector<const xAOD::SpacePoint*> seedSPs;
157 seedSPs.reserve(candidate.size());
158 for (
int spi : candidate) {
159 seedSPs.push_back(allSPPtrs.at(spi));
165 return Acts::fastHypot(
a->x(),
a->y(),
a->z()) <
166 Acts::fastHypot(
b->x(),
b->y(),
b->z());
168 constexpr float quality = 0.f;
169 constexpr float vertexZ = 0.f;
170 seeds.
push_back(ActsTrk::SpacePointRange(seedSPs.data(), seedSPs.size()),
175 <<
" measurements: " << seeds.
size());
176 return StatusCode::SUCCESS;
183 ACTS_ERROR(
"Cannot initialize ActsGnn without the Acts GNN plugin");
184 return StatusCode::FAILURE;
188 const std::vector<const xAOD::SpacePointContainer*>& ,
190 return StatusCode::FAILURE;
#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)
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