Back to home page

EIC code displayed by LXR

 
 

    


File indexing completed on 2026-09-16 08:21:51

0001 // This file is part of the ACTS project.
0002 //
0003 // Copyright (C) 2016 CERN for the benefit of the ACTS project
0004 //
0005 // This Source Code Form is subject to the terms of the Mozilla Public
0006 // License, v. 2.0. If a copy of the MPL was not distributed with this
0007 // file, You can obtain one at https://mozilla.org/MPL/2.0/.
0008 
0009 #include "ActsPlugins/Gnn/DWalkTrackBuilding.hpp"
0010 #include "ActsPlugins/Gnn/detail/ConnectedComponents.cuh"
0011 #include "ActsPlugins/Gnn/detail/CudaUtils.hpp"
0012 
0013 #include <algorithm>
0014 #include <cassert>
0015 #include <cstdint>
0016 #include <functional>
0017 #include <limits>
0018 #include <stdexcept>
0019 #include <unordered_map>
0020 #include <utility>
0021 #include <vector>
0022 
0023 #include <cuda_runtime_api.h>
0024 #include <thrust/copy.h>
0025 #include <thrust/count.h>
0026 #include <thrust/execution_policy.h>
0027 #include <thrust/iterator/counting_iterator.h>
0028 #include <thrust/iterator/zip_iterator.h>
0029 #include <thrust/scan.h>
0030 #include <thrust/sort.h>
0031 #include <thrust/tuple.h>
0032 
0033 namespace {
0034 
0035 constexpr int kBlockSize = 256;
0036 constexpr std::size_t kSmallResidualEdgeThreshold = 512;
0037 
0038 struct IsNonZero {
0039   __host__ __device__ bool operator()(unsigned char value) const {
0040     return value != 0;
0041   }
0042 };
0043 
0044 struct Edge {
0045   int src = 0;
0046   int dst = 0;
0047 };
0048 
0049 template <typename T>
0050 void cudaFreeAsyncNoThrow(T *ptr, cudaStream_t stream) noexcept {
0051   if (ptr != nullptr) {
0052     static_cast<void>(cudaFreeAsync(ptr, stream));
0053   }
0054 }
0055 
0056 struct DeviceOrientedEdges {
0057   int *src = nullptr;
0058   int *dst = nullptr;
0059   float *score = nullptr;
0060   unsigned char *valid = nullptr;
0061   unsigned char *activeNodes = nullptr;
0062   std::size_t numEdges = 0;
0063   std::size_t numNodes = 0;
0064 
0065   DeviceOrientedEdges() = default;
0066   ~DeviceOrientedEdges() { reset(); }
0067 
0068   DeviceOrientedEdges(const DeviceOrientedEdges &) = delete;
0069   DeviceOrientedEdges &operator=(const DeviceOrientedEdges &) = delete;
0070 
0071   DeviceOrientedEdges(DeviceOrientedEdges &&other) noexcept {
0072     *this = std::move(other);
0073   }
0074 
0075   DeviceOrientedEdges &operator=(DeviceOrientedEdges &&other) noexcept {
0076     if (this != &other) {
0077       reset();
0078       src = std::exchange(other.src, nullptr);
0079       dst = std::exchange(other.dst, nullptr);
0080       score = std::exchange(other.score, nullptr);
0081       valid = std::exchange(other.valid, nullptr);
0082       activeNodes = std::exchange(other.activeNodes, nullptr);
0083       numEdges = std::exchange(other.numEdges, 0);
0084       numNodes = std::exchange(other.numNodes, 0);
0085       stream = std::exchange(other.stream, nullptr);
0086     }
0087     return *this;
0088   }
0089 
0090   void setStream(cudaStream_t owningStream) { stream = owningStream; }
0091 
0092   void reset() noexcept {
0093     cudaFreeAsyncNoThrow(src, stream);
0094     cudaFreeAsyncNoThrow(dst, stream);
0095     cudaFreeAsyncNoThrow(score, stream);
0096     cudaFreeAsyncNoThrow(valid, stream);
0097     cudaFreeAsyncNoThrow(activeNodes, stream);
0098     src = nullptr;
0099     dst = nullptr;
0100     score = nullptr;
0101     valid = nullptr;
0102     activeNodes = nullptr;
0103     numEdges = 0;
0104     numNodes = 0;
0105   }
0106 
0107  private:
0108   cudaStream_t stream = nullptr;
0109 };
0110 
0111 struct DeviceCompactEdges {
0112   int *src = nullptr;
0113   int *dst = nullptr;
0114   float *score = nullptr;
0115   std::size_t numEdges = 0;
0116   std::size_t numNodes = 0;
0117 
0118   DeviceCompactEdges() = default;
0119   ~DeviceCompactEdges() { reset(); }
0120 
0121   DeviceCompactEdges(const DeviceCompactEdges &) = delete;
0122   DeviceCompactEdges &operator=(const DeviceCompactEdges &) = delete;
0123 
0124   DeviceCompactEdges(DeviceCompactEdges &&other) noexcept {
0125     *this = std::move(other);
0126   }
0127 
0128   DeviceCompactEdges &operator=(DeviceCompactEdges &&other) noexcept {
0129     if (this != &other) {
0130       reset();
0131       src = std::exchange(other.src, nullptr);
0132       dst = std::exchange(other.dst, nullptr);
0133       score = std::exchange(other.score, nullptr);
0134       numEdges = std::exchange(other.numEdges, 0);
0135       numNodes = std::exchange(other.numNodes, 0);
0136       stream = std::exchange(other.stream, nullptr);
0137     }
0138     return *this;
0139   }
0140 
0141   void setStream(cudaStream_t owningStream) { stream = owningStream; }
0142 
0143   void reset() noexcept {
0144     cudaFreeAsyncNoThrow(src, stream);
0145     cudaFreeAsyncNoThrow(dst, stream);
0146     cudaFreeAsyncNoThrow(score, stream);
0147     src = nullptr;
0148     dst = nullptr;
0149     score = nullptr;
0150     numEdges = 0;
0151     numNodes = 0;
0152   }
0153 
0154  private:
0155   cudaStream_t stream = nullptr;
0156 };
0157 
0158 struct DeviceCsrGraph {
0159   int *rowPtr = nullptr;
0160   int *colIdx = nullptr;
0161   float *edgeWeight = nullptr;
0162   int *incomingRowPtr = nullptr;
0163   int *incomingColIdx = nullptr;
0164   std::size_t numNodes = 0;
0165   std::size_t numEdges = 0;
0166 
0167   DeviceCsrGraph() = default;
0168   ~DeviceCsrGraph() { reset(); }
0169 
0170   DeviceCsrGraph(const DeviceCsrGraph &) = delete;
0171   DeviceCsrGraph &operator=(const DeviceCsrGraph &) = delete;
0172 
0173   DeviceCsrGraph(DeviceCsrGraph &&other) noexcept { *this = std::move(other); }
0174 
0175   DeviceCsrGraph &operator=(DeviceCsrGraph &&other) noexcept {
0176     if (this != &other) {
0177       reset();
0178       rowPtr = std::exchange(other.rowPtr, nullptr);
0179       colIdx = std::exchange(other.colIdx, nullptr);
0180       edgeWeight = std::exchange(other.edgeWeight, nullptr);
0181       incomingRowPtr = std::exchange(other.incomingRowPtr, nullptr);
0182       incomingColIdx = std::exchange(other.incomingColIdx, nullptr);
0183       numNodes = std::exchange(other.numNodes, 0);
0184       numEdges = std::exchange(other.numEdges, 0);
0185       stream = std::exchange(other.stream, nullptr);
0186     }
0187     return *this;
0188   }
0189 
0190   void setStream(cudaStream_t owningStream) { stream = owningStream; }
0191 
0192   void reset() noexcept {
0193     cudaFreeAsyncNoThrow(rowPtr, stream);
0194     cudaFreeAsyncNoThrow(colIdx, stream);
0195     cudaFreeAsyncNoThrow(edgeWeight, stream);
0196     cudaFreeAsyncNoThrow(incomingRowPtr, stream);
0197     cudaFreeAsyncNoThrow(incomingColIdx, stream);
0198     rowPtr = nullptr;
0199     colIdx = nullptr;
0200     edgeWeight = nullptr;
0201     incomingRowPtr = nullptr;
0202     incomingColIdx = nullptr;
0203     numNodes = 0;
0204     numEdges = 0;
0205   }
0206 
0207  private:
0208   cudaStream_t stream = nullptr;
0209 };
0210 
0211 struct DpCudaState {
0212   float *bestScore = nullptr;
0213   int *bestChild = nullptr;
0214   unsigned char *sourceMask = nullptr;
0215   std::size_t numNodes = 0;
0216 
0217   DpCudaState() = default;
0218   ~DpCudaState() { reset(); }
0219 
0220   DpCudaState(const DpCudaState &) = delete;
0221   DpCudaState &operator=(const DpCudaState &) = delete;
0222 
0223   DpCudaState(DpCudaState &&other) noexcept { *this = std::move(other); }
0224 
0225   DpCudaState &operator=(DpCudaState &&other) noexcept {
0226     if (this != &other) {
0227       reset();
0228       bestScore = std::exchange(other.bestScore, nullptr);
0229       bestChild = std::exchange(other.bestChild, nullptr);
0230       sourceMask = std::exchange(other.sourceMask, nullptr);
0231       numNodes = std::exchange(other.numNodes, 0);
0232       stream = std::exchange(other.stream, nullptr);
0233     }
0234     return *this;
0235   }
0236 
0237   void setStream(cudaStream_t owningStream) { stream = owningStream; }
0238 
0239   void reset() noexcept {
0240     cudaFreeAsyncNoThrow(bestScore, stream);
0241     cudaFreeAsyncNoThrow(bestChild, stream);
0242     cudaFreeAsyncNoThrow(sourceMask, stream);
0243     bestScore = nullptr;
0244     bestChild = nullptr;
0245     sourceMask = nullptr;
0246     numNodes = 0;
0247   }
0248 
0249  private:
0250   cudaStream_t stream = nullptr;
0251 };
0252 
0253 std::vector<int> orderedSimpleComponentNodes(const std::vector<int> &nodes,
0254                                              int component,
0255                                              const std::vector<int> &labels,
0256                                              const std::vector<int> &inDegree,
0257                                              const std::vector<int> &nextNode) {
0258   if (nodes.empty()) {
0259     return {};
0260   }
0261 
0262   int start = nodes.front();
0263   for (int node : nodes) {
0264     if (inDegree.at(node) == 0) {
0265       start = node;
0266       break;
0267     }
0268   }
0269 
0270   std::vector<int> ordered;
0271   ordered.reserve(nodes.size());
0272   int node = start;
0273   while (node >= 0 && labels.at(node) == component &&
0274          ordered.size() < nodes.size()) {
0275     ordered.push_back(node);
0276     node = nextNode.at(node);
0277   }
0278 
0279   if (ordered.size() != nodes.size()) {
0280     ordered = nodes;
0281     std::sort(ordered.begin(), ordered.end());
0282   }
0283   return ordered;
0284 }
0285 
0286 float minRootScore(const std::string &pathMetric) {
0287   if (pathMetric == "score_weighted_length") {
0288     return 0.0F;
0289   }
0290   if (pathMetric == "length") {
0291     return 2.0F;
0292   }
0293   throw std::invalid_argument(
0294       "DWalkTrackBuilding pathMetric must be 'score_weighted_length' or "
0295       "'length'");
0296 }
0297 
0298 class DisjointSet {
0299  public:
0300   explicit DisjointSet(std::size_t size) : m_parent(size), m_rank(size, 0) {
0301     for (std::size_t i = 0; i < size; ++i) {
0302       m_parent[i] = static_cast<int>(i);
0303     }
0304   }
0305 
0306   int find(int value) {
0307     int parent = m_parent.at(value);
0308     if (parent != value) {
0309       parent = find(parent);
0310       m_parent.at(value) = parent;
0311     }
0312     return parent;
0313   }
0314 
0315   void unite(int lhs, int rhs) {
0316     int lhsRoot = find(lhs);
0317     int rhsRoot = find(rhs);
0318     if (lhsRoot == rhsRoot) {
0319       return;
0320     }
0321     if (m_rank.at(lhsRoot) < m_rank.at(rhsRoot)) {
0322       std::swap(lhsRoot, rhsRoot);
0323     }
0324     m_parent.at(rhsRoot) = lhsRoot;
0325     if (m_rank.at(lhsRoot) == m_rank.at(rhsRoot)) {
0326       ++m_rank.at(lhsRoot);
0327     }
0328   }
0329 
0330  private:
0331   std::vector<int> m_parent;
0332   std::vector<int> m_rank;
0333 };
0334 
0335 __device__ void atomicMaxFloat(float *address, float value) {
0336   int *addressAsInt = reinterpret_cast<int *>(address);
0337   int old = *addressAsInt;
0338   int assumed;
0339   do {
0340     assumed = old;
0341     old = atomicCAS(addressAsInt, assumed,
0342                     __float_as_int(fmaxf(value, __int_as_float(assumed))));
0343   } while (assumed != old);
0344 }
0345 
0346 __global__ void orientEdgesKernel(
0347     std::size_t numEdges, std::size_t numNodes, std::size_t numFeatures,
0348     std::size_t radialFeatureIndex, const std::int64_t *edgeIndex,
0349     const float *scores, const float *nodeFeatures, int *srcOut, int *dstOut,
0350     float *scoreOut, unsigned char *validEdge) {
0351   std::size_t edge = blockIdx.x * blockDim.x + threadIdx.x;
0352   if (edge >= numEdges) {
0353     return;
0354   }
0355 
0356   int src = static_cast<int>(edgeIndex[edge]);
0357   int dst = static_cast<int>(edgeIndex[numEdges + edge]);
0358   if (src < 0 || dst < 0 || src >= static_cast<int>(numNodes) ||
0359       dst >= static_cast<int>(numNodes) || src == dst) {
0360     srcOut[edge] = 0;
0361     dstOut[edge] = 0;
0362     scoreOut[edge] = 0.0F;
0363     validEdge[edge] = 0;
0364     return;
0365   }
0366 
0367   float srcR = nodeFeatures[src * numFeatures + radialFeatureIndex];
0368   float dstR = nodeFeatures[dst * numFeatures + radialFeatureIndex];
0369   if (srcR > dstR || (srcR == dstR && src > dst)) {
0370     int tmp = src;
0371     src = dst;
0372     dst = tmp;
0373   }
0374 
0375   srcOut[edge] = src;
0376   dstOut[edge] = dst;
0377   scoreOut[edge] = scores[edge];
0378   validEdge[edge] = 1;
0379 }
0380 
0381 __global__ void markActiveNodesKernel(std::size_t numEdges, const int *src,
0382                                       const int *dst,
0383                                       const unsigned char *validEdge,
0384                                       unsigned char *activeNodes) {
0385   std::size_t edge = blockIdx.x * blockDim.x + threadIdx.x;
0386   if (edge >= numEdges || validEdge[edge] == 0) {
0387     return;
0388   }
0389   activeNodes[src[edge]] = 1;
0390   activeNodes[dst[edge]] = 1;
0391 }
0392 
0393 __global__ void computeComponentDegreeKernel(std::size_t numEdges,
0394                                              const int *src, const int *dst,
0395                                              const unsigned char *validEdge,
0396                                              int *inDegree, int *outDegree) {
0397   std::size_t edge = blockIdx.x * blockDim.x + threadIdx.x;
0398   if (edge >= numEdges || validEdge[edge] == 0) {
0399     return;
0400   }
0401 
0402   int s = src[edge];
0403   int d = dst[edge];
0404   atomicAdd(outDegree + s, 1);
0405   atomicAdd(inDegree + d, 1);
0406 }
0407 
0408 __global__ void countActiveComponentNodesKernel(
0409     std::size_t numNodes, const unsigned char *activeNodes, const int *labels,
0410     int *componentSizes) {
0411   std::size_t node = blockIdx.x * blockDim.x + threadIdx.x;
0412   if (node >= numNodes || activeNodes[node] == 0) {
0413     return;
0414   }
0415   atomicAdd(componentSizes + labels[node], 1);
0416 }
0417 
0418 __global__ void markBadComponentsKernel(std::size_t numNodes,
0419                                         const unsigned char *activeNodes,
0420                                         const int *labels, const int *inDegree,
0421                                         const int *outDegree,
0422                                         int *badComponents) {
0423   std::size_t node = blockIdx.x * blockDim.x + threadIdx.x;
0424   if (node >= numNodes || activeNodes[node] == 0) {
0425     return;
0426   }
0427 
0428   if (max(inDegree[node], outDegree[node]) > 1) {
0429     badComponents[labels[node]] = 1;
0430   }
0431 }
0432 
0433 __global__ void buildInitialComponentMasksKernel(
0434     std::size_t numNodes, const unsigned char *activeNodes, const int *labels,
0435     const int *componentSizes, const int *badComponents, int minCandidateSize,
0436     unsigned char *simpleNodeMask, unsigned char *complexNodeMask) {
0437   std::size_t node = blockIdx.x * blockDim.x + threadIdx.x;
0438   if (node >= numNodes || activeNodes[node] == 0) {
0439     if (node < numNodes) {
0440       simpleNodeMask[node] = 0;
0441       complexNodeMask[node] = 0;
0442     }
0443     return;
0444   }
0445 
0446   int label = labels[node];
0447   bool isLarge = componentSizes[label] >= minCandidateSize;
0448   bool isSimple = isLarge && badComponents[label] == 0;
0449   simpleNodeMask[node] = isSimple ? 1 : 0;
0450   complexNodeMask[node] = (isLarge && !isSimple) ? 1 : 0;
0451 }
0452 
0453 __global__ void maskEdgesByActiveNodesKernel(std::size_t numEdges,
0454                                              const int *src, const int *dst,
0455                                              const unsigned char *activeNodes,
0456                                              unsigned char *keepMask) {
0457   std::size_t edge = blockIdx.x * blockDim.x + threadIdx.x;
0458   if (edge >= numEdges) {
0459     return;
0460   }
0461   keepMask[edge] = activeNodes[src[edge]] && activeNodes[dst[edge]] ? 1 : 0;
0462 }
0463 
0464 __global__ void deactivateSelectedNodesKernel(std::size_t numSelected,
0465                                               const int *selectedNodes,
0466                                               unsigned char *activeNodes) {
0467   std::size_t i = blockIdx.x * blockDim.x + threadIdx.x;
0468   if (i >= numSelected) {
0469     return;
0470   }
0471   int node = selectedNodes[i];
0472   if (node >= 0) {
0473     activeNodes[node] = 0;
0474   }
0475 }
0476 
0477 __global__ void fillCountsKernel(std::size_t numEdges, const int *nodes,
0478                                  int *counts) {
0479   std::size_t edge = blockIdx.x * blockDim.x + threadIdx.x;
0480   if (edge >= numEdges) {
0481     return;
0482   }
0483   atomicAdd(counts + nodes[edge], 1);
0484 }
0485 
0486 __global__ void scatterSortedOutgoingKernel(
0487     std::size_t numEdges, const int *sortedSrc, const int *sortedDst,
0488     const float *sortedScore, const int *rowPtr, int *cursor, int *colIdx,
0489     float *edgeWeight, bool useLengthMetric) {
0490   std::size_t edge = blockIdx.x * blockDim.x + threadIdx.x;
0491   if (edge >= numEdges) {
0492     return;
0493   }
0494   int src = sortedSrc[edge];
0495   int slot = atomicAdd(cursor + src, 1);
0496   colIdx[rowPtr[src] + slot] = sortedDst[edge];
0497   edgeWeight[rowPtr[src] + slot] = useLengthMetric ? 1.0F : sortedScore[edge];
0498 }
0499 
0500 __global__ void scatterSortedIncomingKernel(std::size_t numEdges,
0501                                             const int *sortedDst,
0502                                             const int *sortedSrc,
0503                                             const int *incomingRowPtr,
0504                                             int *cursor, int *incomingColIdx) {
0505   std::size_t edge = blockIdx.x * blockDim.x + threadIdx.x;
0506   if (edge >= numEdges) {
0507     return;
0508   }
0509   int dst = sortedDst[edge];
0510   int slot = atomicAdd(cursor + dst, 1);
0511   incomingColIdx[incomingRowPtr[dst] + slot] = sortedSrc[edge];
0512 }
0513 
0514 __global__ void initFloatKernel(std::size_t size, float value, float *array) {
0515   std::size_t i = blockIdx.x * blockDim.x + threadIdx.x;
0516   if (i < size) {
0517     array[i] = value;
0518   }
0519 }
0520 
0521 __global__ void maxAddBestScoresKernel(std::size_t numEdges, const int *src,
0522                                        const int *dst, const float *score,
0523                                        const unsigned char *validEdge,
0524                                        const unsigned char *activeNodes,
0525                                        float *bestOutScore,
0526                                        float *bestInScore) {
0527   std::size_t edge = blockIdx.x * blockDim.x + threadIdx.x;
0528   if (edge >= numEdges) {
0529     return;
0530   }
0531   if (validEdge[edge] == 0) {
0532     return;
0533   }
0534   int s = src[edge];
0535   int d = dst[edge];
0536   if (!activeNodes[s] || !activeNodes[d]) {
0537     return;
0538   }
0539   atomicMaxFloat(bestOutScore + s, score[edge]);
0540   atomicMaxFloat(bestInScore + d, score[edge]);
0541 }
0542 
0543 __global__ void maxAddMaskKernel(std::size_t numEdges, const int *src,
0544                                  const int *dst, const float *score,
0545                                  const unsigned char *validEdge,
0546                                  const unsigned char *activeNodes,
0547                                  const float *bestOutScore,
0548                                  const float *bestInScore, float thMin,
0549                                  float thAdd, unsigned char *keepMask) {
0550   std::size_t edge = blockIdx.x * blockDim.x + threadIdx.x;
0551   if (edge >= numEdges) {
0552     return;
0553   }
0554   if (validEdge[edge] == 0) {
0555     keepMask[edge] = 0;
0556     return;
0557   }
0558 
0559   int s = src[edge];
0560   int d = dst[edge];
0561   if (!activeNodes[s] || !activeNodes[d]) {
0562     keepMask[edge] = 0;
0563     return;
0564   }
0565 
0566   float edgeScore = score[edge];
0567   bool maskAdd = edgeScore > thAdd;
0568   bool outgoingKeep = bestOutScore[s] >= thMin && edgeScore == bestOutScore[s];
0569   bool incomingKeep = bestInScore[d] >= thMin && edgeScore == bestInScore[d];
0570   keepMask[edge] =
0571       ((outgoingKeep || maskAdd) && (incomingKeep || maskAdd)) ? 1 : 0;
0572 }
0573 
0574 __global__ void initializeDpKernel(float *bestScore, int *bestChild,
0575                                    unsigned char *sourceMask, int *inDegree,
0576                                    int *remainingOutDegree,
0577                                    std::size_t numNodes) {
0578   std::size_t node = blockIdx.x * blockDim.x + threadIdx.x;
0579   if (node >= numNodes) {
0580     return;
0581   }
0582   bestScore[node] = 0.0F;
0583   bestChild[node] = -1;
0584   sourceMask[node] = 0;
0585   inDegree[node] = 0;
0586   remainingOutDegree[node] = 0;
0587 }
0588 
0589 __global__ void computeActiveDegreeKernel(const int *rowPtr, const int *colIdx,
0590                                           const unsigned char *activeNodes,
0591                                           int *inDegree,
0592                                           int *remainingOutDegree,
0593                                           std::size_t numNodes) {
0594   std::size_t node = blockIdx.x * blockDim.x + threadIdx.x;
0595   if (node >= numNodes || !activeNodes[node]) {
0596     return;
0597   }
0598 
0599   int activeOutDegree = 0;
0600   for (int edge = rowPtr[node]; edge < rowPtr[node + 1]; ++edge) {
0601     int child = colIdx[edge];
0602     if (activeNodes[child]) {
0603       ++activeOutDegree;
0604       atomicAdd(inDegree + child, 1);
0605     }
0606   }
0607   remainingOutDegree[node] = activeOutDegree;
0608 }
0609 
0610 __global__ void initializeFrontierKernel(const unsigned char *activeNodes,
0611                                          const int *remainingOutDegree,
0612                                          int *frontier, int *frontierSize,
0613                                          std::size_t numNodes) {
0614   std::size_t node = blockIdx.x * blockDim.x + threadIdx.x;
0615   if (node >= numNodes) {
0616     return;
0617   }
0618   if (activeNodes[node] && remainingOutDegree[node] == 0) {
0619     int slot = atomicAdd(frontierSize, 1);
0620     frontier[slot] = static_cast<int>(node);
0621   }
0622 }
0623 
0624 __global__ void processFrontierKernel(
0625     const int *frontier, int frontierSize, const int *rowPtr, const int *colIdx,
0626     const float *edgeWeight, const unsigned char *activeNodes,
0627     const float *bestScore, float *frontierBestScore, int *frontierBestChild) {
0628   int frontierIndex = blockIdx.x * blockDim.x + threadIdx.x;
0629   if (frontierIndex >= frontierSize) {
0630     return;
0631   }
0632 
0633   int node = frontier[frontierIndex];
0634   float bestValue = 0.0F;
0635   int bestNext = -1;
0636   for (int edge = rowPtr[node]; edge < rowPtr[node + 1]; ++edge) {
0637     int child = colIdx[edge];
0638     if (!activeNodes[child]) {
0639       continue;
0640     }
0641     float candidate = edgeWeight[edge] + bestScore[child];
0642     if (candidate > bestValue) {
0643       bestValue = candidate;
0644       bestNext = child;
0645     }
0646   }
0647 
0648   frontierBestScore[frontierIndex] = bestValue;
0649   frontierBestChild[frontierIndex] = bestNext;
0650 }
0651 
0652 __global__ void finalizeFrontierKernel(const int *frontier, int frontierSize,
0653                                        const float *frontierBestScore,
0654                                        const int *frontierBestChild,
0655                                        float *bestScore, int *bestChild) {
0656   int frontierIndex = blockIdx.x * blockDim.x + threadIdx.x;
0657   if (frontierIndex >= frontierSize) {
0658     return;
0659   }
0660 
0661   int node = frontier[frontierIndex];
0662   bestScore[node] = frontierBestScore[frontierIndex];
0663   bestChild[node] = frontierBestChild[frontierIndex];
0664 }
0665 
0666 __global__ void enqueueParentFrontierKernel(
0667     const int *frontier, int frontierSize, const int *incomingRowPtr,
0668     const int *incomingColIdx, const unsigned char *activeNodes,
0669     int *remainingOutDegree, int *nextFrontier, int *nextFrontierSize) {
0670   int frontierIndex = blockIdx.x * blockDim.x + threadIdx.x;
0671   if (frontierIndex >= frontierSize) {
0672     return;
0673   }
0674 
0675   int node = frontier[frontierIndex];
0676   for (int edge = incomingRowPtr[node]; edge < incomingRowPtr[node + 1];
0677        ++edge) {
0678     int parent = incomingColIdx[edge];
0679     if (!activeNodes[parent]) {
0680       continue;
0681     }
0682     int oldValue = atomicSub(remainingOutDegree + parent, 1);
0683     if (oldValue == 1) {
0684       int slot = atomicAdd(nextFrontierSize, 1);
0685       nextFrontier[slot] = parent;
0686     }
0687   }
0688 }
0689 
0690 __global__ void buildSourceMaskKernel(const unsigned char *activeNodes,
0691                                       const int *inDegree,
0692                                       const float *bestScore,
0693                                       unsigned char *sourceMask,
0694                                       std::size_t numNodes) {
0695   std::size_t node = blockIdx.x * blockDim.x + threadIdx.x;
0696   if (node >= numNodes) {
0697     return;
0698   }
0699   sourceMask[node] =
0700       activeNodes[node] && inDegree[node] == 0 && bestScore[node] > 0.0F ? 1
0701                                                                          : 0;
0702 }
0703 
0704 __global__ void selectComponentScoresKernel(std::size_t numNodes,
0705                                             const int *componentLabels,
0706                                             const unsigned char *sourceMask,
0707                                             const float *bestScore,
0708                                             float *componentBestScore) {
0709   std::size_t node = blockIdx.x * blockDim.x + threadIdx.x;
0710   if (node >= numNodes || sourceMask[node] == 0) {
0711     return;
0712   }
0713   int component = componentLabels[node];
0714   if (component < 0) {
0715     return;
0716   }
0717   atomicMaxFloat(componentBestScore + component, bestScore[node]);
0718 }
0719 
0720 __global__ void selectRootsKernel(std::size_t numNodes,
0721                                   const int *componentLabels,
0722                                   const unsigned char *sourceMask,
0723                                   const float *bestScore,
0724                                   const float *componentBestScore,
0725                                   float minRootScore, int *selectedRoots) {
0726   std::size_t node = blockIdx.x * blockDim.x + threadIdx.x;
0727   if (node >= numNodes || sourceMask[node] == 0) {
0728     return;
0729   }
0730   int component = componentLabels[node];
0731   if (component < 0 || componentBestScore[component] <= minRootScore) {
0732     return;
0733   }
0734   if (bestScore[node] == componentBestScore[component]) {
0735     atomicCAS(selectedRoots + component, -1, static_cast<int>(node));
0736   }
0737 }
0738 
0739 __global__ void traceSelectedPathsKernel(const int *bestChild,
0740                                          const int *selectedRoots,
0741                                          int *selectedTrackLabels,
0742                                          int *selectedNodes, int *selectedCount,
0743                                          std::size_t numNodes,
0744                                          std::size_t numPaths) {
0745   std::size_t pathIndex = blockIdx.x * blockDim.x + threadIdx.x;
0746   if (pathIndex >= numPaths) {
0747     return;
0748   }
0749 
0750   int node = selectedRoots[pathIndex];
0751   if (node < 0) {
0752     return;
0753   }
0754   std::size_t step = 0;
0755   while (node != -1 && step < numNodes) {
0756     int slot = atomicAdd(selectedCount, 1);
0757     selectedTrackLabels[slot] = static_cast<int>(pathIndex);
0758     selectedNodes[slot] = node;
0759     node = bestChild[node];
0760     ++step;
0761   }
0762 }
0763 
0764 __global__ void remapEdgesKernel(std::size_t numEdges, int *src, int *dst,
0765                                  const int *oldToNew) {
0766   std::size_t edge = blockIdx.x * blockDim.x + threadIdx.x;
0767   if (edge >= numEdges) {
0768     return;
0769   }
0770   src[edge] = oldToNew[src[edge]];
0771   dst[edge] = oldToNew[dst[edge]];
0772 }
0773 
0774 std::vector<Edge> createOrientedEdgesCuda(
0775     const ActsPlugins::Tensor<std::int64_t> &edgeTensor,
0776     const ActsPlugins::Tensor<float> &scoreTensor,
0777     const ActsPlugins::Tensor<float> &featureTensor,
0778     std::size_t radialFeatureIndex, cudaStream_t stream,
0779     DeviceOrientedEdges &deviceGraph, int **cudaInitialLabels = nullptr,
0780     int *initialNumComponents = nullptr,
0781     std::vector<int> *initialLabels = nullptr) {
0782   const auto numNodes = featureTensor.shape().at(0);
0783   const auto numFeatures = featureTensor.shape().at(1);
0784   const auto numEdges = edgeTensor.shape().at(1);
0785 
0786   deviceGraph.setStream(stream);
0787   deviceGraph.numEdges = numEdges;
0788   deviceGraph.numNodes = numNodes;
0789   ACTS_CUDA_CHECK(
0790       cudaMallocAsync(&deviceGraph.src, numEdges * sizeof(int), stream));
0791   ACTS_CUDA_CHECK(
0792       cudaMallocAsync(&deviceGraph.dst, numEdges * sizeof(int), stream));
0793   ACTS_CUDA_CHECK(
0794       cudaMallocAsync(&deviceGraph.score, numEdges * sizeof(float), stream));
0795   ACTS_CUDA_CHECK(cudaMallocAsync(&deviceGraph.valid,
0796                                   numEdges * sizeof(unsigned char), stream));
0797   ACTS_CUDA_CHECK(cudaMallocAsync(&deviceGraph.activeNodes,
0798                                   numNodes * sizeof(unsigned char), stream));
0799   ACTS_CUDA_CHECK(cudaMemsetAsync(deviceGraph.activeNodes, 0,
0800                                   numNodes * sizeof(unsigned char), stream));
0801 
0802   const dim3 edgeGrid((numEdges + kBlockSize - 1) / kBlockSize);
0803   orientEdgesKernel<<<edgeGrid, kBlockSize, 0, stream>>>(
0804       numEdges, numNodes, numFeatures, radialFeatureIndex, edgeTensor.data(),
0805       scoreTensor.data(), featureTensor.data(), deviceGraph.src,
0806       deviceGraph.dst, deviceGraph.score, deviceGraph.valid);
0807   ACTS_CUDA_CHECK(cudaGetLastError());
0808   markActiveNodesKernel<<<edgeGrid, kBlockSize, 0, stream>>>(
0809       numEdges, deviceGraph.src, deviceGraph.dst, deviceGraph.valid,
0810       deviceGraph.activeNodes);
0811   ACTS_CUDA_CHECK(cudaGetLastError());
0812 
0813   if (cudaInitialLabels != nullptr || initialLabels != nullptr ||
0814       initialNumComponents != nullptr) {
0815     int *cudaLabels{};
0816     ACTS_CUDA_CHECK(
0817         cudaMallocAsync(&cudaLabels, numNodes * sizeof(int), stream));
0818     int numComponents = ActsPlugins::detail::connectedComponentsCuda(
0819         numEdges, deviceGraph.src, deviceGraph.dst, numNodes, cudaLabels,
0820         stream, false);
0821     if (initialNumComponents != nullptr) {
0822       *initialNumComponents = numComponents;
0823     }
0824     if (initialLabels != nullptr) {
0825       initialLabels->assign(numNodes, -1);
0826       ACTS_CUDA_CHECK(cudaMemcpyAsync(initialLabels->data(), cudaLabels,
0827                                       numNodes * sizeof(int),
0828                                       cudaMemcpyDeviceToHost, stream));
0829     }
0830     if (cudaInitialLabels != nullptr) {
0831       *cudaInitialLabels = cudaLabels;
0832     } else {
0833       ACTS_CUDA_CHECK(cudaFreeAsync(cudaLabels, stream));
0834     }
0835   }
0836 
0837   std::vector<int> src(numEdges);
0838   std::vector<int> dst(numEdges);
0839   std::vector<unsigned char> valid(numEdges);
0840 
0841   ACTS_CUDA_CHECK(cudaMemcpyAsync(src.data(), deviceGraph.src,
0842                                   numEdges * sizeof(int),
0843                                   cudaMemcpyDeviceToHost, stream));
0844   ACTS_CUDA_CHECK(cudaMemcpyAsync(dst.data(), deviceGraph.dst,
0845                                   numEdges * sizeof(int),
0846                                   cudaMemcpyDeviceToHost, stream));
0847   ACTS_CUDA_CHECK(cudaMemcpyAsync(valid.data(), deviceGraph.valid,
0848                                   numEdges * sizeof(unsigned char),
0849                                   cudaMemcpyDeviceToHost, stream));
0850 
0851   ACTS_CUDA_CHECK(cudaStreamSynchronize(stream));
0852 
0853   std::vector<Edge> edges;
0854   edges.reserve(numEdges);
0855   for (std::size_t edge = 0; edge < numEdges; ++edge) {
0856     if (valid.at(edge) != 0) {
0857       edges.push_back({src.at(edge), dst.at(edge)});
0858     }
0859   }
0860   return edges;
0861 }
0862 
0863 DeviceCompactEdges compactDeviceEdges(const int *src, const int *dst,
0864                                       const float *score,
0865                                       const unsigned char *keepMask,
0866                                       std::size_t numEdges,
0867                                       std::size_t numNodes,
0868                                       cudaStream_t stream) {
0869   DeviceCompactEdges compact;
0870   compact.setStream(stream);
0871   compact.numNodes = numNodes;
0872   compact.numEdges = thrust::count(thrust::device.on(stream), keepMask,
0873                                    keepMask + numEdges, 1);
0874   if (compact.numEdges == 0) {
0875     return compact;
0876   }
0877 
0878   ACTS_CUDA_CHECK(
0879       cudaMallocAsync(&compact.src, compact.numEdges * sizeof(int), stream));
0880   ACTS_CUDA_CHECK(
0881       cudaMallocAsync(&compact.dst, compact.numEdges * sizeof(int), stream));
0882   ACTS_CUDA_CHECK(cudaMallocAsync(&compact.score,
0883                                   compact.numEdges * sizeof(float), stream));
0884 
0885   thrust::copy_if(thrust::device.on(stream), src, src + numEdges, keepMask,
0886                   compact.src, IsNonZero{});
0887   thrust::copy_if(thrust::device.on(stream), dst, dst + numEdges, keepMask,
0888                   compact.dst, IsNonZero{});
0889   thrust::copy_if(thrust::device.on(stream), score, score + numEdges, keepMask,
0890                   compact.score, IsNonZero{});
0891   return compact;
0892 }
0893 
0894 DeviceCompactEdges maxAddCompactDeviceEdgesCuda(
0895     const DeviceOrientedEdges &deviceGraph, const unsigned char *activeNodes,
0896     float thMin, float thAdd, cudaStream_t stream) {
0897   if (deviceGraph.numEdges == 0) {
0898     return {};
0899   }
0900 
0901   float *cudaBestOut{};
0902   float *cudaBestIn{};
0903   unsigned char *cudaKeep{};
0904   ACTS_CUDA_CHECK(cudaMallocAsync(
0905       &cudaBestOut, deviceGraph.numNodes * sizeof(float), stream));
0906   ACTS_CUDA_CHECK(cudaMallocAsync(
0907       &cudaBestIn, deviceGraph.numNodes * sizeof(float), stream));
0908   ACTS_CUDA_CHECK(cudaMallocAsync(
0909       &cudaKeep, deviceGraph.numEdges * sizeof(unsigned char), stream));
0910 
0911   const dim3 nodeGrid((deviceGraph.numNodes + kBlockSize - 1) / kBlockSize);
0912   const dim3 edgeGrid((deviceGraph.numEdges + kBlockSize - 1) / kBlockSize);
0913   initFloatKernel<<<nodeGrid, kBlockSize, 0, stream>>>(
0914       deviceGraph.numNodes, -std::numeric_limits<float>::infinity(),
0915       cudaBestOut);
0916   initFloatKernel<<<nodeGrid, kBlockSize, 0, stream>>>(
0917       deviceGraph.numNodes, -std::numeric_limits<float>::infinity(),
0918       cudaBestIn);
0919   ACTS_CUDA_CHECK(cudaGetLastError());
0920   maxAddBestScoresKernel<<<edgeGrid, kBlockSize, 0, stream>>>(
0921       deviceGraph.numEdges, deviceGraph.src, deviceGraph.dst, deviceGraph.score,
0922       deviceGraph.valid, activeNodes, cudaBestOut, cudaBestIn);
0923   ACTS_CUDA_CHECK(cudaGetLastError());
0924   maxAddMaskKernel<<<edgeGrid, kBlockSize, 0, stream>>>(
0925       deviceGraph.numEdges, deviceGraph.src, deviceGraph.dst, deviceGraph.score,
0926       deviceGraph.valid, activeNodes, cudaBestOut, cudaBestIn, thMin, thAdd,
0927       cudaKeep);
0928   ACTS_CUDA_CHECK(cudaGetLastError());
0929 
0930   auto compact = compactDeviceEdges(
0931       deviceGraph.src, deviceGraph.dst, deviceGraph.score, cudaKeep,
0932       deviceGraph.numEdges, deviceGraph.numNodes, stream);
0933 
0934   ACTS_CUDA_CHECK(cudaFreeAsync(cudaBestOut, stream));
0935   ACTS_CUDA_CHECK(cudaFreeAsync(cudaBestIn, stream));
0936   ACTS_CUDA_CHECK(cudaFreeAsync(cudaKeep, stream));
0937   return compact;
0938 }
0939 
0940 DeviceCompactEdges compactActiveEdgesCuda(const DeviceCompactEdges &edges,
0941                                           const unsigned char *activeNodes,
0942                                           cudaStream_t stream) {
0943   if (edges.numEdges == 0) {
0944     return {};
0945   }
0946 
0947   unsigned char *cudaKeep{};
0948   ACTS_CUDA_CHECK(cudaMallocAsync(
0949       &cudaKeep, edges.numEdges * sizeof(unsigned char), stream));
0950   const dim3 edgeGrid((edges.numEdges + kBlockSize - 1) / kBlockSize);
0951   maskEdgesByActiveNodesKernel<<<edgeGrid, kBlockSize, 0, stream>>>(
0952       edges.numEdges, edges.src, edges.dst, activeNodes, cudaKeep);
0953   ACTS_CUDA_CHECK(cudaGetLastError());
0954   auto compact = compactDeviceEdges(edges.src, edges.dst, edges.score, cudaKeep,
0955                                     edges.numEdges, edges.numNodes, stream);
0956   ACTS_CUDA_CHECK(cudaFreeAsync(cudaKeep, stream));
0957   return compact;
0958 }
0959 
0960 void classifyInitialComponentsCuda(const DeviceOrientedEdges &deviceGraph,
0961                                    const int *cudaLabels, int numComponents,
0962                                    int minCandidateSize,
0963                                    unsigned char **cudaSimpleNodeMask,
0964                                    unsigned char **cudaComplexNodeMask,
0965                                    cudaStream_t stream) {
0966   int *cudaInDegree{};
0967   int *cudaOutDegree{};
0968   int *cudaComponentSizes{};
0969   int *cudaBadComponents{};
0970   ACTS_CUDA_CHECK(cudaMallocAsync(&cudaInDegree,
0971                                   deviceGraph.numNodes * sizeof(int), stream));
0972   ACTS_CUDA_CHECK(cudaMallocAsync(&cudaOutDegree,
0973                                   deviceGraph.numNodes * sizeof(int), stream));
0974   ACTS_CUDA_CHECK(cudaMallocAsync(&cudaComponentSizes,
0975                                   numComponents * sizeof(int), stream));
0976   ACTS_CUDA_CHECK(
0977       cudaMallocAsync(&cudaBadComponents, numComponents * sizeof(int), stream));
0978   ACTS_CUDA_CHECK(cudaMallocAsync(cudaSimpleNodeMask,
0979                                   deviceGraph.numNodes * sizeof(unsigned char),
0980                                   stream));
0981   ACTS_CUDA_CHECK(cudaMallocAsync(cudaComplexNodeMask,
0982                                   deviceGraph.numNodes * sizeof(unsigned char),
0983                                   stream));
0984 
0985   ACTS_CUDA_CHECK(cudaMemsetAsync(cudaInDegree, 0,
0986                                   deviceGraph.numNodes * sizeof(int), stream));
0987   ACTS_CUDA_CHECK(cudaMemsetAsync(cudaOutDegree, 0,
0988                                   deviceGraph.numNodes * sizeof(int), stream));
0989   ACTS_CUDA_CHECK(cudaMemsetAsync(cudaComponentSizes, 0,
0990                                   numComponents * sizeof(int), stream));
0991   ACTS_CUDA_CHECK(cudaMemsetAsync(cudaBadComponents, 0,
0992                                   numComponents * sizeof(int), stream));
0993 
0994   const dim3 edgeGrid((deviceGraph.numEdges + kBlockSize - 1) / kBlockSize);
0995   const dim3 nodeGrid((deviceGraph.numNodes + kBlockSize - 1) / kBlockSize);
0996   computeComponentDegreeKernel<<<edgeGrid, kBlockSize, 0, stream>>>(
0997       deviceGraph.numEdges, deviceGraph.src, deviceGraph.dst, deviceGraph.valid,
0998       cudaInDegree, cudaOutDegree);
0999   ACTS_CUDA_CHECK(cudaGetLastError());
1000   countActiveComponentNodesKernel<<<nodeGrid, kBlockSize, 0, stream>>>(
1001       deviceGraph.numNodes, deviceGraph.activeNodes, cudaLabels,
1002       cudaComponentSizes);
1003   ACTS_CUDA_CHECK(cudaGetLastError());
1004   markBadComponentsKernel<<<nodeGrid, kBlockSize, 0, stream>>>(
1005       deviceGraph.numNodes, deviceGraph.activeNodes, cudaLabels, cudaInDegree,
1006       cudaOutDegree, cudaBadComponents);
1007   ACTS_CUDA_CHECK(cudaGetLastError());
1008   buildInitialComponentMasksKernel<<<nodeGrid, kBlockSize, 0, stream>>>(
1009       deviceGraph.numNodes, deviceGraph.activeNodes, cudaLabels,
1010       cudaComponentSizes, cudaBadComponents, minCandidateSize,
1011       *cudaSimpleNodeMask, *cudaComplexNodeMask);
1012   ACTS_CUDA_CHECK(cudaGetLastError());
1013 
1014   ACTS_CUDA_CHECK(cudaFreeAsync(cudaInDegree, stream));
1015   ACTS_CUDA_CHECK(cudaFreeAsync(cudaOutDegree, stream));
1016   ACTS_CUDA_CHECK(cudaFreeAsync(cudaComponentSizes, stream));
1017   ACTS_CUDA_CHECK(cudaFreeAsync(cudaBadComponents, stream));
1018 }
1019 
1020 DeviceCsrGraph buildCsrGraphCuda(const DeviceCompactEdges &edges,
1021                                  const std::string &pathMetric,
1022                                  cudaStream_t stream) {
1023   DeviceCsrGraph graph;
1024   graph.setStream(stream);
1025   graph.numNodes = edges.numNodes;
1026   graph.numEdges = edges.numEdges;
1027   if (edges.numEdges == 0) {
1028     return graph;
1029   }
1030 
1031   int *sortedSrc{};
1032   int *sortedDst{};
1033   float *sortedScore{};
1034   int *sortedIncomingDst{};
1035   int *sortedIncomingSrc{};
1036   int *rowCounts{};
1037   int *incomingCounts{};
1038   int *cursor{};
1039   int *incomingCursor{};
1040 
1041   ACTS_CUDA_CHECK(
1042       cudaMallocAsync(&sortedSrc, edges.numEdges * sizeof(int), stream));
1043   ACTS_CUDA_CHECK(
1044       cudaMallocAsync(&sortedDst, edges.numEdges * sizeof(int), stream));
1045   ACTS_CUDA_CHECK(
1046       cudaMallocAsync(&sortedScore, edges.numEdges * sizeof(float), stream));
1047   ACTS_CUDA_CHECK(cudaMallocAsync(&sortedIncomingDst,
1048                                   edges.numEdges * sizeof(int), stream));
1049   ACTS_CUDA_CHECK(cudaMallocAsync(&sortedIncomingSrc,
1050                                   edges.numEdges * sizeof(int), stream));
1051   ACTS_CUDA_CHECK(
1052       cudaMallocAsync(&rowCounts, edges.numNodes * sizeof(int), stream));
1053   ACTS_CUDA_CHECK(
1054       cudaMallocAsync(&incomingCounts, edges.numNodes * sizeof(int), stream));
1055   ACTS_CUDA_CHECK(
1056       cudaMallocAsync(&cursor, edges.numNodes * sizeof(int), stream));
1057   ACTS_CUDA_CHECK(
1058       cudaMallocAsync(&incomingCursor, edges.numNodes * sizeof(int), stream));
1059   ACTS_CUDA_CHECK(cudaMallocAsync(&graph.rowPtr,
1060                                   (edges.numNodes + 1) * sizeof(int), stream));
1061   ACTS_CUDA_CHECK(
1062       cudaMallocAsync(&graph.colIdx, edges.numEdges * sizeof(int), stream));
1063   ACTS_CUDA_CHECK(cudaMallocAsync(&graph.edgeWeight,
1064                                   edges.numEdges * sizeof(float), stream));
1065   ACTS_CUDA_CHECK(cudaMallocAsync(&graph.incomingRowPtr,
1066                                   (edges.numNodes + 1) * sizeof(int), stream));
1067   ACTS_CUDA_CHECK(cudaMallocAsync(&graph.incomingColIdx,
1068                                   edges.numEdges * sizeof(int), stream));
1069 
1070   ACTS_CUDA_CHECK(
1071       cudaMemsetAsync(rowCounts, 0, edges.numNodes * sizeof(int), stream));
1072   ACTS_CUDA_CHECK(
1073       cudaMemsetAsync(incomingCounts, 0, edges.numNodes * sizeof(int), stream));
1074   const dim3 edgeGrid((edges.numEdges + kBlockSize - 1) / kBlockSize);
1075   fillCountsKernel<<<edgeGrid, kBlockSize, 0, stream>>>(edges.numEdges,
1076                                                         edges.src, rowCounts);
1077   fillCountsKernel<<<edgeGrid, kBlockSize, 0, stream>>>(
1078       edges.numEdges, edges.dst, incomingCounts);
1079   ACTS_CUDA_CHECK(cudaGetLastError());
1080 
1081   ACTS_CUDA_CHECK(cudaMemsetAsync(graph.rowPtr, 0, sizeof(int), stream));
1082   ACTS_CUDA_CHECK(
1083       cudaMemsetAsync(graph.incomingRowPtr, 0, sizeof(int), stream));
1084   thrust::inclusive_scan(thrust::device.on(stream), rowCounts,
1085                          rowCounts + edges.numNodes, graph.rowPtr + 1);
1086   thrust::inclusive_scan(thrust::device.on(stream), incomingCounts,
1087                          incomingCounts + edges.numNodes,
1088                          graph.incomingRowPtr + 1);
1089 
1090   ACTS_CUDA_CHECK(cudaMemcpyAsync(sortedSrc, edges.src,
1091                                   edges.numEdges * sizeof(int),
1092                                   cudaMemcpyDeviceToDevice, stream));
1093   ACTS_CUDA_CHECK(cudaMemcpyAsync(sortedDst, edges.dst,
1094                                   edges.numEdges * sizeof(int),
1095                                   cudaMemcpyDeviceToDevice, stream));
1096   ACTS_CUDA_CHECK(cudaMemcpyAsync(sortedScore, edges.score,
1097                                   edges.numEdges * sizeof(float),
1098                                   cudaMemcpyDeviceToDevice, stream));
1099   thrust::sort_by_key(
1100       thrust::device.on(stream), sortedSrc, sortedSrc + edges.numEdges,
1101       thrust::make_zip_iterator(thrust::make_tuple(sortedDst, sortedScore)));
1102 
1103   ACTS_CUDA_CHECK(cudaMemcpyAsync(sortedIncomingDst, edges.dst,
1104                                   edges.numEdges * sizeof(int),
1105                                   cudaMemcpyDeviceToDevice, stream));
1106   ACTS_CUDA_CHECK(cudaMemcpyAsync(sortedIncomingSrc, edges.src,
1107                                   edges.numEdges * sizeof(int),
1108                                   cudaMemcpyDeviceToDevice, stream));
1109   thrust::sort_by_key(thrust::device.on(stream), sortedIncomingDst,
1110                       sortedIncomingDst + edges.numEdges, sortedIncomingSrc);
1111 
1112   ACTS_CUDA_CHECK(
1113       cudaMemsetAsync(cursor, 0, edges.numNodes * sizeof(int), stream));
1114   ACTS_CUDA_CHECK(
1115       cudaMemsetAsync(incomingCursor, 0, edges.numNodes * sizeof(int), stream));
1116   scatterSortedOutgoingKernel<<<edgeGrid, kBlockSize, 0, stream>>>(
1117       edges.numEdges, sortedSrc, sortedDst, sortedScore, graph.rowPtr, cursor,
1118       graph.colIdx, graph.edgeWeight, pathMetric == "length");
1119   scatterSortedIncomingKernel<<<edgeGrid, kBlockSize, 0, stream>>>(
1120       edges.numEdges, sortedIncomingDst, sortedIncomingSrc,
1121       graph.incomingRowPtr, incomingCursor, graph.incomingColIdx);
1122   ACTS_CUDA_CHECK(cudaGetLastError());
1123 
1124   ACTS_CUDA_CHECK(cudaFreeAsync(sortedSrc, stream));
1125   ACTS_CUDA_CHECK(cudaFreeAsync(sortedDst, stream));
1126   ACTS_CUDA_CHECK(cudaFreeAsync(sortedScore, stream));
1127   ACTS_CUDA_CHECK(cudaFreeAsync(sortedIncomingDst, stream));
1128   ACTS_CUDA_CHECK(cudaFreeAsync(sortedIncomingSrc, stream));
1129   ACTS_CUDA_CHECK(cudaFreeAsync(rowCounts, stream));
1130   ACTS_CUDA_CHECK(cudaFreeAsync(incomingCounts, stream));
1131   ACTS_CUDA_CHECK(cudaFreeAsync(cursor, stream));
1132   ACTS_CUDA_CHECK(cudaFreeAsync(incomingCursor, stream));
1133   return graph;
1134 }
1135 
1136 DpCudaState runDpOnCsrCuda(const DeviceCsrGraph &graph,
1137                            const unsigned char *cudaActiveNodes,
1138                            cudaStream_t stream) {
1139   DpCudaState state;
1140   state.setStream(stream);
1141   state.numNodes = graph.numNodes;
1142   ACTS_CUDA_CHECK(cudaMallocAsync(&state.bestScore,
1143                                   graph.numNodes * sizeof(float), stream));
1144   ACTS_CUDA_CHECK(
1145       cudaMallocAsync(&state.bestChild, graph.numNodes * sizeof(int), stream));
1146   ACTS_CUDA_CHECK(cudaMallocAsync(
1147       &state.sourceMask, graph.numNodes * sizeof(unsigned char), stream));
1148   if (graph.numEdges == 0) {
1149     return state;
1150   }
1151 
1152   int *cudaInDegree{};
1153   int *cudaRemainingOutDegree{};
1154   int *cudaFrontier{};
1155   int *cudaNextFrontier{};
1156   int *cudaFrontierSize{};
1157   int *cudaNextFrontierSize{};
1158   float *cudaFrontierBestScore{};
1159   int *cudaFrontierBestChild{};
1160 
1161   ACTS_CUDA_CHECK(
1162       cudaMallocAsync(&cudaInDegree, graph.numNodes * sizeof(int), stream));
1163   ACTS_CUDA_CHECK(cudaMallocAsync(&cudaRemainingOutDegree,
1164                                   graph.numNodes * sizeof(int), stream));
1165   ACTS_CUDA_CHECK(
1166       cudaMallocAsync(&cudaFrontier, graph.numNodes * sizeof(int), stream));
1167   ACTS_CUDA_CHECK(
1168       cudaMallocAsync(&cudaNextFrontier, graph.numNodes * sizeof(int), stream));
1169   ACTS_CUDA_CHECK(cudaMallocAsync(&cudaFrontierSize, sizeof(int), stream));
1170   ACTS_CUDA_CHECK(cudaMallocAsync(&cudaNextFrontierSize, sizeof(int), stream));
1171   ACTS_CUDA_CHECK(cudaMallocAsync(&cudaFrontierBestScore,
1172                                   graph.numNodes * sizeof(float), stream));
1173   ACTS_CUDA_CHECK(cudaMallocAsync(&cudaFrontierBestChild,
1174                                   graph.numNodes * sizeof(int), stream));
1175 
1176   const dim3 nodeGrid((graph.numNodes + kBlockSize - 1) / kBlockSize);
1177   initializeDpKernel<<<nodeGrid, kBlockSize, 0, stream>>>(
1178       state.bestScore, state.bestChild, state.sourceMask, cudaInDegree,
1179       cudaRemainingOutDegree, graph.numNodes);
1180   ACTS_CUDA_CHECK(cudaGetLastError());
1181   computeActiveDegreeKernel<<<nodeGrid, kBlockSize, 0, stream>>>(
1182       graph.rowPtr, graph.colIdx, cudaActiveNodes, cudaInDegree,
1183       cudaRemainingOutDegree, graph.numNodes);
1184   ACTS_CUDA_CHECK(cudaGetLastError());
1185   ACTS_CUDA_CHECK(cudaMemsetAsync(cudaFrontierSize, 0, sizeof(int), stream));
1186   initializeFrontierKernel<<<nodeGrid, kBlockSize, 0, stream>>>(
1187       cudaActiveNodes, cudaRemainingOutDegree, cudaFrontier, cudaFrontierSize,
1188       graph.numNodes);
1189   ACTS_CUDA_CHECK(cudaGetLastError());
1190 
1191   int currentFrontierSize = 0;
1192   ACTS_CUDA_CHECK(cudaMemcpyAsync(&currentFrontierSize, cudaFrontierSize,
1193                                   sizeof(int), cudaMemcpyDeviceToHost, stream));
1194   ACTS_CUDA_CHECK(cudaStreamSynchronize(stream));
1195   while (currentFrontierSize > 0) {
1196     const dim3 frontierGrid((currentFrontierSize + kBlockSize - 1) /
1197                             kBlockSize);
1198     processFrontierKernel<<<frontierGrid, kBlockSize, 0, stream>>>(
1199         cudaFrontier, currentFrontierSize, graph.rowPtr, graph.colIdx,
1200         graph.edgeWeight, cudaActiveNodes, state.bestScore,
1201         cudaFrontierBestScore, cudaFrontierBestChild);
1202     ACTS_CUDA_CHECK(cudaGetLastError());
1203     finalizeFrontierKernel<<<frontierGrid, kBlockSize, 0, stream>>>(
1204         cudaFrontier, currentFrontierSize, cudaFrontierBestScore,
1205         cudaFrontierBestChild, state.bestScore, state.bestChild);
1206     ACTS_CUDA_CHECK(cudaGetLastError());
1207     ACTS_CUDA_CHECK(
1208         cudaMemsetAsync(cudaNextFrontierSize, 0, sizeof(int), stream));
1209     enqueueParentFrontierKernel<<<frontierGrid, kBlockSize, 0, stream>>>(
1210         cudaFrontier, currentFrontierSize, graph.incomingRowPtr,
1211         graph.incomingColIdx, cudaActiveNodes, cudaRemainingOutDegree,
1212         cudaNextFrontier, cudaNextFrontierSize);
1213     ACTS_CUDA_CHECK(cudaGetLastError());
1214     std::swap(cudaFrontier, cudaNextFrontier);
1215     std::swap(cudaFrontierSize, cudaNextFrontierSize);
1216     ACTS_CUDA_CHECK(cudaMemcpyAsync(&currentFrontierSize, cudaFrontierSize,
1217                                     sizeof(int), cudaMemcpyDeviceToHost,
1218                                     stream));
1219     ACTS_CUDA_CHECK(cudaStreamSynchronize(stream));
1220   }
1221 
1222   buildSourceMaskKernel<<<nodeGrid, kBlockSize, 0, stream>>>(
1223       cudaActiveNodes, cudaInDegree, state.bestScore, state.sourceMask,
1224       graph.numNodes);
1225   ACTS_CUDA_CHECK(cudaGetLastError());
1226 
1227   ACTS_CUDA_CHECK(cudaFreeAsync(cudaInDegree, stream));
1228   ACTS_CUDA_CHECK(cudaFreeAsync(cudaRemainingOutDegree, stream));
1229   ACTS_CUDA_CHECK(cudaFreeAsync(cudaFrontier, stream));
1230   ACTS_CUDA_CHECK(cudaFreeAsync(cudaNextFrontier, stream));
1231   ACTS_CUDA_CHECK(cudaFreeAsync(cudaFrontierSize, stream));
1232   ACTS_CUDA_CHECK(cudaFreeAsync(cudaNextFrontierSize, stream));
1233   ACTS_CUDA_CHECK(cudaFreeAsync(cudaFrontierBestScore, stream));
1234   ACTS_CUDA_CHECK(cudaFreeAsync(cudaFrontierBestChild, stream));
1235   return state;
1236 }
1237 
1238 std::pair<std::vector<int>, std::vector<int>> selectAndTracePathsCuda(
1239     const DpCudaState &dpState, const int *cudaComponentLabels,
1240     int numComponents, float minRootScore, unsigned char *cudaActiveNodes,
1241     cudaStream_t stream) {
1242   if (numComponents <= 0) {
1243     return {};
1244   }
1245 
1246   const std::size_t numNodes = dpState.numNodes;
1247   int *cudaSelectedRoots{};
1248   float *cudaComponentBestScore{};
1249   int *cudaSelectedTrackLabels{};
1250   int *cudaSelectedNodes{};
1251   int *cudaSelectedCount{};
1252   ACTS_CUDA_CHECK(
1253       cudaMallocAsync(&cudaSelectedRoots, numComponents * sizeof(int), stream));
1254   ACTS_CUDA_CHECK(cudaMallocAsync(&cudaComponentBestScore,
1255                                   numComponents * sizeof(float), stream));
1256   ACTS_CUDA_CHECK(cudaMallocAsync(&cudaSelectedTrackLabels,
1257                                   numNodes * sizeof(int), stream));
1258   ACTS_CUDA_CHECK(
1259       cudaMallocAsync(&cudaSelectedNodes, numNodes * sizeof(int), stream));
1260   ACTS_CUDA_CHECK(cudaMallocAsync(&cudaSelectedCount, sizeof(int), stream));
1261 
1262   const dim3 componentGrid((numComponents + kBlockSize - 1) / kBlockSize);
1263   initFloatKernel<<<componentGrid, kBlockSize, 0, stream>>>(
1264       numComponents, -std::numeric_limits<float>::infinity(),
1265       cudaComponentBestScore);
1266   ACTS_CUDA_CHECK(cudaGetLastError());
1267   ACTS_CUDA_CHECK(cudaMemsetAsync(cudaSelectedRoots, 0xff,
1268                                   numComponents * sizeof(int), stream));
1269   const dim3 nodeGrid((numNodes + kBlockSize - 1) / kBlockSize);
1270   selectComponentScoresKernel<<<nodeGrid, kBlockSize, 0, stream>>>(
1271       numNodes, cudaComponentLabels, dpState.sourceMask, dpState.bestScore,
1272       cudaComponentBestScore);
1273   ACTS_CUDA_CHECK(cudaGetLastError());
1274   selectRootsKernel<<<nodeGrid, kBlockSize, 0, stream>>>(
1275       numNodes, cudaComponentLabels, dpState.sourceMask, dpState.bestScore,
1276       cudaComponentBestScore, minRootScore, cudaSelectedRoots);
1277   ACTS_CUDA_CHECK(cudaGetLastError());
1278   ACTS_CUDA_CHECK(cudaMemsetAsync(cudaSelectedCount, 0, sizeof(int), stream));
1279 
1280   const dim3 pathGrid((numComponents + kBlockSize - 1) / kBlockSize);
1281   traceSelectedPathsKernel<<<pathGrid, kBlockSize, 0, stream>>>(
1282       dpState.bestChild, cudaSelectedRoots, cudaSelectedTrackLabels,
1283       cudaSelectedNodes, cudaSelectedCount, numNodes, numComponents);
1284   ACTS_CUDA_CHECK(cudaGetLastError());
1285 
1286   int selectedCount = 0;
1287   ACTS_CUDA_CHECK(cudaMemcpyAsync(&selectedCount, cudaSelectedCount,
1288                                   sizeof(int), cudaMemcpyDeviceToHost, stream));
1289   ACTS_CUDA_CHECK(cudaStreamSynchronize(stream));
1290 
1291   if (selectedCount > 0 && cudaActiveNodes != nullptr) {
1292     const dim3 selectedGrid((selectedCount + kBlockSize - 1) / kBlockSize);
1293     deactivateSelectedNodesKernel<<<selectedGrid, kBlockSize, 0, stream>>>(
1294         selectedCount, cudaSelectedNodes, cudaActiveNodes);
1295     ACTS_CUDA_CHECK(cudaGetLastError());
1296   }
1297 
1298   std::vector<int> labels(selectedCount);
1299   std::vector<int> nodes(selectedCount);
1300   if (selectedCount > 0) {
1301     ACTS_CUDA_CHECK(cudaMemcpyAsync(labels.data(), cudaSelectedTrackLabels,
1302                                     selectedCount * sizeof(int),
1303                                     cudaMemcpyDeviceToHost, stream));
1304     ACTS_CUDA_CHECK(cudaMemcpyAsync(nodes.data(), cudaSelectedNodes,
1305                                     selectedCount * sizeof(int),
1306                                     cudaMemcpyDeviceToHost, stream));
1307     ACTS_CUDA_CHECK(cudaStreamSynchronize(stream));
1308   }
1309 
1310   ACTS_CUDA_CHECK(cudaFreeAsync(cudaSelectedRoots, stream));
1311   ACTS_CUDA_CHECK(cudaFreeAsync(cudaComponentBestScore, stream));
1312   ACTS_CUDA_CHECK(cudaFreeAsync(cudaSelectedTrackLabels, stream));
1313   ACTS_CUDA_CHECK(cudaFreeAsync(cudaSelectedNodes, stream));
1314   ACTS_CUDA_CHECK(cudaFreeAsync(cudaSelectedCount, stream));
1315 
1316   return {std::move(labels), std::move(nodes)};
1317 }
1318 
1319 void appendSmallResidualDWalkTracks(
1320     const DeviceCompactEdges &edges,
1321     const std::vector<int> &complexSpacePointIds, std::size_t minCandidateSize,
1322     const std::string &pathMetric, cudaStream_t stream,
1323     std::vector<std::vector<int>> &trackCandidates) {
1324   if (edges.numEdges == 0) {
1325     return;
1326   }
1327 
1328   std::vector<int> src(edges.numEdges);
1329   std::vector<int> dst(edges.numEdges);
1330   std::vector<float> score(edges.numEdges);
1331   ACTS_CUDA_CHECK(cudaMemcpyAsync(src.data(), edges.src,
1332                                   edges.numEdges * sizeof(int),
1333                                   cudaMemcpyDeviceToHost, stream));
1334   ACTS_CUDA_CHECK(cudaMemcpyAsync(dst.data(), edges.dst,
1335                                   edges.numEdges * sizeof(int),
1336                                   cudaMemcpyDeviceToHost, stream));
1337   ACTS_CUDA_CHECK(cudaMemcpyAsync(score.data(), edges.score,
1338                                   edges.numEdges * sizeof(float),
1339                                   cudaMemcpyDeviceToHost, stream));
1340   ACTS_CUDA_CHECK(cudaStreamSynchronize(stream));
1341 
1342   std::vector<unsigned char> activeEdge(edges.numEdges, 1);
1343   const float rootScoreCut = minRootScore(pathMetric);
1344   const bool useLengthMetric = pathMetric == "length";
1345 
1346   while (true) {
1347     std::vector<std::size_t> currentEdges;
1348     currentEdges.reserve(edges.numEdges);
1349     std::vector<unsigned char> activeNode(edges.numNodes, 0);
1350     for (std::size_t edge = 0; edge < edges.numEdges; ++edge) {
1351       if (activeEdge.at(edge) == 0) {
1352         continue;
1353       }
1354       currentEdges.push_back(edge);
1355       activeNode.at(src.at(edge)) = 1;
1356       activeNode.at(dst.at(edge)) = 1;
1357     }
1358     if (currentEdges.empty()) {
1359       break;
1360     }
1361 
1362     DisjointSet components(edges.numNodes);
1363     std::vector<int> inDegree(edges.numNodes, 0);
1364     std::vector<std::vector<std::pair<int, float>>> adjacency(edges.numNodes);
1365     for (std::size_t edge : currentEdges) {
1366       components.unite(src.at(edge), dst.at(edge));
1367       ++inDegree.at(dst.at(edge));
1368       adjacency.at(src.at(edge))
1369           .push_back({dst.at(edge), useLengthMetric ? 1.0F : score.at(edge)});
1370     }
1371 
1372     std::unordered_map<int, int> componentMap;
1373     std::vector<int> componentLabel(edges.numNodes, -1);
1374     int numComponents = 0;
1375     for (std::size_t node = 0; node < edges.numNodes; ++node) {
1376       if (activeNode.at(node) == 0) {
1377         continue;
1378       }
1379       auto [it, inserted] = componentMap.emplace(
1380           components.find(static_cast<int>(node)), numComponents);
1381       if (inserted) {
1382         ++numComponents;
1383       }
1384       componentLabel.at(node) = it->second;
1385     }
1386 
1387     std::vector<float> bestScore(edges.numNodes, 0.0F);
1388     std::vector<int> bestChild(edges.numNodes, -1);
1389     std::vector<unsigned char> visitState(edges.numNodes, 0);
1390     std::function<float(int)> solve = [&](int node) -> float {
1391       if (visitState.at(node) == 2) {
1392         return bestScore.at(node);
1393       }
1394       if (visitState.at(node) == 1) {
1395         return 0.0F;
1396       }
1397       visitState.at(node) = 1;
1398       float bestValue = 0.0F;
1399       int bestNext = -1;
1400       for (auto [child, weight] : adjacency.at(node)) {
1401         if (activeNode.at(child) == 0) {
1402           continue;
1403         }
1404         float candidate = weight + solve(child);
1405         if (candidate > bestValue) {
1406           bestValue = candidate;
1407           bestNext = child;
1408         }
1409       }
1410       bestScore.at(node) = bestValue;
1411       bestChild.at(node) = bestNext;
1412       visitState.at(node) = 2;
1413       return bestValue;
1414     };
1415 
1416     std::vector<float> componentBest(numComponents,
1417                                      -std::numeric_limits<float>::infinity());
1418     std::vector<int> componentRoot(numComponents, -1);
1419     for (std::size_t node = 0; node < edges.numNodes; ++node) {
1420       if (activeNode.at(node) == 0 || inDegree.at(node) != 0) {
1421         continue;
1422       }
1423       float value = solve(static_cast<int>(node));
1424       int component = componentLabel.at(node);
1425       if (component >= 0 && value > rootScoreCut &&
1426           value > componentBest.at(component)) {
1427         componentBest.at(component) = value;
1428         componentRoot.at(component) = static_cast<int>(node);
1429       }
1430     }
1431 
1432     std::vector<unsigned char> selectedNode(edges.numNodes, 0);
1433     bool selectedAny = false;
1434     for (int root : componentRoot) {
1435       if (root < 0) {
1436         continue;
1437       }
1438       std::vector<int> path;
1439       int node = root;
1440       std::size_t guard = 0;
1441       while (node >= 0 && guard < edges.numNodes &&
1442              selectedNode.at(node) == 0) {
1443         path.push_back(node);
1444         selectedNode.at(node) = 1;
1445         selectedAny = true;
1446         node = bestChild.at(node);
1447         ++guard;
1448       }
1449       if (path.size() < minCandidateSize) {
1450         continue;
1451       }
1452       std::vector<int> track;
1453       track.reserve(path.size());
1454       for (int pathNode : path) {
1455         track.push_back(complexSpacePointIds.at(pathNode));
1456       }
1457       trackCandidates.push_back(std::move(track));
1458     }
1459 
1460     if (!selectedAny) {
1461       break;
1462     }
1463     for (std::size_t edge = 0; edge < edges.numEdges; ++edge) {
1464       if (activeEdge.at(edge) != 0 && (selectedNode.at(src.at(edge)) != 0 ||
1465                                        selectedNode.at(dst.at(edge)) != 0)) {
1466         activeEdge.at(edge) = 0;
1467       }
1468     }
1469   }
1470 }
1471 
1472 }  // namespace
1473 
1474 namespace ActsPlugins {
1475 
1476 std::vector<std::vector<int>> DWalkTrackBuilding::operator()(
1477     PipelineTensors tensors, std::vector<int> &spacePointIds,
1478     const ExecutionContext &execContext) {
1479   ACTS_VERBOSE("Start CUDA D-WALK track building");
1480 
1481   if (!tensors.edgeScores.has_value()) {
1482     throw std::runtime_error("DWalkTrackBuilding expects edge scores");
1483   }
1484   if (!(tensors.edgeIndex.device().isCuda() &&
1485         tensors.edgeScores->device().isCuda() &&
1486         tensors.nodeFeatures.device().isCuda())) {
1487     throw std::runtime_error(
1488         "DWalkTrackBuilding expects tensors to be on CUDA");
1489   }
1490   if (m_cfg.pathMetric != "score_weighted_length" &&
1491       m_cfg.pathMetric != "length") {
1492     throw std::invalid_argument(
1493         "DWalkTrackBuilding pathMetric must be 'score_weighted_length' or "
1494         "'length'");
1495   }
1496 
1497   assert(tensors.edgeIndex.shape().at(0) == 2);
1498   assert(tensors.edgeIndex.shape().at(1) == tensors.edgeScores->shape().at(0));
1499 
1500   const auto numNodes = tensors.nodeFeatures.shape().at(0);
1501   const auto numFeatures = tensors.nodeFeatures.shape().at(1);
1502   const auto numEdges = tensors.edgeIndex.shape().at(1);
1503   if (m_cfg.radialFeatureIndex >= numFeatures) {
1504     throw std::out_of_range("DWalkTrackBuilding radialFeatureIndex is invalid");
1505   }
1506   if (numNodes > spacePointIds.size()) {
1507     throw std::runtime_error(
1508         "DWalkTrackBuilding received more graph nodes than space point IDs");
1509   }
1510   if (numEdges == 0 || numNodes < m_cfg.minCandidateSize) {
1511     return {};
1512   }
1513 
1514   auto stream = execContext.stream.value();
1515   std::vector<int> initialLabels;
1516   DeviceOrientedEdges deviceGraph;
1517   int *cudaInitialLabels{};
1518   int initialNumComponents = 0;
1519   auto edges = createOrientedEdgesCuda(
1520       tensors.edgeIndex, *tensors.edgeScores, tensors.nodeFeatures,
1521       m_cfg.radialFeatureIndex, stream, deviceGraph, &cudaInitialLabels,
1522       &initialNumComponents, &initialLabels);
1523   if (edges.empty()) {
1524     deviceGraph.reset();
1525     ACTS_CUDA_CHECK(cudaFreeAsync(cudaInitialLabels, stream));
1526     ACTS_CUDA_CHECK(cudaStreamSynchronize(stream));
1527     return {};
1528   }
1529 
1530   unsigned char *cudaSimpleNodeMask{};
1531   unsigned char *cudaComplexNodeMask{};
1532   classifyInitialComponentsCuda(
1533       deviceGraph, cudaInitialLabels, initialNumComponents,
1534       static_cast<int>(m_cfg.minCandidateSize), &cudaSimpleNodeMask,
1535       &cudaComplexNodeMask, stream);
1536 
1537   std::vector<unsigned char> simpleNodeMask(numNodes, 0);
1538   std::vector<unsigned char> complexNodeMask(numNodes, 0);
1539   ACTS_CUDA_CHECK(cudaMemcpyAsync(simpleNodeMask.data(), cudaSimpleNodeMask,
1540                                   numNodes * sizeof(unsigned char),
1541                                   cudaMemcpyDeviceToHost, stream));
1542   ACTS_CUDA_CHECK(cudaMemcpyAsync(complexNodeMask.data(), cudaComplexNodeMask,
1543                                   numNodes * sizeof(unsigned char),
1544                                   cudaMemcpyDeviceToHost, stream));
1545   ACTS_CUDA_CHECK(cudaStreamSynchronize(stream));
1546 
1547   std::vector<int> inDegree(numNodes, 0);
1548   std::vector<int> outDegree(numNodes, 0);
1549   std::vector<int> nextNode(numNodes, -1);
1550   for (const auto &edge : edges) {
1551     ++outDegree.at(edge.src);
1552     ++inDegree.at(edge.dst);
1553     nextNode.at(edge.src) = edge.dst;
1554   }
1555 
1556   std::unordered_map<int, std::vector<int>> componentNodes;
1557   for (std::size_t node = 0; node < numNodes; ++node) {
1558     if (simpleNodeMask.at(node) != 0) {
1559       const int label = initialLabels.at(node);
1560       componentNodes[label].push_back(static_cast<int>(node));
1561     }
1562   }
1563   std::vector<std::vector<int>> trackCandidates;
1564   for (const auto &[component, nodes] : componentNodes) {
1565     auto orderedNodes = orderedSimpleComponentNodes(
1566         nodes, component, initialLabels, inDegree, nextNode);
1567     std::vector<int> track;
1568     track.reserve(orderedNodes.size());
1569     for (int node : orderedNodes) {
1570       track.push_back(spacePointIds.at(node));
1571     }
1572     trackCandidates.push_back(std::move(track));
1573   }
1574 
1575   std::vector<int> originalToComplex(numNodes, -1);
1576   std::vector<int> complexSpacePointIds;
1577   for (std::size_t node = 0; node < numNodes; ++node) {
1578     if (complexNodeMask.at(node) != 0) {
1579       originalToComplex.at(node) =
1580           static_cast<int>(complexSpacePointIds.size());
1581       complexSpacePointIds.push_back(spacePointIds.at(node));
1582     }
1583   }
1584   const std::size_t numComplexNodes = complexSpacePointIds.size();
1585   ACTS_DEBUG("CUDA D-WALK complex nodes: " << numComplexNodes);
1586 
1587   if (numComplexNodes == 0) {
1588     deviceGraph.reset();
1589     ACTS_CUDA_CHECK(cudaFreeAsync(cudaInitialLabels, stream));
1590     ACTS_CUDA_CHECK(cudaFreeAsync(cudaSimpleNodeMask, stream));
1591     ACTS_CUDA_CHECK(cudaFreeAsync(cudaComplexNodeMask, stream));
1592     ACTS_CUDA_CHECK(cudaStreamSynchronize(stream));
1593     ACTS_DEBUG("CUDA D-WALK found " << trackCandidates.size()
1594                                     << " track candidates");
1595     return trackCandidates;
1596   }
1597 
1598   int *cudaOriginalToComplex{};
1599   ACTS_CUDA_CHECK(
1600       cudaMallocAsync(&cudaOriginalToComplex, numNodes * sizeof(int), stream));
1601   ACTS_CUDA_CHECK(
1602       cudaMemcpyAsync(cudaOriginalToComplex, originalToComplex.data(),
1603                       numNodes * sizeof(int), cudaMemcpyHostToDevice, stream));
1604 
1605   auto complexEdges = maxAddCompactDeviceEdgesCuda(
1606       deviceGraph, cudaComplexNodeMask, m_cfg.thMin, m_cfg.thAdd, stream);
1607   deviceGraph.reset();
1608   if (complexEdges.numEdges != 0) {
1609     const dim3 remapGrid((complexEdges.numEdges + kBlockSize - 1) / kBlockSize);
1610     remapEdgesKernel<<<remapGrid, kBlockSize, 0, stream>>>(
1611         complexEdges.numEdges, complexEdges.src, complexEdges.dst,
1612         cudaOriginalToComplex);
1613     ACTS_CUDA_CHECK(cudaGetLastError());
1614     complexEdges.numNodes = numComplexNodes;
1615   }
1616   ACTS_CUDA_CHECK(cudaFreeAsync(cudaInitialLabels, stream));
1617   ACTS_CUDA_CHECK(cudaFreeAsync(cudaSimpleNodeMask, stream));
1618   ACTS_CUDA_CHECK(cudaFreeAsync(cudaComplexNodeMask, stream));
1619   ACTS_CUDA_CHECK(cudaFreeAsync(cudaOriginalToComplex, stream));
1620   ACTS_DEBUG("CUDA D-WALK initial complex edges: " << complexEdges.numEdges);
1621 
1622   unsigned char *cudaComplexActiveNodes{};
1623   ACTS_CUDA_CHECK(cudaMallocAsync(&cudaComplexActiveNodes,
1624                                   numComplexNodes * sizeof(unsigned char),
1625                                   stream));
1626   ACTS_CUDA_CHECK(cudaMemsetAsync(cudaComplexActiveNodes, 1,
1627                                   numComplexNodes * sizeof(unsigned char),
1628                                   stream));
1629 
1630   int iteration = 0;
1631   while (complexEdges.numEdges != 0) {
1632     if (complexEdges.numEdges <= kSmallResidualEdgeThreshold) {
1633       appendSmallResidualDWalkTracks(complexEdges, complexSpacePointIds,
1634                                      m_cfg.minCandidateSize, m_cfg.pathMetric,
1635                                      stream, trackCandidates);
1636       break;
1637     }
1638 
1639     ++iteration;
1640     int *cudaCurrentLabels{};
1641     ACTS_CUDA_CHECK(cudaMallocAsync(&cudaCurrentLabels,
1642                                     numComplexNodes * sizeof(int), stream));
1643     int currentNumComponents = ActsPlugins::detail::connectedComponentsCuda(
1644         complexEdges.numEdges, complexEdges.src, complexEdges.dst,
1645         numComplexNodes, cudaCurrentLabels, stream, false);
1646     ACTS_DEBUG("CUDA D-WALK iteration "
1647                << iteration << ": components=" << currentNumComponents
1648                << ", edges=" << complexEdges.numEdges);
1649     if (currentNumComponents == 0) {
1650       ACTS_CUDA_CHECK(cudaFreeAsync(cudaCurrentLabels, stream));
1651       break;
1652     }
1653 
1654     auto csrGraph = buildCsrGraphCuda(complexEdges, m_cfg.pathMetric, stream);
1655     auto dpState = runDpOnCsrCuda(csrGraph, cudaComplexActiveNodes, stream);
1656     csrGraph.reset();
1657 
1658     auto [selectedTrackLabels, selectedNodes] = selectAndTracePathsCuda(
1659         dpState, cudaCurrentLabels, currentNumComponents,
1660         minRootScore(m_cfg.pathMetric), cudaComplexActiveNodes, stream);
1661     dpState.reset();
1662     ACTS_CUDA_CHECK(cudaFreeAsync(cudaCurrentLabels, stream));
1663     if (selectedNodes.empty()) {
1664       break;
1665     }
1666 
1667     std::vector<std::vector<int>> paths(currentNumComponents);
1668     for (std::size_t selected = 0; selected < selectedNodes.size();
1669          ++selected) {
1670       int label = selectedTrackLabels.at(selected);
1671       int node = selectedNodes.at(selected);
1672       if (label < 0 || label >= static_cast<int>(paths.size()) || node < 0 ||
1673           node >= static_cast<int>(numComplexNodes)) {
1674         continue;
1675       }
1676       paths.at(label).push_back(node);
1677     }
1678 
1679     for (const auto &path : paths) {
1680       if (path.size() < m_cfg.minCandidateSize) {
1681         continue;
1682       }
1683       std::vector<int> track;
1684       track.reserve(path.size());
1685       for (int node : path) {
1686         track.push_back(complexSpacePointIds.at(node));
1687       }
1688       trackCandidates.push_back(std::move(track));
1689     }
1690 
1691     auto updatedEdges =
1692         compactActiveEdgesCuda(complexEdges, cudaComplexActiveNodes, stream);
1693     complexEdges = std::move(updatedEdges);
1694   }
1695 
1696   complexEdges.reset();
1697   ACTS_CUDA_CHECK(cudaFreeAsync(cudaComplexActiveNodes, stream));
1698   ACTS_CUDA_CHECK(cudaStreamSynchronize(stream));
1699 
1700   ACTS_DEBUG("CUDA D-WALK found " << trackCandidates.size()
1701                                   << " track candidates");
1702   return trackCandidates;
1703 }
1704 
1705 }  // namespace ActsPlugins