Back to home page

EIC code displayed by LXR

 
 

    


Warning, /acts/Traccc/device/sycl/src/fitting/kalman_fitting_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) 2022-2026 CERN for the benefit of the ACTS project
0004  *
0005  * Mozilla Public License Version 2.0
0006  */
0007 
0008 // SYCL include(s).
0009 #include <sycl/sycl.hpp>
0010 
0011 // Local include(s).
0012 #include "../utils/calculate1DimNdRange.hpp"
0013 #include "../utils/detector_types.hpp"
0014 #include "../utils/get_queue.hpp"
0015 #include "../utils/global_index.hpp"
0016 #include "../utils/magnetic_field_types.hpp"
0017 #include "../utils/oneDPL.hpp"
0018 #include "traccc/sycl/fitting/kalman_fitting_algorithm.hpp"
0019 
0020 // Project include(s).
0021 #include "traccc/bfield/magnetic_field_types.hpp"
0022 #include "traccc/fitting/device/fill_fitting_sort_keys.hpp"
0023 #include "traccc/fitting/device/fit_backward.hpp"
0024 #include "traccc/fitting/device/fit_forward.hpp"
0025 #include "traccc/fitting/device/fit_prelude.hpp"
0026 #include "traccc/utils/detector_buffer_bfield_visitor.hpp"
0027 
0028 namespace traccc::sycl {
0029 namespace kernels {
0030 
0031 struct fill_fitting_sort_keys;
0032 struct fit_prelude;
0033 template <typename detector_t, typename bfield_t>
0034 struct fit_forward;
0035 template <typename detector_t, typename bfield_t>
0036 struct fit_backward;
0037 
0038 }  // namespace kernels
0039 
0040 void kalman_fitting_algorithm::prepare_track_fit_order(
0041     const edm::track_collection<default_algebra>::const_view& tracks,
0042     vecmem::data::vector_view<device::sort_key>& track_sort_keys,
0043     vecmem::data::vector_view<unsigned int>& track_indices) const {
0044   // Get the number of tracks.
0045   const unsigned int n_tracks = tracks.capacity();
0046   assert(n_tracks == copy().get_size(tracks));
0047   assert(n_tracks == track_indices.capacity());
0048   assert(track_indices.size_ptr() == nullptr);
0049 
0050   // Launch parameters for the kernel.
0051   static constexpr unsigned int localSize = 64;
0052   ::sycl::nd_range<1> range =
0053       details::calculate1DimNdRange(n_tracks, localSize);
0054 
0055   // Fill the keys and indices buffers.
0056   details::get_queue(queue()).submit([&](::sycl::handler& h) {
0057     h.parallel_for<kernels::fill_fitting_sort_keys>(
0058         range,
0059         [tracks, track_sort_keys, track_indices](::sycl::nd_item<1> item) {
0060           device::fill_fitting_sort_keys(details::global_index(item), tracks,
0061                                          track_sort_keys, track_indices);
0062         });
0063   });
0064 
0065   // Sort the key to get the sorted parameter ids
0066   vecmem::device_vector<device::sort_key> keys_device(track_sort_keys);
0067   vecmem::device_vector<unsigned int> track_indices_device(track_indices);
0068   oneapi::dpl::sort_by_key(
0069       oneapi::dpl::execution::device_policy{details::get_queue(queue())},
0070       keys_device.begin(), keys_device.end(), track_indices_device.begin());
0071 }
0072 
0073 void kalman_fitting_algorithm::fit_prelude_kernel(
0074     const device::fit_prelude_payload& payload) const {
0075   // Get the number of tracks.
0076   const unsigned int n_tracks = payload.input_tracks.tracks.capacity();
0077   assert(n_tracks == copy().get_size(payload.input_tracks.tracks));
0078   assert(n_tracks == payload.track_indices.capacity());
0079   assert(payload.track_indices.size_ptr() == nullptr);
0080   assert(n_tracks == copy().get_size(payload.output_tracks.tracks));
0081 
0082   // Launch parameters for the kernel.
0083   static constexpr unsigned int localSize = 64;
0084   ::sycl::nd_range<1> range =
0085       details::calculate1DimNdRange(n_tracks, localSize);
0086 
0087   // Run the fitting, using the sorted parameter IDs.
0088   details::get_queue(queue()).submit([&](::sycl::handler& h) {
0089     h.parallel_for<kernels::fit_prelude>(
0090         range, [payload](::sycl::nd_item<1> item) {
0091           device::fit_prelude(details::global_index(item), payload);
0092         });
0093   });
0094 }
0095 
0096 auto kalman_fitting_algorithm::prepare_fit_payload(
0097     const detector_buffer& det, const magnetic_field& field,
0098     const std::vector<unsigned int>& n_surfaces,
0099     const device::fit_payload& payload) const -> fit_payload {
0100   return prepare_fit_payload_helper<detector_type_list,
0101                                     sycl::bfield_type_list<scalar>>(
0102       det, field, n_surfaces, payload);
0103 }
0104 
0105 void kalman_fitting_algorithm::fit_forward_kernel(
0106     const fitting_config& config, const fit_payload& payload) const {
0107   return detector_buffer_magnetic_field_visitor<detector_type_list,
0108                                                 sycl::bfield_type_list<scalar>>(
0109       payload.detector, payload.field,
0110       [&]<typename detector_traits_t, typename bfield_view_t>(
0111           const typename detector_traits_t::view&, const bfield_view_t&) {
0112         // Get the number of tracks.
0113         const unsigned int n_tracks = payload.payload.tracks.tracks.capacity();
0114         assert(n_tracks == copy().get_size(payload.payload.tracks.tracks));
0115 
0116         // Launch parameters for the kernel.
0117         static constexpr unsigned int localSize = 64;
0118         ::sycl::nd_range<1> range =
0119             details::calculate1DimNdRange(n_tracks, localSize);
0120 
0121         // Fitter type to use.
0122         using fitter_t =
0123             traccc::details::kalman_fitter_t<typename detector_traits_t::device,
0124                                              bfield_view_t>;
0125 
0126         // Run the track fitting
0127         details::get_queue(queue()).submit([&](::sycl::handler& h) {
0128           h.parallel_for<kernels::fit_forward<
0129               detector_tag_selector_t<detector_traits_t>,
0130               bfield_tag_selector_t<typename bfield_view_t::backend_t>>>(
0131               range, [config, payload = payload.payload,
0132                       tpayload = payload.get_tpayload<fitter_t>()](
0133                          ::sycl::nd_item<1> item) {
0134                 device::fit_forward<fitter_t>(details::global_index(item),
0135                                               config, payload, *tpayload);
0136               });
0137         });
0138       });
0139 }
0140 
0141 void kalman_fitting_algorithm::fit_backward_kernel(
0142     const fitting_config& config, const fit_payload& payload) const {
0143   return detector_buffer_magnetic_field_visitor<detector_type_list,
0144                                                 sycl::bfield_type_list<scalar>>(
0145       payload.detector, payload.field,
0146       [&]<typename detector_traits_t, typename bfield_view_t>(
0147           const typename detector_traits_t::view&, const bfield_view_t&) {
0148         // Get the number of tracks.
0149         const unsigned int n_tracks = payload.payload.tracks.tracks.capacity();
0150         assert(n_tracks == copy().get_size(payload.payload.tracks.tracks));
0151 
0152         // Launch parameters for the kernel.
0153         static constexpr unsigned int localSize = 64;
0154         ::sycl::nd_range<1> range =
0155             details::calculate1DimNdRange(n_tracks, localSize);
0156 
0157         // Fitter type to use.
0158         using fitter_t =
0159             traccc::details::kalman_fitter_t<typename detector_traits_t::device,
0160                                              bfield_view_t>;
0161 
0162         // Run the track fitting
0163         details::get_queue(queue()).submit([&](::sycl::handler& h) {
0164           h.parallel_for<kernels::fit_backward<
0165               detector_tag_selector_t<detector_traits_t>,
0166               bfield_tag_selector_t<typename bfield_view_t::backend_t>>>(
0167               range, [config, payload = payload.payload,
0168                       tpayload = payload.get_tpayload<fitter_t>()](
0169                          ::sycl::nd_item<1> item) {
0170                 device::fit_backward<fitter_t>(details::global_index(item),
0171                                                config, payload, *tpayload);
0172               });
0173         });
0174       });
0175 }
0176 
0177 }  // namespace traccc::sycl