Back to home page

EIC code displayed by LXR

 
 

    


Warning, /acts/Traccc/device/sycl/src/seeding/triplet_seeding_algorithm.sycl 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 // SYCL include
0009 #include <sycl/sycl.hpp>
0010 
0011 // Library include(s).
0012 #include "../utils/calculate1DimNdRange.hpp"
0013 #include "../utils/get_queue.hpp"
0014 #include "../utils/global_index.hpp"
0015 #include "traccc/sycl/seeding/triplet_seeding_algorithm.hpp"
0016 
0017 // Project include(s).
0018 #include "traccc/seeding/device/count_doublets.hpp"
0019 #include "traccc/seeding/device/count_grid_capacities.hpp"
0020 #include "traccc/seeding/device/count_triplets.hpp"
0021 #include "traccc/seeding/device/find_doublets.hpp"
0022 #include "traccc/seeding/device/find_triplets.hpp"
0023 #include "traccc/seeding/device/populate_grid.hpp"
0024 #include "traccc/seeding/device/reduce_triplet_counts.hpp"
0025 #include "traccc/seeding/device/select_seeds.hpp"
0026 #include "traccc/seeding/device/update_triplet_weights.hpp"
0027 
0028 // VecMem include(s).
0029 #include <vecmem/utils/sycl/local_accessor.hpp>
0030 
0031 namespace traccc::sycl {
0032 namespace kernels {
0033 
0034 /// Class identifying the SYCL kernel that runs
0035 /// @c traccc::device::count_grid_capacities
0036 class count_grid_capacities;
0037 
0038 /// Class identifying the SYCL kernel that runs @c traccc::device::populate_grid
0039 class populate_grid;
0040 
0041 /// Class identifying the kernel running @c traccc::device::count_doublets
0042 class count_doublets;
0043 
0044 /// Class identifying the kernel running @c traccc::device::find_doublets
0045 class find_doublets;
0046 
0047 /// Class identifying the kernel running @c traccc::device::count_triplets
0048 class count_triplets;
0049 
0050 /// Class identifying the kernel running @c
0051 /// traccc::device::reduce_triplet_counts
0052 class reduce_triplet_counts;
0053 
0054 /// Class identifying the kernel running @c traccc::device::find_triplets
0055 class find_triplets;
0056 
0057 /// Class identifying the kernel running @c
0058 /// traccc::device::update_triplet_weights
0059 class update_triplet_weights;
0060 
0061 /// Class identifying the kernel running @c traccc::device::select_seeds
0062 class select_seeds;
0063 
0064 }  // namespace kernels
0065 
0066 triplet_seeding_algorithm::triplet_seeding_algorithm(
0067     const seedfinder_config& finder_config,
0068     const spacepoint_grid_config& grid_config,
0069     const seedfilter_config& filter_config, const traccc::memory_resource& mr,
0070     const vecmem::copy& copy, queue_wrapper& queue,
0071     std::unique_ptr<const Logger> logger)
0072     : device::triplet_seeding_algorithm(finder_config, grid_config,
0073                                         filter_config, mr, copy,
0074                                         std::move(logger)),
0075       sycl::algorithm_base{queue} {}
0076 
0077 void triplet_seeding_algorithm::count_grid_capacities_kernel(
0078     const count_grid_capacities_kernel_payload& payload) const {
0079   ::sycl::queue& squeue = details::get_queue(queue());
0080   squeue.throw_asynchronous();
0081   squeue.submit([&](::sycl::handler& h) {
0082     h.parallel_for<kernels::count_grid_capacities>(
0083         details::calculate1DimNdRange(payload.n_spacepoints, warp_size() * 4),
0084         [config = payload.config, phi_axis = payload.phi_axis,
0085          z_axis = payload.z_axis, spacepoints = payload.spacepoints,
0086          grid_capacities = payload.grid_capacities](::sycl::nd_item<1> item) {
0087           device::count_grid_capacities(details::global_index(item), config,
0088                                         phi_axis, z_axis, spacepoints,
0089                                         grid_capacities);
0090         });
0091   });
0092 }
0093 
0094 void triplet_seeding_algorithm::populate_grid_kernel(
0095     const populate_grid_kernel_payload& payload) const {
0096   ::sycl::queue& squeue = details::get_queue(queue());
0097   squeue.throw_asynchronous();
0098   squeue.submit([&](::sycl::handler& h) {
0099     h.parallel_for<kernels::populate_grid>(
0100         details::calculate1DimNdRange(payload.n_spacepoints, warp_size() * 4),
0101         [config = payload.config, spacepoints = payload.spacepoints,
0102          grid = payload.grid,
0103          grid_prefix_sum = payload.grid_prefix_sum](::sycl::nd_item<1> item) {
0104           device::populate_grid(details::global_index(item), config,
0105                                 spacepoints, grid, grid_prefix_sum);
0106         });
0107   });
0108 }
0109 
0110 void triplet_seeding_algorithm::count_doublets_kernel(
0111     const count_doublets_kernel_payload& payload) const {
0112   ::sycl::queue& squeue = details::get_queue(queue());
0113   squeue.throw_asynchronous();
0114   squeue.submit([&](::sycl::handler& h) {
0115     h.parallel_for<kernels::count_doublets>(
0116         details::calculate1DimNdRange(payload.n_spacepoints, warp_size() * 2),
0117         [config = payload.config, spacepoints = payload.spacepoints,
0118          grid = payload.grid, grid_prefix_sum = payload.grid_prefix_sum,
0119          doublet_counter = payload.doublet_counter,
0120          nMidBot = &(payload.nMidBot),
0121          nMidTop = &(payload.nMidTop)](::sycl::nd_item<1> item) {
0122           device::count_doublets(details::global_index(item), config,
0123                                  spacepoints, grid, grid_prefix_sum,
0124                                  doublet_counter, *nMidBot, *nMidTop);
0125         });
0126   });
0127 }
0128 
0129 void triplet_seeding_algorithm::find_doublets_kernel(
0130     const find_doublets_kernel_payload& payload) const {
0131   ::sycl::queue& squeue = details::get_queue(queue());
0132   squeue.throw_asynchronous();
0133   squeue.submit([&](::sycl::handler& h) {
0134     h.parallel_for<kernels::find_doublets>(
0135         details::calculate1DimNdRange(payload.n_doublets, warp_size() * 2),
0136         [config = payload.config, spacepoints = payload.spacepoints,
0137          grid = payload.grid, doublet_counter = payload.doublet_counter,
0138          mb_doublets = payload.mb_doublets,
0139          mt_doublets = payload.mt_doublets](::sycl::nd_item<1> item) {
0140           device::find_doublets(details::global_index(item), config,
0141                                 spacepoints, grid, doublet_counter, mb_doublets,
0142                                 mt_doublets);
0143         });
0144   });
0145 }
0146 
0147 void triplet_seeding_algorithm::count_triplets_kernel(
0148     const count_triplets_kernel_payload& payload) const {
0149   ::sycl::queue& squeue = details::get_queue(queue());
0150   squeue.throw_asynchronous();
0151   squeue.submit([&](::sycl::handler& h) {
0152     h.parallel_for<kernels::count_triplets>(
0153         details::calculate1DimNdRange(payload.nMidBot, warp_size() * 2),
0154         [config = payload.config, spacepoints = payload.spacepoints,
0155          grid = payload.grid, doublet_counter = payload.doublet_counter,
0156          mb_doublets = payload.mb_doublets, mt_doublets = payload.mt_doublets,
0157          spM_counter = payload.spM_counter,
0158          midBot_counter = payload.midBot_counter](::sycl::nd_item<1> item) {
0159           device::count_triplets(details::global_index(item), config,
0160                                  spacepoints, grid, doublet_counter,
0161                                  mb_doublets, mt_doublets, spM_counter,
0162                                  midBot_counter);
0163         });
0164   });
0165 }
0166 
0167 void triplet_seeding_algorithm::triplet_counts_reduction_kernel(
0168     const triplet_counts_reduction_kernel_payload& payload) const {
0169   ::sycl::queue& squeue = details::get_queue(queue());
0170   squeue.throw_asynchronous();
0171   squeue.submit([&](::sycl::handler& h) {
0172     h.parallel_for<kernels::reduce_triplet_counts>(
0173         details::calculate1DimNdRange(payload.n_doublets, warp_size() * 2),
0174         [doublet_counter = payload.doublet_counter,
0175          spM_counter = payload.spM_counter,
0176          nTriplets = &(payload.nTriplets)](::sycl::nd_item<1> item) {
0177           device::reduce_triplet_counts(details::global_index(item),
0178                                         doublet_counter, spM_counter,
0179                                         *nTriplets);
0180         });
0181   });
0182 }
0183 
0184 void triplet_seeding_algorithm::find_triplets_kernel(
0185     const find_triplets_kernel_payload& payload) const {
0186   ::sycl::queue& squeue = details::get_queue(queue());
0187   squeue.throw_asynchronous();
0188   squeue.submit([&](::sycl::handler& h) {
0189     h.parallel_for<kernels::find_triplets>(
0190         details::calculate1DimNdRange(payload.nMidBot, warp_size() * 2),
0191         [finding_config = payload.finding_config,
0192          filter_config = payload.filter_config,
0193          spacepoints = payload.spacepoints, grid = payload.grid,
0194          doublet_counter = payload.doublet_counter,
0195          mt_doublets = payload.mt_doublets, spM_tc = payload.spM_tc,
0196          midBot_tc = payload.midBot_tc,
0197          triplets = payload.triplets](::sycl::nd_item<1> item) {
0198           device::find_triplets(details::global_index(item), finding_config,
0199                                 filter_config, spacepoints, grid,
0200                                 doublet_counter, mt_doublets, spM_tc, midBot_tc,
0201                                 triplets);
0202         });
0203   });
0204 }
0205 
0206 void triplet_seeding_algorithm::update_triplet_weights_kernel(
0207     const update_triplet_weights_kernel_payload& payload) const {
0208   ::sycl::queue& squeue = details::get_queue(queue());
0209   squeue.throw_asynchronous();
0210 
0211   const unsigned int n_threads = warp_size() * 2;
0212 
0213   // Check if device is capable of allocating sufficient local memory
0214   assert(sizeof(scalar) * payload.config.compatSeedLimit * n_threads <
0215          squeue.get_device().get_info<::sycl::info::device::local_mem_size>());
0216 
0217   squeue.submit([&](::sycl::handler& h) {
0218     // Array for temporary storage of triplet weights for comparing
0219     // within kernel
0220     vecmem::sycl::local_accessor<scalar> local_mem(
0221         payload.config.compatSeedLimit * n_threads, h);
0222 
0223     h.parallel_for<kernels::update_triplet_weights>(
0224         details::calculate1DimNdRange(payload.n_triplets, n_threads),
0225         [config = payload.config, spacepoints = payload.spacepoints,
0226          spM_tc = payload.spM_tc, midBot_tc = payload.midBot_tc,
0227          triplets = payload.triplets, local_mem](::sycl::nd_item<1> item) {
0228           // Each thread uses compatSeedLimit elements of the array
0229           scalar* dataPos =
0230               &local_mem[item.get_local_id() * config.compatSeedLimit];
0231           device::update_triplet_weights(details::global_index(item), config,
0232                                          spacepoints, spM_tc, midBot_tc,
0233                                          dataPos, triplets);
0234         });
0235   });
0236 }
0237 
0238 void triplet_seeding_algorithm::select_seeds_kernel(
0239     const select_seeds_kernel_payload& payload) const {
0240   ::sycl::queue& squeue = details::get_queue(queue());
0241   squeue.throw_asynchronous();
0242 
0243   const unsigned int n_threads = warp_size() * 2;
0244 
0245   // Check if device is capable of allocating sufficient local memory
0246   assert(sizeof(triplet) * payload.finder_config.maxSeedsPerSpM * n_threads <
0247          squeue.get_device().get_info<::sycl::info::device::local_mem_size>());
0248 
0249   squeue.submit([&](::sycl::handler& h) {
0250     // Array for temporary storage of triplets for comparing within
0251     // kernel
0252     vecmem::sycl::local_accessor<device::device_triplet> local_mem(
0253         payload.finder_config.maxSeedsPerSpM * n_threads, h);
0254 
0255     h.parallel_for<kernels::select_seeds>(
0256         details::calculate1DimNdRange(payload.n_doublets, n_threads),
0257         [finder_config = payload.finder_config,
0258          filter_config = payload.filter_config,
0259          spacepoints = payload.spacepoints, grid = payload.grid,
0260          spM_tc = payload.spM_tc, midBot_tc = payload.midBot_tc,
0261          triplets = payload.triplets, seeds = payload.seeds,
0262          local_mem](::sycl::nd_item<1> item) {
0263           // Each thread uses compatSeedLimit elements of the array
0264           device::device_triplet* dataPos =
0265               &local_mem[item.get_local_id() * finder_config.maxSeedsPerSpM];
0266 
0267           device::select_seeds(details::global_index(item), finder_config,
0268                                filter_config, spacepoints, grid, spM_tc,
0269                                midBot_tc, triplets, dataPos, seeds);
0270         });
0271   });
0272 }
0273 
0274 }  // namespace traccc::sycl