7#ifdef ACTS_GNN_WITH_MODULEMAP
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"
27 const std::string& name,
28 const IInterface* parent)
29 : base_class(
type, name, parent) {}
33#ifdef ACTS_GNN_WITH_MODULEMAP
42 ActsPlugins::ModuleMapCuda::Config gcCfg;
43 gcCfg.rScale = 1000.f;
44 gcCfg.zScale = 1000.f;
45 gcCfg.phiScale = std::numbers::pi_v<float>;
47 gcCfg.gpuBlocks = 512;
48 std::shared_ptr<ActsPlugins::ModuleMapCuda> gc =
49 std::make_shared<ActsPlugins::ModuleMapCuda>(
50 gcCfg,
m_logger->cloneWithSuffix(
"ModuleMap"));
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;
58 gnn = std::make_shared<ActsPlugins::OnnxEdgeClassifier>(
59 gnnCfg,
m_logger->cloneWithSuffix(
"GNN"));
61 ATH_MSG_ERROR(
"GNN .onnx selected but build lacks ONNX backend");
62 return StatusCode::FAILURE;
64 }
else if (
m_gnnPath.value().find(
".pt") != std::string::npos) {
65#ifdef ACTS_GNN_TORCH_BACKEND
66 ActsPlugins::TorchEdgeClassifier::Config gnnCfg;
69 gnnCfg.useEdgeFeatures =
true;
70 gnn = std::make_shared<ActsPlugins::TorchEdgeClassifier>(
71 gnnCfg,
m_logger->cloneWithSuffix(
"GNN"));
73 ATH_MSG_ERROR(
"GNN .pt selected but build lacks libtorch backend");
74 return StatusCode::FAILURE;
76 }
else if (
m_gnnPath.value().find(
".engine") != std::string::npos) {
77#ifdef ACTS_GNN_WITH_TENSORRT
78 ActsPlugins::TensorRTEdgeClassifier::Config gnnCfg;
82 gnn = std::make_shared<ActsPlugins::TensorRTEdgeClassifier>(
83 gnnCfg,
m_logger->cloneWithSuffix(
"GNN"));
85 ATH_MSG_ERROR(
"GNN .engine selected but build lacks TensorRT backend");
86 return StatusCode::FAILURE;
90 return StatusCode::FAILURE;
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"));
100 gc, std::vector{gnn},
tb,
m_logger->cloneWithSuffix(
"Pipeline"));
105 ACTS_INFO(
"Use phi overlap spacepoints: " << std::boolalpha
107 return StatusCode::SUCCESS;
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;
120 std::optional<Athena::Chrono>
timer;
123 m_gpuInstanceCount->acquire();
127 m_gpuInstanceCount->release();
129 ACTS_DEBUG(
"Have " <<
candidates.size() <<
" candidates after GNN");
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();
140 bool noPixelHits = !std::ranges::any_of(c, [&](
auto spi) {
141 return allSPPtrs.at(spi)->measurements().size() == 1;
143 return tooFewMeasurements || noPixelHits;
146 for (
const auto& candidate : candidates) {
147 if (candidateSelector(candidate)) {
150 std::vector<const xAOD::SpacePoint*> seedSPs;
151 seedSPs.reserve(candidate.size());
152 for (
int spi : candidate) {
153 seedSPs.push_back(allSPPtrs.at(spi));
155 constexpr float quality = 0.f;
156 constexpr float vertexZ = 0.f;
157 seeds.
push_back(ActsTrk::SpacePointRange(seedSPs.data(), seedSPs.size()),
162 <<
" measurements: " << seeds.
size());
163 return StatusCode::SUCCESS;
170 ACTS_ERROR(
"Cannot initialize ActsGnn without the Acts GNN plugin");
171 return StatusCode::FAILURE;
175 const std::vector<const xAOD::SpacePointContainer*>& ,
177 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