5#ifdef ACTS_GNN_WITH_MODULEMAP
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"
36 ActsPlugins::ModuleMapCuda::Config gcCfg;
41 gcCfg.gpuBlocks = 512;
42 auto gc = std::make_shared<ActsPlugins::ModuleMapCuda>(
43 gcCfg,
m_logger->cloneWithSuffix(
"ModuleMap"));
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;
54 gnn = std::make_shared<ActsPlugins::OnnxEdgeClassifier>(
55 gnnCfg,
m_logger->cloneWithSuffix(
"GNN"));
58 ATH_MSG_FATAL(
"Not compiled with ONNX, cannot interpret *.onnx files");
59 return StatusCode::FAILURE;
61 }
else if (
m_gnnPath.value().find(
".pt") != std::string::npos) {
62#ifdef ACTS_GNN_TORCH_BACKEND
63 ActsPlugins::TorchEdgeClassifier::Config gnnCfg;
66 gnn = std::make_shared<ActsPlugins::TorchEdgeClassifier>(
67 gnnCfg,
m_logger->cloneWithSuffix(
"GNN"));
70 ATH_MSG_FATAL(
"Not compiled with Torch, cannot interpret *.pt files");
71 return StatusCode::FAILURE;
73 }
else if (
m_gnnPath.value().find(
".engine") != std::string::npos) {
74#ifdef ACTS_GNN_WITH_TENSORRT
75 ActsPlugins::TensorRTEdgeClassifier::Config gnnCfg;
79 gnn = std::make_shared<ActsPlugins::TensorRTEdgeClassifier>(
80 gnnCfg,
m_logger->cloneWithSuffix(
"GNN"));
82 ATH_MSG_FATAL(
"Not compiled with TensorRT, cannot interpret *.engine files");
83 return StatusCode::FAILURE;
87 return StatusCode::FAILURE;
91 std::shared_ptr<ActsPlugins::TrackBuildingBase>
tb;
92 ATH_MSG_INFO(
"Configure CC&JunctionRemoval as graph segmentation algorithm");
94 ActsPlugins::EdgeLayerConnector::Config tbCfg;
96 tbCfg.blockSize = 512;
98 tb = std::make_shared<ActsPlugins::EdgeLayerConnector>(
99 tbCfg,
m_logger->cloneWithSuffix(
"ELC"));
101 ActsPlugins::CudaTrackBuilding::Config tbCfg;
102 tbCfg.doJunctionRemoval =
true;
103 tb = std::make_shared<ActsPlugins::CudaTrackBuilding>(
104 tbCfg,
m_logger->cloneWithSuffix(
"CC&JR"));
109 gc, std::vector{std::move(gnn)},
tb,
m_logger->cloneWithSuffix(
"Pipeline"));
111 return StatusCode::SUCCESS;
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 {
119 const std::size_t nSP = spacepoints.size();
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();
132 std::vector<std::uint64_t> moduleIds(nSP);
133 std::vector<int>
ids(nSP);
135 for (std::size_t k = 0;
k < nSP; ++
k) {
136 const std::size_t origIdx = sortIdx[
k];
139 moduleIds[
k] = spacepoints[origIdx]->clusterList().first->detectorElement()->identify().get_compact();
140 ids[
k] =
static_cast<int>(
k);
148 std::unique_lock<std::mutex>
lock;
149 if (m_runMutex)
lock = std::unique_lock<std::mutex>(*m_runMutex);
151 if (edgeMap !=
nullptr) {
152 ScoredGraphHook hook;
153 auto result =
m_gnnPipeline->run(features, moduleIds, ids, ActsPlugins::Device::Cuda(0), 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();
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];
169 return m_gnnPipeline->run(features, moduleIds, ids, ActsPlugins::Device::Cuda(0));;
178 for (
const auto& candidate : candidates) {
184 std::vector<uint32_t>
track;
185 track.reserve(candidate.size());
186 for (
int sortedIdx : candidate) {
189 tracks.push_back(std::move(track));
192 ATH_MSG_DEBUG(
"Returning " << tracks.size() <<
" track candidates after filtering (>= "
195 return StatusCode::SUCCESS;
200 out <<
"|---------------------------------------------------------------------|\n" ;
201 out <<
"| ActsGnnModuleMapFinderTool |\n" ;
202 out <<
"|---------------------------------------------------------------------|\n" ;
#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)
::StatusCode StatusCode
StatusCode definition for legacy code.
float j(const xAOD::IParticle &, const xAOD::TrackMeasurementValidation &hit, const Eigen::Matrix3d &jab_inv)