Back to home page

EIC code displayed by LXR

 
 

    


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

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 // Local include(s).
0009 #include "traccc/alpaka/fitting/kalman_fitting_algorithm.hpp"
0010 
0011 #include "../utils/get_queue.hpp"
0012 #include "../utils/magnetic_field_types.hpp"
0013 #include "../utils/parallel_algorithms.hpp"
0014 #include "../utils/utils.hpp"
0015 
0016 // Project include(s).
0017 #include "traccc/fitting/details/kalman_fitting_types.hpp"
0018 #include "traccc/fitting/device/fill_fitting_sort_keys.hpp"
0019 #include "traccc/fitting/device/fit_backward.hpp"
0020 #include "traccc/fitting/device/fit_forward.hpp"
0021 #include "traccc/fitting/device/fit_prelude.hpp"
0022 #include "traccc/utils/detector_buffer_bfield_visitor.hpp"
0023 
0024 namespace traccc::alpaka {
0025 namespace kernels {
0026 
0027 /// Alpaka kernel functor for @c traccc::device::fill_fitting_sort_keys
0028 struct fill_fitting_sort_keys {
0029   template <typename TAcc>
0030   ALPAKA_FN_ACC void operator()(
0031       TAcc const& acc,
0032       edm::track_collection<default_algebra>::const_view track_candidates_view,
0033       vecmem::data::vector_view<device::sort_key> keys_view,
0034       vecmem::data::vector_view<unsigned int> ids_view) const {
0035     const device::global_index_t globalThreadIdx =
0036         ::alpaka::getIdx<::alpaka::Grid, ::alpaka::Threads>(acc)[0];
0037     device::fill_fitting_sort_keys(globalThreadIdx, track_candidates_view,
0038                                    keys_view, ids_view);
0039   }
0040 };
0041 
0042 /// Alpaka kernel functor for @c traccc::device::fit_prelude
0043 struct fit_prelude {
0044   template <typename TAcc>
0045   ALPAKA_FN_ACC void operator()(TAcc const& acc,
0046                                 device::fit_prelude_payload payload) const {
0047     const device::global_index_t globalThreadIdx =
0048         ::alpaka::getIdx<::alpaka::Grid, ::alpaka::Threads>(acc)[0];
0049     device::fit_prelude(globalThreadIdx, payload);
0050   }
0051 };
0052 
0053 /// Alpaka kernel functor for @c traccc::device::fit_forward
0054 template <typename fitter_t>
0055 struct fit_forward {
0056   template <typename TAcc>
0057   ALPAKA_FN_ACC void operator()(
0058       TAcc const& acc, const typename fitter_t::config_type cfg,
0059       const device::fit_payload payload,
0060       const device::fit_tpayload<
0061           typename fitter_t::detector_type::const_view_type,
0062           typename fitter_t::bfield_type, typename fitter_t::surface_type>*
0063           tpayload) const {
0064     const device::global_index_t globalThreadIdx =
0065         ::alpaka::getIdx<::alpaka::Grid, ::alpaka::Threads>(acc)[0];
0066     device::fit_forward<fitter_t>(globalThreadIdx, cfg, payload, *tpayload);
0067   }
0068 };
0069 
0070 /// Alpaka kernel functor for @c traccc::device::fit_backward
0071 template <typename fitter_t>
0072 struct fit_backward {
0073   template <typename TAcc>
0074   ALPAKA_FN_ACC void operator()(
0075       TAcc const& acc, const typename fitter_t::config_type cfg,
0076       const device::fit_payload payload,
0077       const device::fit_tpayload<
0078           typename fitter_t::detector_type::const_view_type,
0079           typename fitter_t::bfield_type, typename fitter_t::surface_type>*
0080           tpayload) const {
0081     const device::global_index_t globalThreadIdx =
0082         ::alpaka::getIdx<::alpaka::Grid, ::alpaka::Threads>(acc)[0];
0083     device::fit_backward<fitter_t>(globalThreadIdx, cfg, payload, *tpayload);
0084   }
0085 };
0086 
0087 }  // namespace kernels
0088 
0089 kalman_fitting_algorithm::kalman_fitting_algorithm(
0090     const config_type& config, const traccc::memory_resource& mr,
0091     const vecmem::copy& copy, alpaka::queue& q,
0092     std::unique_ptr<const Logger> logger)
0093     : device::kalman_fitting_algorithm{config, mr, copy, std::move(logger)},
0094       alpaka::algorithm_base{q} {}
0095 
0096 void kalman_fitting_algorithm::prepare_track_fit_order(
0097     const edm::track_collection<default_algebra>::const_view& tracks,
0098     vecmem::data::vector_view<device::sort_key>& track_sort_keys,
0099     vecmem::data::vector_view<unsigned int>& track_indices) const {
0100   // Get the number of tracks.
0101   const unsigned int n_tracks = tracks.capacity();
0102   assert(n_tracks == copy().get_size(tracks));
0103   assert(n_tracks == track_indices.capacity());
0104   assert(track_indices.size_ptr() == nullptr);
0105 
0106   // Launch parameters for the kernel.
0107   const unsigned int nThreads = warp_size() * 4;
0108   const unsigned int nBlocks = (n_tracks + nThreads - 1) / nThreads;
0109   auto workDiv = makeWorkDiv<Acc>(nBlocks, nThreads);
0110 
0111   // Fill the keys and indices buffers.
0112   ::alpaka::exec<Acc>(details::get_queue(queue()), workDiv,
0113                       kernels::fill_fitting_sort_keys{}, tracks,
0114                       track_sort_keys, track_indices);
0115 
0116   // Sort the key to get the sorted parameter ids
0117   vecmem::device_vector<device::sort_key> keys_device(track_sort_keys);
0118   vecmem::device_vector<unsigned int> track_indices_device(track_indices);
0119   details::sort_by_key(details::get_queue(queue()), mr(), keys_device.begin(),
0120                        keys_device.end(), track_indices_device.begin());
0121 }
0122 
0123 void kalman_fitting_algorithm::fit_prelude_kernel(
0124     const device::fit_prelude_payload& payload) const {
0125   // Get the number of tracks.
0126   const unsigned int n_tracks = payload.input_tracks.tracks.capacity();
0127   assert(n_tracks == copy().get_size(payload.input_tracks.tracks));
0128   assert(n_tracks == payload.track_indices.capacity());
0129   assert(payload.track_indices.size_ptr() == nullptr);
0130   assert(n_tracks == copy().get_size(payload.output_tracks.tracks));
0131 
0132   // Launch parameters for the kernel.
0133   const unsigned int nThreads = warp_size() * 4;
0134   const unsigned int nBlocks = (n_tracks + nThreads - 1) / nThreads;
0135   auto workDiv = makeWorkDiv<Acc>(nBlocks, nThreads);
0136 
0137   // Run the fitting, using the sorted parameter IDs.
0138   ::alpaka::exec<Acc>(details::get_queue(queue()), workDiv,
0139                       kernels::fit_prelude{}, payload);
0140 }
0141 
0142 auto kalman_fitting_algorithm::prepare_fit_payload(
0143     const detector_buffer& det, const magnetic_field& field,
0144     const std::vector<unsigned int>& n_surfaces,
0145     const device::fit_payload& payload) const -> fit_payload {
0146   return prepare_fit_payload_helper<detector_type_list,
0147                                     alpaka::bfield_type_list<scalar>>(
0148       det, field, n_surfaces, payload);
0149 }
0150 
0151 void kalman_fitting_algorithm::fit_forward_kernel(
0152     const fitting_config& config, const fit_payload& payload) const {
0153   return detector_buffer_magnetic_field_visitor<
0154       detector_type_list, alpaka::bfield_type_list<scalar>>(
0155       payload.detector, payload.field,
0156       [&]<typename detector_traits_t, typename bfield_view_t>(
0157           const typename detector_traits_t::view&, const bfield_view_t&) {
0158         // Get the number of tracks.
0159         const unsigned int n_tracks = payload.payload.tracks.tracks.capacity();
0160         assert(n_tracks == copy().get_size(payload.payload.tracks.tracks));
0161 
0162         // Launch parameters for the kernel.
0163         const unsigned int nThreads = warp_size() * 4;
0164         const unsigned int nBlocks = (n_tracks + nThreads - 1) / nThreads;
0165         auto workDiv = makeWorkDiv<Acc>(nBlocks, nThreads);
0166 
0167         // Fitter type to use.
0168         using fitter_t =
0169             traccc::details::kalman_fitter_t<typename detector_traits_t::device,
0170                                              bfield_view_t>;
0171 
0172         // Run the track fitting
0173         ::alpaka::exec<Acc>(details::get_queue(queue()), workDiv,
0174                             kernels::fit_forward<fitter_t>{}, config,
0175                             payload.payload, payload.get_tpayload<fitter_t>());
0176       });
0177 }
0178 
0179 void kalman_fitting_algorithm::fit_backward_kernel(
0180     const fitting_config& config, const fit_payload& payload) const {
0181   return detector_buffer_magnetic_field_visitor<
0182       detector_type_list, alpaka::bfield_type_list<scalar>>(
0183       payload.detector, payload.field,
0184       [&]<typename detector_traits_t, typename bfield_view_t>(
0185           const typename detector_traits_t::view&, const bfield_view_t&) {
0186         // Get the number of tracks.
0187         const unsigned int n_tracks = payload.payload.tracks.tracks.capacity();
0188         assert(n_tracks == copy().get_size(payload.payload.tracks.tracks));
0189 
0190         // Launch parameters for the kernel.
0191         const unsigned int nThreads = warp_size() * 4;
0192         const unsigned int nBlocks = (n_tracks + nThreads - 1) / nThreads;
0193         auto workDiv = makeWorkDiv<Acc>(nBlocks, nThreads);
0194 
0195         // Fitter type to use.
0196         using fitter_t =
0197             traccc::details::kalman_fitter_t<typename detector_traits_t::device,
0198                                              bfield_view_t>;
0199 
0200         // Run the track fitting
0201         ::alpaka::exec<Acc>(details::get_queue(queue()), workDiv,
0202                             kernels::fit_backward<fitter_t>{}, config,
0203                             payload.payload, payload.get_tpayload<fitter_t>());
0204       });
0205 }
0206 
0207 }  // namespace traccc::alpaka