Back to home page

EIC code displayed by LXR

 
 

    


Warning, /acts/Traccc/device/hip/src/seeding/triplet_seeding_algorithm.hip is written in an unsupported language. File is not indexed.

0001 /** TRACCC library, part of the ACTS project (R&D line)
0002  *
0003  * (c) 2021-2026 CERN for the benefit of the ACTS project
0004  *
0005  * Mozilla Public License Version 2.0
0006  */
0007 
0008 // Library include(s).
0009 #include "../utils/global_index.hpp"
0010 #include "../utils/hip_error_handling.hpp"
0011 #include "../utils/utils.hpp"
0012 #include "traccc/hip/seeding/triplet_seeding_algorithm.hpp"
0013 
0014 // Project include(s).
0015 #include "traccc/seeding/detail/spacepoint_grid.hpp"
0016 #include "traccc/seeding/device/count_doublets.hpp"
0017 #include "traccc/seeding/device/count_grid_capacities.hpp"
0018 #include "traccc/seeding/device/count_triplets.hpp"
0019 #include "traccc/seeding/device/find_doublets.hpp"
0020 #include "traccc/seeding/device/find_triplets.hpp"
0021 #include "traccc/seeding/device/populate_grid.hpp"
0022 #include "traccc/seeding/device/reduce_triplet_counts.hpp"
0023 #include "traccc/seeding/device/select_seeds.hpp"
0024 #include "traccc/seeding/device/update_triplet_weights.hpp"
0025 
0026 namespace traccc::hip {
0027 namespace kernels {
0028 
0029 /// HIP kernel for running @c traccc::device::count_grid_capacities
0030 __global__ void count_grid_capacities(
0031     seedfinder_config config,
0032     traccc::details::spacepoint_grid_types::host::axis_p0_type phi_axis,
0033     traccc::details::spacepoint_grid_types::host::axis_p1_type z_axis,
0034     edm::spacepoint_collection::const_view spacepoints,
0035     vecmem::data::vector_view<unsigned int> grid_capacities) {
0036   device::count_grid_capacities(details::global_index1(), config, phi_axis,
0037                                 z_axis, spacepoints, grid_capacities);
0038 }
0039 
0040 /// HIP kernel for running @c traccc::device::populate_grid
0041 __global__ void populate_grid(
0042     seedfinder_config config,
0043     edm::spacepoint_collection::const_view spacepoints,
0044     traccc::details::spacepoint_grid_types::view grid,
0045     vecmem::data::vector_view<device::prefix_sum_element_t> grid_prefix_sum) {
0046   device::populate_grid(details::global_index1(), config, spacepoints, grid,
0047                         grid_prefix_sum);
0048 }
0049 
0050 /// HIP kernel for running @c traccc::device::count_doublets
0051 __global__ void count_doublets(
0052     seedfinder_config config,
0053     edm::spacepoint_collection::const_view spacepoints,
0054     traccc::details::spacepoint_grid_types::const_view sp_grid,
0055     vecmem::data::vector_view<const device::prefix_sum_element_t> sp_prefix_sum,
0056     device::doublet_counter_collection_types::view doublet_counter,
0057     unsigned int& nMidBot, unsigned int& nMidTop) {
0058   device::count_doublets(details::global_index1(), config, spacepoints, sp_grid,
0059                          sp_prefix_sum, doublet_counter, nMidBot, nMidTop);
0060 }
0061 
0062 /// HIP kernel for running @c traccc::device::find_doublets
0063 __global__ void find_doublets(
0064     seedfinder_config config,
0065     edm::spacepoint_collection::const_view spacepoints,
0066     traccc::details::spacepoint_grid_types::const_view sp_grid,
0067     device::doublet_counter_collection_types::const_view doublet_counter,
0068     device::device_doublet_collection_types::view mb_doublets,
0069     device::device_doublet_collection_types::view mt_doublets) {
0070   device::find_doublets(details::global_index1(), config, spacepoints, sp_grid,
0071                         doublet_counter, mb_doublets, mt_doublets);
0072 }
0073 
0074 /// HIP kernel for running @c traccc::device::count_triplets
0075 __global__ void count_triplets(
0076     seedfinder_config config,
0077     edm::spacepoint_collection::const_view spacepoints,
0078     traccc::details::spacepoint_grid_types::const_view sp_grid,
0079     device::doublet_counter_collection_types::const_view doublet_counter,
0080     device::device_doublet_collection_types::const_view mb_doublets,
0081     device::device_doublet_collection_types::const_view mt_doublets,
0082     device::triplet_counter_spM_collection_types::view spM_counter,
0083     device::triplet_counter_collection_types::view midBot_counter) {
0084   device::count_triplets(details::global_index1(), config, spacepoints, sp_grid,
0085                          doublet_counter, mb_doublets, mt_doublets, spM_counter,
0086                          midBot_counter);
0087 }
0088 
0089 /// HIP kernel for running @c traccc::device::reduce_triplet_counts
0090 __global__ void reduce_triplet_counts(
0091     device::doublet_counter_collection_types::const_view doublet_counter,
0092     device::triplet_counter_spM_collection_types::view spM_counter,
0093     unsigned int& num_triplets) {
0094   device::reduce_triplet_counts(details::global_index1(), doublet_counter,
0095                                 spM_counter, num_triplets);
0096 }
0097 
0098 /// HIP kernel for running @c traccc::device::find_triplets
0099 __global__ void find_triplets(
0100     seedfinder_config config, seedfilter_config filter_config,
0101     edm::spacepoint_collection::const_view spacepoints,
0102     traccc::details::spacepoint_grid_types::const_view sp_grid,
0103     device::doublet_counter_collection_types::const_view doublet_counter,
0104     device::device_doublet_collection_types::const_view mt_doublets,
0105     device::triplet_counter_spM_collection_types::const_view spM_tc,
0106     device::triplet_counter_collection_types::const_view midBot_tc,
0107     device::device_triplet_collection_types::view triplet_view) {
0108   device::find_triplets(details::global_index1(), config, filter_config,
0109                         spacepoints, sp_grid, doublet_counter, mt_doublets,
0110                         spM_tc, midBot_tc, triplet_view);
0111 }
0112 
0113 /// HIP kernel for running @c traccc::device::update_triplet_weights
0114 __global__ void update_triplet_weights(
0115     seedfilter_config filter_config,
0116     edm::spacepoint_collection::const_view spacepoints,
0117     device::triplet_counter_spM_collection_types::const_view spM_tc,
0118     device::triplet_counter_collection_types::const_view midBot_tc,
0119     device::device_triplet_collection_types::view triplet_view) {
0120   // Array for temporary storage of quality parameters for comparing triplets
0121   // within weight updating kernel
0122   extern __shared__ scalar data[];
0123   // Each thread uses compatSeedLimit elements of the array
0124   scalar* dataPos = &data[threadIdx.x * filter_config.compatSeedLimit];
0125 
0126   device::update_triplet_weights(details::global_index1(), filter_config,
0127                                  spacepoints, spM_tc, midBot_tc, dataPos,
0128                                  triplet_view);
0129 }
0130 
0131 /// HIP kernel for running @c traccc::device::select_seeds
0132 __global__ void select_seeds(
0133     seedfinder_config finder_config, seedfilter_config filter_config,
0134     edm::spacepoint_collection::const_view spacepoints,
0135     traccc::details::spacepoint_grid_types::const_view sp_view,
0136     device::triplet_counter_spM_collection_types::const_view spM_tc,
0137     device::triplet_counter_collection_types::const_view midBot_tc,
0138     device::device_triplet_collection_types::const_view triplet_view,
0139     edm::seed_collection::view seed_view) {
0140   // Array for temporary storage of triplets for comparing within seed
0141   // selecting kernel
0142   extern __shared__ device::device_triplet data2[];
0143   // Each thread uses max_triplets_per_spM elements of the array
0144   device::device_triplet* dataPos =
0145       &data2[threadIdx.x * finder_config.maxSeedsPerSpM];
0146 
0147   device::select_seeds(details::global_index1(), finder_config, filter_config,
0148                        spacepoints, sp_view, spM_tc, midBot_tc, triplet_view,
0149                        dataPos, seed_view);
0150 }
0151 
0152 }  // namespace kernels
0153 
0154 triplet_seeding_algorithm::triplet_seeding_algorithm(
0155     const seedfinder_config& finder_config,
0156     const spacepoint_grid_config& grid_config,
0157     const seedfilter_config& filter_config, const traccc::memory_resource& mr,
0158     const vecmem::copy& copy, const stream_wrapper& str,
0159     std::unique_ptr<const Logger> logger)
0160     : device::triplet_seeding_algorithm(finder_config, grid_config,
0161                                         filter_config, mr, copy,
0162                                         std::move(logger)),
0163       hip::algorithm_base{str} {}
0164 
0165 void triplet_seeding_algorithm::count_grid_capacities_kernel(
0166     const count_grid_capacities_kernel_payload& payload) const {
0167   const unsigned int n_threads = warp_size() * 8;
0168   const unsigned int n_blocks =
0169       (payload.n_spacepoints + n_threads - 1) / n_threads;
0170   hipLaunchKernelGGL(kernels::count_grid_capacities, n_blocks, n_threads, 0,
0171                      details::get_stream(stream()), payload.config,
0172                      payload.phi_axis, payload.z_axis, payload.spacepoints,
0173                      payload.grid_capacities);
0174   TRACCC_HIP_ERROR_CHECK(hipGetLastError());
0175 }
0176 
0177 void triplet_seeding_algorithm::populate_grid_kernel(
0178     const populate_grid_kernel_payload& payload) const {
0179   const unsigned int n_threads = warp_size() * 8;
0180   const unsigned int n_blocks =
0181       (payload.n_spacepoints + n_threads - 1) / n_threads;
0182   hipLaunchKernelGGL(kernels::populate_grid, n_blocks, n_threads, 0,
0183                      details::get_stream(stream()), payload.config,
0184                      payload.spacepoints, payload.grid,
0185                      payload.grid_prefix_sum);
0186   TRACCC_HIP_ERROR_CHECK(hipGetLastError());
0187 }
0188 
0189 void triplet_seeding_algorithm::count_doublets_kernel(
0190     const count_doublets_kernel_payload& payload) const {
0191   const unsigned int n_threads = warp_size() * 2;
0192   const unsigned int n_blocks =
0193       (payload.n_spacepoints + n_threads - 1) / n_threads;
0194   hipLaunchKernelGGL(kernels::count_doublets, n_blocks, n_threads, 0,
0195                      details::get_stream(stream()), payload.config,
0196                      payload.spacepoints, payload.grid, payload.grid_prefix_sum,
0197                      payload.doublet_counter, payload.nMidBot, payload.nMidTop);
0198   TRACCC_HIP_ERROR_CHECK(hipGetLastError());
0199 }
0200 
0201 void triplet_seeding_algorithm::find_doublets_kernel(
0202     const find_doublets_kernel_payload& payload) const {
0203   const unsigned int n_threads = warp_size() * 2;
0204   const unsigned int n_blocks =
0205       (payload.n_doublets + n_threads - 1) / n_threads;
0206   hipLaunchKernelGGL(kernels::find_doublets, n_blocks, n_threads, 0,
0207                      details::get_stream(stream()), payload.config,
0208                      payload.spacepoints, payload.grid, payload.doublet_counter,
0209                      payload.mb_doublets, payload.mt_doublets);
0210   TRACCC_HIP_ERROR_CHECK(hipGetLastError());
0211 }
0212 
0213 void triplet_seeding_algorithm::count_triplets_kernel(
0214     const count_triplets_kernel_payload& payload) const {
0215   const unsigned int n_threads = warp_size() * 2;
0216   const unsigned int n_blocks = (payload.nMidBot + n_threads - 1) / n_threads;
0217   hipLaunchKernelGGL(kernels::count_triplets, n_blocks, n_threads, 0,
0218                      details::get_stream(stream()), payload.config,
0219                      payload.spacepoints, payload.grid, payload.doublet_counter,
0220                      payload.mb_doublets, payload.mt_doublets,
0221                      payload.spM_counter, payload.midBot_counter);
0222   TRACCC_HIP_ERROR_CHECK(hipGetLastError());
0223 }
0224 
0225 void triplet_seeding_algorithm::triplet_counts_reduction_kernel(
0226     const triplet_counts_reduction_kernel_payload& payload) const {
0227   const unsigned int n_threads = warp_size() * 2;
0228   const unsigned int n_blocks =
0229       (payload.n_doublets + n_threads - 1) / n_threads;
0230   hipLaunchKernelGGL(kernels::reduce_triplet_counts, n_blocks, n_threads, 0,
0231                      details::get_stream(stream()), payload.doublet_counter,
0232                      payload.spM_counter, payload.nTriplets);
0233   TRACCC_HIP_ERROR_CHECK(hipGetLastError());
0234 }
0235 
0236 void triplet_seeding_algorithm::find_triplets_kernel(
0237     const find_triplets_kernel_payload& payload) const {
0238   const unsigned int n_threads = warp_size() * 2;
0239   const unsigned int n_blocks = (payload.nMidBot + n_threads - 1) / n_threads;
0240   hipLaunchKernelGGL(kernels::find_triplets, n_blocks, n_threads, 0,
0241                      details::get_stream(stream()), payload.finding_config,
0242                      payload.filter_config, payload.spacepoints, payload.grid,
0243                      payload.doublet_counter, payload.mt_doublets,
0244                      payload.spM_tc, payload.midBot_tc, payload.triplets);
0245   TRACCC_HIP_ERROR_CHECK(hipGetLastError());
0246 }
0247 
0248 void triplet_seeding_algorithm::update_triplet_weights_kernel(
0249     const update_triplet_weights_kernel_payload& payload) const {
0250   const unsigned int n_threads = warp_size() * 2;
0251   const unsigned int n_blocks =
0252       (payload.n_triplets + n_threads - 1) / n_threads;
0253   hipLaunchKernelGGL(
0254       kernels::update_triplet_weights, n_blocks, n_threads,
0255       sizeof(scalar) * payload.config.compatSeedLimit * n_threads,
0256       details::get_stream(stream()), payload.config, payload.spacepoints,
0257       payload.spM_tc, payload.midBot_tc, payload.triplets);
0258   TRACCC_HIP_ERROR_CHECK(hipGetLastError());
0259 }
0260 
0261 void triplet_seeding_algorithm::select_seeds_kernel(
0262     const select_seeds_kernel_payload& payload) const {
0263   const unsigned int n_threads = warp_size() * 2;
0264   const unsigned int n_blocks =
0265       (payload.n_doublets + n_threads - 1) / n_threads;
0266   hipLaunchKernelGGL(kernels::select_seeds, n_blocks, n_threads, 0,
0267                      details::get_stream(stream()), payload.finder_config,
0268                      payload.filter_config, payload.spacepoints, payload.grid,
0269                      payload.spM_tc, payload.midBot_tc, payload.triplets,
0270                      payload.seeds);
0271   TRACCC_HIP_ERROR_CHECK(hipGetLastError());
0272 }
0273 
0274 }  // namespace traccc::hip