File indexing completed on 2026-07-26 08:22:04
0001
0002
0003
0004
0005
0006
0007
0008
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
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
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
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
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
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 }
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
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
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
0112 ::alpaka::exec<Acc>(details::get_queue(queue()), workDiv,
0113 kernels::fill_fitting_sort_keys{}, tracks,
0114 track_sort_keys, track_indices);
0115
0116
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
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
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
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
0159 const unsigned int n_tracks = payload.payload.tracks.tracks.capacity();
0160 assert(n_tracks == copy().get_size(payload.payload.tracks.tracks));
0161
0162
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
0168 using fitter_t =
0169 traccc::details::kalman_fitter_t<typename detector_traits_t::device,
0170 bfield_view_t>;
0171
0172
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
0187 const unsigned int n_tracks = payload.payload.tracks.tracks.capacity();
0188 assert(n_tracks == copy().get_size(payload.payload.tracks.tracks));
0189
0190
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
0196 using fitter_t =
0197 traccc::details::kalman_fitter_t<typename detector_traits_t::device,
0198 bfield_view_t>;
0199
0200
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 }