Back to home page

EIC code displayed by LXR

 
 

    


File indexing completed on 2026-07-26 08:22:10

0001 /** TRACCC library, part of the ACTS project (R&D line)
0002  *
0003  * (c) 2025-2026 CERN for the benefit of the ACTS project
0004  *
0005  * Mozilla Public License Version 2.0
0006  */
0007 
0008 // Project include(s).
0009 #include "../utils/cuda_error_handling.hpp"
0010 #include "../utils/utils.hpp"
0011 #include "./kernels/add_block_offset.cuh"
0012 #include "./kernels/block_inclusive_scan.cuh"
0013 #include "./kernels/count_shared_measurements.cuh"
0014 #include "./kernels/fill_inverted_ids.cuh"
0015 #include "./kernels/fill_track_candidates.cuh"
0016 #include "./kernels/fill_tracks_per_measurement.cuh"
0017 #include "./kernels/fill_unique_meas_id_map.cuh"
0018 #include "./kernels/fill_vectors.cuh"
0019 #include "./kernels/rearrange_tracks.cuh"
0020 #include "./kernels/remove_tracks.cuh"
0021 #include "./kernels/scan_block_offsets.cuh"
0022 #include "./kernels/sort_tracks_per_measurement.cuh"
0023 #include "./kernels/sort_updated_tracks.cuh"
0024 #include "./kernels/update_status.cuh"
0025 #include "traccc/cuda/ambiguity_resolution/greedy_ambiguity_resolution_algorithm.hpp"
0026 #include "traccc/definitions/math.hpp"
0027 
0028 // Thrust include(s).
0029 #include <thrust/execution_policy.h>
0030 // Suppress warning (error at -Werror) from CUB/Thrust.
0031 #pragma GCC diagnostic push
0032 #pragma GCC diagnostic ignored "-Wmaybe-uninitialized"
0033 #include <thrust/extrema.h>
0034 #pragma GCC diagnostic pop
0035 #include <thrust/fill.h>
0036 #include <thrust/functional.h>
0037 #include <thrust/iterator/constant_iterator.h>
0038 #include <thrust/reduce.h>
0039 #include <thrust/sort.h>
0040 #include <thrust/transform.h>
0041 #include <thrust/unique.h>
0042 namespace traccc::cuda {
0043 
0044 struct identity_op {
0045   template <typename T>
0046   TRACCC_HOST_DEVICE T operator()(T i) const {
0047     return i;
0048   }
0049 };
0050 
0051 // Device operator to calculate relative number of shared measurements
0052 struct devide_op {
0053   TRACCC_HOST_DEVICE
0054   traccc::scalar operator()(unsigned int a, unsigned int b) const {
0055     return math::div_ieee754(static_cast<traccc::scalar>(a),
0056                              static_cast<traccc::scalar>(b));
0057   }
0058 };
0059 
0060 // Track comparator to sort the track ids
0061 struct track_comparator {
0062   const traccc::scalar* rel_shared;
0063   const traccc::scalar* pvals;
0064 
0065   TRACCC_HOST_DEVICE track_comparator(const traccc::scalar* rel_shared_,
0066                                       const traccc::scalar* pvals_)
0067       : rel_shared(rel_shared_), pvals(pvals_) {}
0068 
0069   TRACCC_HOST_DEVICE bool operator()(unsigned int a, unsigned int b) const {
0070     if (rel_shared[a] != rel_shared[b]) {
0071       return rel_shared[a] < rel_shared[b];
0072     }
0073     return pvals[a] > pvals[b];
0074   }
0075 };
0076 
0077 greedy_ambiguity_resolution_algorithm::greedy_ambiguity_resolution_algorithm(
0078     const config_type& cfg, const traccc::memory_resource& mr,
0079     const vecmem::copy& copy, const stream_wrapper& str,
0080     std::unique_ptr<const Logger> logger)
0081     : messaging(std::move(logger)),
0082       m_config(cfg),
0083       m_mr(mr),
0084       m_copy(copy),
0085       m_stream(str),
0086       m_warp_size(details::get_warp_size(str.device())) {}
0087 
0088 greedy_ambiguity_resolution_algorithm::output_type
0089 greedy_ambiguity_resolution_algorithm::operator()(
0090     const edm::track_container<default_algebra>::const_view& tracks_view)
0091     const {
0092   const edm::measurement_collection::const_device measurements(
0093       tracks_view.measurements);
0094 
0095   auto n_meas_total = m_copy.get().get_size(tracks_view.measurements);
0096 
0097   // Make sure that max_measurement_id = number_of_measurement -1
0098   // @TODO: More robust way is to assert that measurement id ranges from 0, 1,
0099   // ..., number_of_measurement - 1
0100   [[maybe_unused]] auto max_meas_it = thrust::max_element(
0101       thrust::device, measurements.identifier().begin(),
0102       // We have to use this ugly form here, because if the measurement
0103       // collection is resizable (which it often is), the end() function
0104       // cannot be used in host code.
0105       measurements.identifier().begin() + n_meas_total);
0106 
0107   unsigned int max_meas_id;
0108   cudaMemcpy(&max_meas_id, thrust::raw_pointer_cast(&(*max_meas_it)),
0109              sizeof(unsigned int), cudaMemcpyDeviceToHost);
0110 
0111   if (max_meas_id != n_meas_total - 1) {
0112     throw std::runtime_error(
0113         "max measurement id should be equal to (the number of measurements "
0114         "- 1)");
0115   }
0116 
0117   // Get a convenience variable for the stream that we'll be using.
0118   cudaStream_t stream = details::get_stream(m_stream);
0119 
0120   // The Thrust policy to use.
0121   auto thrust_policy =
0122       thrust::cuda::par_nosync(std::pmr::polymorphic_allocator(&(m_mr.main)))
0123           .on(stream);
0124 
0125   const unsigned int n_tracks = tracks_view.tracks.capacity();
0126 
0127   if (n_tracks == 0) {
0128     return {};
0129   }
0130 
0131   // Make sure that max_shared_meas is largen than zero
0132   assert(m_config.max_shared_meas > 0u);
0133 
0134   // Status (1 = Accept, 0 = Reject) vector to count the number of acceptable
0135   // tracks based on the number of candidates (measurements)
0136   vecmem::data::vector_buffer<int> status_buffer{n_tracks, m_mr.main};
0137 
0138   vecmem::device_vector<int> status_device(status_buffer);
0139   thrust::fill(thrust_policy, status_device.begin(), status_device.end(), 1);
0140 
0141   // Get the sizes of the measurement index vector in each track
0142   const std::vector<unsigned int> candidate_sizes =
0143       m_copy.get().get_sizes(tracks_view.tracks);
0144 
0145   // Declare the buffer for meas_ids which is a jagged vector
0146   // Each sub-vector of meas_ids represent measurement IDs of each track
0147   vecmem::data::jagged_vector_buffer<measurement_id_type> meas_ids_buffer{
0148       candidate_sizes, m_mr.main, m_mr.host,
0149       vecmem::data::buffer_type::resizable};
0150   m_copy.get().setup(meas_ids_buffer)->ignore();
0151 
0152   // The sum of the number of candidates (measurements) of all tracks
0153   const unsigned int n_cands_total =
0154       std::accumulate(candidate_sizes.begin(), candidate_sizes.end(), 0u);
0155 
0156   // Declare flat_meas_ids which is just a flattening version of meas_ids with
0157   // a single vector container. It is used to count the number of unique
0158   // measurements
0159   vecmem::data::vector_buffer<measurement_id_type> flat_meas_ids_buffer{
0160       n_cands_total, m_mr.main, vecmem::data::buffer_type::resizable};
0161   m_copy.get().setup(flat_meas_ids_buffer)->ignore();
0162   vecmem::data::vector_buffer<traccc::scalar> pvals_buffer{n_tracks, m_mr.main};
0163   vecmem::data::vector_buffer<unsigned int> n_meas_buffer{n_tracks, m_mr.main};
0164   thrust::fill(thrust_policy, n_meas_buffer.ptr(),
0165                n_meas_buffer.ptr() + n_tracks, 0);
0166 
0167   {
0168     const unsigned int nThreads = m_warp_size * 2;
0169     const unsigned int nBlocks = (n_tracks + nThreads - 1) / nThreads;
0170 
0171     // Fill the vectors
0172     kernels::fill_vectors<<<nBlocks, nThreads, 0, stream>>>(
0173         m_config,
0174         device::fill_vectors_payload{.tracks_view = tracks_view,
0175                                      .meas_ids_view = meas_ids_buffer,
0176                                      .flat_meas_ids_view = flat_meas_ids_buffer,
0177                                      .pvals_view = pvals_buffer,
0178                                      .n_meas_view = n_meas_buffer,
0179                                      .status_view = status_buffer});
0180     TRACCC_CUDA_ERROR_CHECK(cudaGetLastError());
0181 
0182     m_stream.synchronize();
0183   }
0184 
0185   // Count the number of pre-accepted tracks
0186   unsigned int n_accepted = static_cast<unsigned int>(thrust::count(
0187       thrust_policy, status_buffer.ptr(), status_buffer.ptr() + n_tracks, 1));
0188 
0189   vecmem::unique_alloc_ptr<unsigned int> n_accepted_device =
0190       vecmem::make_unique_alloc<unsigned int>(m_mr.main);
0191   TRACCC_CUDA_ERROR_CHECK(cudaMemcpyAsync(n_accepted_device.get(), &n_accepted,
0192                                           sizeof(unsigned int),
0193                                           cudaMemcpyHostToDevice, stream));
0194 
0195   m_stream.synchronize();
0196 
0197   if (n_accepted == 0) {
0198     return {};
0199   }
0200 
0201   // Indices of pre-accepted tracks
0202   vecmem::data::vector_buffer<unsigned int> pre_accepted_ids_buffer{n_accepted,
0203                                                                     m_mr.main};
0204 
0205   m_copy.get().setup(pre_accepted_ids_buffer)->ignore();
0206 
0207   // Find the indices of pre-accepted tracks by checking if status is 1
0208   auto cit_begin = thrust::counting_iterator<int>(0);
0209   auto cit_end = cit_begin + n_tracks;
0210   thrust::copy_if(thrust_policy, cit_begin, cit_end, status_buffer.ptr(),
0211                   pre_accepted_ids_buffer.ptr(), identity_op{});
0212 
0213   // Sort the flat measurement id vector, which is required to count the
0214   // number of unique measurements
0215   thrust::sort(thrust_policy, flat_meas_ids_buffer.ptr(),
0216                flat_meas_ids_buffer.ptr() + n_cands_total);
0217 
0218   // Count the number of unique measurements
0219   const unsigned int meas_count =
0220       static_cast<unsigned int>(thrust::unique_count(
0221           thrust_policy, flat_meas_ids_buffer.ptr(),
0222           flat_meas_ids_buffer.ptr() + n_cands_total, thrust::equal_to<int>()));
0223 
0224   // Unique measurement ids
0225   vecmem::data::vector_buffer<measurement_id_type> unique_meas_buffer{
0226       meas_count, m_mr.main};
0227 
0228   // Counts of unique measurement id in flat id vector.
0229   // This information is used to know the number of tracks associated with a
0230   // measurement ID.
0231   vecmem::data::vector_buffer<std::size_t> unique_meas_counts_buffer{meas_count,
0232                                                                      m_mr.main};
0233   m_copy.get().setup(unique_meas_counts_buffer)->ignore();
0234 
0235   // Counting can be done using reduce_by_key and constant iterator
0236   thrust::reduce_by_key(thrust_policy, flat_meas_ids_buffer.ptr(),
0237                         flat_meas_ids_buffer.ptr() + n_cands_total,
0238                         thrust::make_constant_iterator(1),
0239                         unique_meas_buffer.ptr(),
0240                         unique_meas_counts_buffer.ptr());
0241 
0242   // Sort unique meas ids
0243   thrust::sort_by_key(thrust_policy, unique_meas_buffer.ptr(),
0244                       unique_meas_buffer.ptr() + meas_count,
0245                       unique_meas_counts_buffer.ptr());
0246 
0247   // Unique measurement ids
0248   vecmem::data::vector_buffer<measurement_id_type> meas_id_to_unique_id_buffer{
0249       max_meas_id + 1, m_mr.main};
0250 
0251   // Make meas_id to unique_meas_id vector
0252   {
0253     const unsigned int nThreads = m_warp_size * 2;
0254     const unsigned int nBlocks = (meas_count + nThreads - 1) / nThreads;
0255 
0256     kernels::fill_unique_meas_id_map<<<nBlocks, nThreads, 0, stream>>>(
0257         device::fill_unique_meas_id_map_payload{
0258             .unique_meas_view = unique_meas_buffer,
0259             .meas_id_to_unique_id_view = meas_id_to_unique_id_buffer});
0260     TRACCC_CUDA_ERROR_CHECK(cudaGetLastError());
0261 
0262     m_stream.synchronize();
0263   }
0264 
0265   // Retreive the counting vector to host for the size allocation of
0266   // tracks_per_measurement
0267   std::vector<std::size_t> unique_meas_counts;
0268   m_copy
0269       .get()(unique_meas_counts_buffer, unique_meas_counts,
0270              vecmem::copy::type::device_to_host)
0271       ->wait();
0272 
0273   // Make the tracks_per_measurement vector
0274   // Each sub vector contains track ids associated with the unique measurement
0275   vecmem::data::jagged_vector_buffer<unsigned int>
0276       tracks_per_measurement_buffer(unique_meas_counts, m_mr.main, m_mr.host,
0277                                     vecmem::data::buffer_type::resizable);
0278   m_copy.get().setup(tracks_per_measurement_buffer)->ignore();
0279 
0280   // Make the track_status_per_measurement vector
0281   // Each sub vector contains whether the track ids is still associated with
0282   // the unique measurements For example, the value turns into 0 (false) if
0283   // the track is rejected during the ambiguity solver
0284   vecmem::data::jagged_vector_buffer<int> track_status_per_measurement_buffer(
0285       unique_meas_counts, m_mr.main, m_mr.host,
0286       vecmem::data::buffer_type::resizable);
0287 
0288   m_copy.get().setup(track_status_per_measurement_buffer)->ignore();
0289 
0290   // Make the number of accetped_tracks_per_measurement vector
0291   // Each element represents the number of associated tracks with the unique
0292   // measurement (the number of track_status whose value is 1 (true))
0293   vecmem::data::vector_buffer<unsigned int>
0294       n_accepted_tracks_per_measurement_buffer(meas_count, m_mr.main);
0295   thrust::fill(thrust_policy, n_accepted_tracks_per_measurement_buffer.ptr(),
0296                n_accepted_tracks_per_measurement_buffer.ptr() + meas_count, 0);
0297 
0298   // Fill tracks_per_measurement, track_status_per_measurement and
0299   // n_accepted_tracks_per_measurement vectors
0300   {
0301     const unsigned int nThreads = m_warp_size * 2;
0302     const unsigned int nBlocks = (n_accepted + nThreads - 1) / nThreads;
0303 
0304     kernels::fill_tracks_per_measurement<<<nBlocks, nThreads, 0, stream>>>(
0305         device::fill_tracks_per_measurement_payload{
0306             .accepted_ids_view = pre_accepted_ids_buffer,
0307             .meas_ids_view = meas_ids_buffer,
0308             .meas_id_to_unique_id_view = meas_id_to_unique_id_buffer,
0309             .tracks_per_measurement_view = tracks_per_measurement_buffer,
0310             .track_status_per_measurement_view =
0311                 track_status_per_measurement_buffer,
0312             .n_accepted_tracks_per_measurement_view =
0313                 n_accepted_tracks_per_measurement_buffer});
0314     TRACCC_CUDA_ERROR_CHECK(cudaGetLastError());
0315 
0316     m_stream.synchronize();
0317   }
0318 
0319   // Sort tracks per measurement vector
0320   // @TODO: For the case where the measurement is shared by more than 1024
0321   // tracks, the tracks need to be sorted again using thrust::sort
0322   {
0323     const unsigned int nThreads = 1024;
0324     const unsigned int nBlocks = meas_count;
0325 
0326     kernels::sort_tracks_per_measurement<<<nBlocks, nThreads, 0, stream>>>(
0327         device::sort_tracks_per_measurement_payload{
0328             .tracks_per_measurement_view = tracks_per_measurement_buffer,
0329         });
0330     TRACCC_CUDA_ERROR_CHECK(cudaGetLastError());
0331 
0332     m_stream.synchronize();
0333   }
0334 
0335   // Make vector buffer for the number of shared measurements for each track
0336   vecmem::data::vector_buffer<unsigned int> n_shared_buffer{n_tracks,
0337                                                             m_mr.main};
0338   thrust::fill(thrust_policy, n_shared_buffer.ptr(),
0339                n_shared_buffer.ptr() + n_tracks, 0);
0340   m_copy.get().setup(n_shared_buffer)->ignore();
0341 
0342   // Count the number of shared measurements
0343   {
0344     const unsigned int nThreads = m_warp_size * 2;
0345     const unsigned int nBlocks = (n_accepted + nThreads - 1) / nThreads;
0346 
0347     kernels::count_shared_measurements<<<nBlocks, nThreads, 0, stream>>>(
0348         device::count_shared_measurements_payload{
0349             .accepted_ids_view = pre_accepted_ids_buffer,
0350             .meas_ids_view = meas_ids_buffer,
0351             .meas_id_to_unique_id_view = meas_id_to_unique_id_buffer,
0352             .n_accepted_tracks_per_measurement_view =
0353                 n_accepted_tracks_per_measurement_buffer,
0354             .n_shared_view = n_shared_buffer});
0355     TRACCC_CUDA_ERROR_CHECK(cudaGetLastError());
0356 
0357     m_stream.synchronize();
0358   }
0359 
0360   // Make relative number of shared measurements vector
0361   // The relative number of shared measurement is defined as the number of
0362   // shared measurement divided the number of measurements of the track
0363   vecmem::data::vector_buffer<traccc::scalar> rel_shared_buffer{n_tracks,
0364                                                                 m_mr.main};
0365 
0366   // Fill the relative shared number of measurements vector
0367   thrust::transform(thrust_policy, n_shared_buffer.ptr(),
0368                     n_shared_buffer.ptr() + n_tracks, n_meas_buffer.ptr(),
0369                     rel_shared_buffer.ptr(), devide_op{});
0370 
0371   // Make a buffer for track ids sorted based on the relative number of shared
0372   // measurements and pvalues
0373   vecmem::data::vector_buffer<unsigned int> sorted_ids_buffer{n_accepted,
0374                                                               m_mr.main};
0375   m_copy.get().setup(sorted_ids_buffer)->ignore();
0376 
0377   // Make a temporary buffer for sorted track ids
0378   vecmem::data::vector_buffer<unsigned int> temp_sorted_ids_buffer{n_accepted,
0379                                                                    m_mr.main};
0380   m_copy.get().setup(temp_sorted_ids_buffer)->ignore();
0381 
0382   // track id to the index of sorted ids
0383   vecmem::data::vector_buffer<unsigned int> inverted_ids_buffer{n_tracks,
0384                                                                 m_mr.main};
0385   m_copy.get().setup(inverted_ids_buffer)->ignore();
0386 
0387   // Make a buffer of boolean elements (Whether a corresponding track id is
0388   // updated after an iteration)
0389   vecmem::data::vector_buffer<int> is_updated_buffer{n_tracks, m_mr.main};
0390   m_copy.get().setup(is_updated_buffer)->ignore();
0391   m_copy.get().memset(is_updated_buffer, 0)->ignore();
0392 
0393   // Count track id apperance during removal process
0394   vecmem::data::vector_buffer<int> track_count_buffer{n_tracks, m_mr.main};
0395   m_copy.get().setup(track_count_buffer)->ignore();
0396   m_copy.get().memset(track_count_buffer, 0)->ignore();
0397 
0398   // Prefix sum buffer used for the insertion sort during an iteration
0399   vecmem::data::vector_buffer<int> prefix_sums_buffer{n_tracks, m_mr.main};
0400   m_copy.get().setup(prefix_sums_buffer)->ignore();
0401 
0402   // Fill the sorted ids vector
0403   thrust::copy(thrust_policy, pre_accepted_ids_buffer.ptr(),
0404                pre_accepted_ids_buffer.ptr() + n_accepted,
0405                sorted_ids_buffer.ptr());
0406   m_stream.synchronize();
0407 
0408   track_comparator trk_comp(rel_shared_buffer.ptr(), pvals_buffer.ptr());
0409 
0410   // Sort the sorted ids vector based on the relative number of shared
0411   // measurements and pvalues
0412   thrust::sort(thrust_policy, sorted_ids_buffer.ptr(),
0413                sorted_ids_buffer.ptr() + n_accepted, trk_comp);
0414 
0415   // Make a buffer of track ids whose number of shared measurements are
0416   // updated during an iteration
0417   vecmem::data::vector_buffer<unsigned int> updated_tracks_buffer{n_accepted,
0418                                                                   m_mr.main};
0419   m_copy.get().setup(updated_tracks_buffer)->ignore();
0420 
0421   // Device objects
0422   vecmem::unique_alloc_ptr<unsigned int> n_removable_tracks_device =
0423       vecmem::make_unique_alloc<unsigned int>(m_mr.main);
0424   vecmem::unique_alloc_ptr<unsigned int> n_meas_to_remove_device =
0425       vecmem::make_unique_alloc<unsigned int>(m_mr.main);
0426   vecmem::unique_alloc_ptr<unsigned int> n_valid_threads_device =
0427       vecmem::make_unique_alloc<unsigned int>(m_mr.main);
0428 
0429   // Whether to terminate the iteration process
0430   int terminate = 0;
0431   vecmem::unique_alloc_ptr<int> terminate_device =
0432       vecmem::make_unique_alloc<int>(m_mr.main);
0433   cudaMemsetAsync(terminate_device.get(), 0, sizeof(int), stream);
0434   auto max_shared = thrust::max_element(thrust::device, n_shared_buffer.ptr(),
0435                                         n_shared_buffer.ptr() + n_tracks);
0436 
0437   // The maximum number of shared measurements. The process is terminated if
0438   // this value is zero
0439   vecmem::unique_alloc_ptr<unsigned int> max_shared_device =
0440       vecmem::make_unique_alloc<unsigned int>(m_mr.main);
0441   cudaMemcpyAsync(max_shared_device.get(), max_shared, sizeof(unsigned int),
0442                   cudaMemcpyDeviceToDevice, stream);
0443 
0444   // The number of tracks whose number of share measurements is updated
0445   vecmem::unique_alloc_ptr<unsigned int> n_updated_tracks_device =
0446       vecmem::make_unique_alloc<unsigned int>(m_mr.main);
0447 
0448   // Thread block size
0449   unsigned int nThreads_adaptive = m_warp_size;
0450   unsigned int nBlocks_adaptive =
0451       (n_accepted + nThreads_adaptive - 1) / nThreads_adaptive;
0452 
0453   unsigned int nThreads_rearrange = 1024;
0454   unsigned int nBlocks_rearrange =
0455       (n_accepted + (nThreads_rearrange / kernels::nThreads_per_track) - 1) /
0456       (nThreads_rearrange / kernels::nThreads_per_track);
0457 
0458   // Compute the threadblock dimension for scanning kernels
0459   auto compute_scan_config = [&](unsigned int n_accepted) {
0460     unsigned int nThreads_scan = m_warp_size * 4;
0461     unsigned int nBlocks_scan =
0462         (n_accepted + nThreads_scan - 1) / nThreads_scan;
0463 
0464     while (nThreads_scan <= 1024) {
0465       if (nBlocks_scan > 1024) {
0466         nThreads_scan *= 2;
0467         nBlocks_scan = (n_accepted + nThreads_scan - 1) / nThreads_scan;
0468       } else {
0469         break;
0470       }
0471     }
0472 
0473     return std::make_pair(nThreads_scan, nBlocks_scan);
0474   };
0475 
0476   auto scan_dim = compute_scan_config(n_accepted);
0477   unsigned int nThreads_scan = scan_dim.first;
0478   unsigned int nBlocks_scan = scan_dim.second;
0479 
0480   assert(nBlocks_scan <= 1024 &&
0481          "nBlocks_scan larger than 1024 will cause invalid arguments in "
0482          "scan_block_offsets kernel");
0483 
0484   // Make buffers used in prefix sum calculation
0485   vecmem::data::vector_buffer<int> block_offsets_buffer{nBlocks_scan,
0486                                                         m_mr.main};
0487   m_copy.get().setup(block_offsets_buffer)->ignore();
0488   vecmem::data::vector_buffer<int> scanned_block_offsets_buffer{nBlocks_scan,
0489                                                                 m_mr.main};
0490   m_copy.get().setup(scanned_block_offsets_buffer)->ignore();
0491 
0492   // Start the iteration
0493   while (!terminate && n_accepted > 0) {
0494     nBlocks_adaptive = (n_accepted + nThreads_adaptive - 1) / nThreads_adaptive;
0495 
0496     scan_dim = compute_scan_config(n_accepted);
0497     nThreads_scan = scan_dim.first;
0498     nBlocks_scan = scan_dim.second;
0499     nBlocks_rearrange =
0500         (n_accepted + (nThreads_rearrange / kernels::nThreads_per_track) - 1) /
0501         (nThreads_rearrange / kernels::nThreads_per_track);
0502 
0503     // Make a CUDA Graph. We use CUDA graph to minimize the overheads from
0504     // kernel launches
0505     cudaGraph_t graph;
0506     cudaGraphExec_t graphExec;
0507 
0508     cudaStreamBeginCapture(stream, cudaStreamCaptureModeGlobal);
0509 
0510     // Counts the number of removable tracks in the current iteration and
0511     // remove them from the track pool
0512     kernels::remove_tracks<<<1, 512, 0, stream>>>(device::remove_tracks_payload{
0513         .sorted_ids_view = sorted_ids_buffer,
0514         .n_accepted = n_accepted_device.get(),
0515         .meas_ids_view = meas_ids_buffer,
0516         .n_meas_view = n_meas_buffer,
0517         .meas_id_to_unique_id_view = meas_id_to_unique_id_buffer,
0518         .tracks_per_measurement_view = tracks_per_measurement_buffer,
0519         .track_status_per_measurement_view =
0520             track_status_per_measurement_buffer,
0521         .n_accepted_tracks_per_measurement_view =
0522             n_accepted_tracks_per_measurement_buffer,
0523         .n_shared_view = n_shared_buffer,
0524         .rel_shared_view = rel_shared_buffer,
0525         .n_removable_tracks = n_removable_tracks_device.get(),
0526         .n_meas_to_remove = n_meas_to_remove_device.get(),
0527         .terminate = terminate_device.get(),
0528         .max_shared = max_shared_device.get(),
0529         .n_updated_tracks = n_updated_tracks_device.get(),
0530         .updated_tracks_view = updated_tracks_buffer,
0531         .is_updated_view = is_updated_buffer,
0532         .n_valid_threads = n_valid_threads_device.get(),
0533         .track_count_view = track_count_buffer});
0534 
0535     // After the kernel "remove_tracks", sorted_ids_view is not sorted
0536     // anymore as the number of measurements of a few tracks might change.
0537     // We can consider using thrust::sort() like the following:
0538     /*
0539     cudaMemcpyAsync(&n_accepted, n_accepted_device.get(),
0540                     sizeof(unsigned int), cudaMemcpyDeviceToHost,
0541                     stream);
0542     thrust::sort(thrust_policy, sorted_ids_buffer.ptr(),
0543                  sorted_ids_buffer.ptr() + n_accepted,
0544                  trk_comp);
0545     */
0546     // However, thrust::sort (Radix sort) is not optimized for our case
0547     // where we only need to rearrange the indices of a few tracks whose
0548     // number of measurement changed. In such case, insertion sort would be
0549     // a good choice and the following seven kernels are collaborating each
0550     // other to do insertion sort
0551 
0552     // This kernel sort the tracks whose number of measurement changed
0553     // during remove_track kernel. The number of such tracks are very small
0554     // and we apply bitoncic sort here. The purpose is to make each of
0555     // updated tracks not interfere each other when we rearrange them during
0556     // the insertion sort
0557     kernels::sort_updated_tracks<<<1, 512, 0, stream>>>(
0558         device::sort_updated_tracks_payload{
0559             .rel_shared_view = rel_shared_buffer,
0560             .pvals_view = pvals_buffer,
0561             .terminate = terminate_device.get(),
0562             .n_updated_tracks = n_updated_tracks_device.get(),
0563             .updated_tracks_view = updated_tracks_buffer,
0564         });
0565 
0566     // Fill the inverted_ids vector which converts a track id to the index
0567     // of sorted ids, which is for the fast lookup
0568     kernels::
0569         fill_inverted_ids<<<nBlocks_adaptive, nThreads_adaptive, 0, stream>>>(
0570             device::fill_inverted_ids_payload{
0571                 .sorted_ids_view = sorted_ids_buffer,
0572                 .terminate = terminate_device.get(),
0573                 .n_accepted = n_accepted_device.get(),
0574                 .n_updated_tracks = n_updated_tracks_device.get(),
0575                 .inverted_ids_view = inverted_ids_buffer,
0576             });
0577 
0578     // The three kernels (block_inclusive_scan, scan_block_offsets, and
0579     // add_block_offset) work together to compute the prefix sum of track
0580     // IDs, with respect to the number of updated tracks, based on the
0581     // indices of sorted_ids. The kernels are splitted as it is not
0582     // efficient to calculate this in a single kernel. This prefix sums are
0583     // used during the insertion sorting in rearrange_tracks to precisely
0584     // calculate the new index of updated tracks
0585 
0586     // Caculate the prefix sum of the number of updated tracks block-wisely.
0587     // block_offset is the last element of block-wise prefix sums, which is
0588     // used to get the real prefix sum later
0589     kernels::block_inclusive_scan<<<nBlocks_scan, nThreads_scan,
0590                                     nThreads_scan * sizeof(int), stream>>>(
0591         device::block_inclusive_scan_payload{
0592             .sorted_ids_view = sorted_ids_buffer,
0593             .terminate = terminate_device.get(),
0594             .n_accepted = n_accepted_device.get(),
0595             .n_updated_tracks = n_updated_tracks_device.get(),
0596             .is_updated_view = is_updated_buffer,
0597             .block_offsets_view = block_offsets_buffer,
0598             .prefix_sums_view = prefix_sums_buffer});
0599 
0600     // Calculate the scanned block offsets which is the prefix sum of block
0601     // offsets
0602     kernels::scan_block_offsets<<<1, nBlocks_scan, nBlocks_scan * sizeof(int),
0603                                   stream>>>(device::scan_block_offsets_payload{
0604         .terminate = terminate_device.get(),
0605         .n_accepted = n_accepted_device.get(),
0606         .n_updated_tracks = n_updated_tracks_device.get(),
0607         .block_offsets_view = block_offsets_buffer,
0608         .scanned_block_offsets_view = scanned_block_offsets_buffer});
0609 
0610     // To calculate the real prefix-sum, add the scanned block offsets to
0611     // block-wise prefix sums of the number of updated tracks.
0612     kernels::add_block_offset<<<nBlocks_scan, nThreads_scan, 0, stream>>>(
0613         device::add_block_offset_payload{
0614             .terminate = terminate_device.get(),
0615             .n_accepted = n_accepted_device.get(),
0616             .n_updated_tracks = n_updated_tracks_device.get(),
0617             .block_offsets_view = scanned_block_offsets_buffer,
0618             .prefix_sums_view = prefix_sums_buffer});
0619 
0620     // Apply the insertion sort algorithm to sorted_ids_view using the
0621     // sorted updated tracks and prefix sums. The sorted elements are stored
0622     // in temp_sorted_ids_view
0623     kernels::
0624         rearrange_tracks<<<nBlocks_rearrange, nThreads_rearrange, 0, stream>>>(
0625             device::rearrange_tracks_payload{
0626                 .sorted_ids_view = sorted_ids_buffer,
0627                 .inverted_ids_view = inverted_ids_buffer,
0628                 .rel_shared_view = rel_shared_buffer,
0629                 .pvals_view = pvals_buffer,
0630                 .terminate = terminate_device.get(),
0631                 .n_accepted = n_accepted_device.get(),
0632                 .n_updated_tracks = n_updated_tracks_device.get(),
0633                 .updated_tracks_view = updated_tracks_buffer,
0634                 .is_updated_view = is_updated_buffer,
0635                 .prefix_sums_view = prefix_sums_buffer,
0636                 .temp_sorted_ids_view = temp_sorted_ids_buffer,
0637             });
0638 
0639     // Find the max shared number of measurements to decide whether to
0640     // terminate the process. Also Move the elements in temp_sorted_ids to
0641     // sorted_ids
0642     kernels::update_status<<<nBlocks_adaptive, nThreads_adaptive, 0, stream>>>(
0643         device::update_status_payload{
0644             .terminate = terminate_device.get(),
0645             .n_accepted = n_accepted_device.get(),
0646             .n_updated_tracks = n_updated_tracks_device.get(),
0647             .temp_sorted_ids_view = temp_sorted_ids_buffer,
0648             .sorted_ids_view = sorted_ids_buffer,
0649             .updated_tracks_view = updated_tracks_buffer,
0650             .is_updated_view = is_updated_buffer,
0651             .n_shared_view = n_shared_buffer,
0652             .max_shared = max_shared_device.get()});
0653 
0654     cudaStreamEndCapture(stream, &graph);
0655     cudaGraphInstantiate(&graphExec, graph, nullptr, nullptr, 0);
0656 
0657     // TODO: Make n_it adaptive based on the average track length, bound
0658     // value in remove_tracks, etc.
0659     const unsigned int n_it = 100;
0660     for (unsigned int iter = 0; iter < n_it; iter++) {
0661       cudaGraphLaunch(graphExec, stream);
0662     }
0663 
0664     cudaMemcpyAsync(&terminate, terminate_device.get(), sizeof(int),
0665                     cudaMemcpyDeviceToHost, stream);
0666     cudaMemcpyAsync(&n_accepted, n_accepted_device.get(), sizeof(unsigned int),
0667                     cudaMemcpyDeviceToHost, stream);
0668     m_stream.synchronize();
0669   }
0670 
0671   cudaMemcpyAsync(&n_accepted, n_accepted_device.get(), sizeof(unsigned int),
0672                   cudaMemcpyDeviceToHost, stream);
0673 
0674   auto max_it =
0675       std::max_element(candidate_sizes.begin(), candidate_sizes.end());
0676   const unsigned int max_cands_size = *max_it;
0677 
0678   // Create resolved candidate buffer
0679   edm::track_container<default_algebra>::buffer res_track_candidates_buffer{
0680       {std::vector<std::size_t>(n_accepted, max_cands_size), m_mr.main,
0681        m_mr.host, vecmem::data::buffer_type::resizable},
0682       {},
0683       tracks_view.measurements};
0684   m_copy.get().setup(res_track_candidates_buffer.tracks)->ignore();
0685 
0686   // Fill the output track candidates
0687   {
0688     if (n_accepted > 0) {
0689       kernels::fill_track_candidates<<<
0690           static_cast<unsigned int>((n_accepted + 63) / 64), 64, 0, stream>>>(
0691           device::fill_track_candidates_payload{
0692               .tracks_view = tracks_view,
0693               .n_accepted = n_accepted,
0694               .sorted_ids_view = sorted_ids_buffer,
0695               .res_tracks_view = res_track_candidates_buffer});
0696       TRACCC_CUDA_ERROR_CHECK(cudaGetLastError());
0697 
0698       m_stream.synchronize();
0699     }
0700   }
0701 
0702   return res_track_candidates_buffer;
0703 }
0704 
0705 }  // namespace traccc::cuda