Back to home page

EIC code displayed by LXR

 
 

    


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

0001 /** TRACCC library, part of the ACTS project (R&D line)
0002  *
0003  * (c) 2023-2026 CERN for the benefit of the ACTS project
0004  *
0005  * Mozilla Public License Version 2.0
0006  */
0007 
0008 // Test include(s).
0009 #include "tests/ckf_toy_detector_test.hpp"
0010 #include "traccc/utils/seed_generator.hpp"
0011 
0012 // Project include(s).
0013 #include "traccc/bfield/construct_const_bfield.hpp"
0014 #include "traccc/bfield/magnetic_field_types.hpp"
0015 #include "traccc/finding/combinatorial_kalman_filter_algorithm.hpp"
0016 #include "traccc/io/read_detector.hpp"
0017 #include "traccc/io/read_measurements.hpp"
0018 #include "traccc/io/utils.hpp"
0019 #include "traccc/performance/container_comparator.hpp"
0020 #include "traccc/simulation/event_generators.hpp"
0021 #include "traccc/simulation/simulator.hpp"
0022 #include "traccc/sycl/finding/combinatorial_kalman_filter_algorithm.hpp"
0023 #include "traccc/utils/event_data.hpp"
0024 #include "traccc/utils/ranges.hpp"
0025 
0026 // VecMem include(s).
0027 #include <vecmem/memory/host_memory_resource.hpp>
0028 #include <vecmem/memory/sycl/device_memory_resource.hpp>
0029 #include <vecmem/memory/sycl/shared_memory_resource.hpp>
0030 #include <vecmem/utils/sycl/async_copy.hpp>
0031 
0032 // GTest include(s).
0033 #include <gtest/gtest.h>
0034 
0035 // System include(s).
0036 #include <filesystem>
0037 #include <string>
0038 
0039 namespace traccc {
0040 
0041 TEST_P(CkfToyDetectorTests, Run) {
0042   // Get the parameters
0043   const std::string name = std::get<0>(GetParam());
0044   const traccc::pdg_particle<traccc::scalar> ptc = std::get<6>(GetParam());
0045   const unsigned int n_truth_tracks = std::get<7>(GetParam());
0046   const unsigned int n_events = std::get<8>(GetParam());
0047   const bool random_charge = std::get<9>(GetParam());
0048 
0049   /*****************************
0050    * Build a toy detector
0051    *****************************/
0052 
0053   // SYCL queue.
0054   vecmem::sycl::queue_wrapper vecmem_queue;
0055   traccc::sycl::queue_wrapper traccc_queue{vecmem_queue.queue()};
0056 
0057   // Only run this test on NVIDIA backends.
0058   if (!vecmem_queue.is_cuda()) {
0059     GTEST_SKIP();
0060   }
0061 
0062   // Memory resources used by the application.
0063   vecmem::host_memory_resource host_mr;
0064   vecmem::sycl::device_memory_resource device_mr{vecmem_queue};
0065   traccc::memory_resource mr{device_mr, &host_mr};
0066   vecmem::sycl::shared_memory_resource shared_mr{vecmem_queue};
0067 
0068   // Copy objects
0069   vecmem::sycl::async_copy copy{vecmem_queue};
0070 
0071   // Path to the working directory.
0072   const std::filesystem::path path = std::filesystem::current_path() / name;
0073 
0074   constexpr bool use_material_maps = false;
0075   WriteDetector(use_material_maps, name);
0076 
0077   // Read in the detector geometry that was generated by the test fixture.
0078   detray::io::detector_reader_config reader_cfg{};
0079   reader_cfg.add_file((path / "toy_detector_geometry.json").native())
0080       .add_file((path / "toy_detector_surface_grids.json").native())
0081       .add_file((path / "toy_detector_homogeneous_material.json").native())
0082       .do_check(true);
0083 
0084   auto [io_det, names] =
0085       detray::io::read_detector<traccc::default_detector::host>(host_mr,
0086                                                                 reader_cfg);
0087   traccc::host_detector host_detector{};
0088   host_detector.template set<
0089       traccc::detector_traits<traccc::default_detector::host::metadata>>(
0090       std::move(io_det));
0091 
0092   const traccc::detector_buffer detector_buffer =
0093       traccc::buffer_from_host_detector(host_detector, device_mr, copy);
0094   vecmem_queue.synchronize();
0095 
0096   const auto field = traccc::construct_const_bfield(B);
0097 
0098   /***************************
0099    * Generate simulation data
0100    ***************************/
0101 
0102   // Track generator
0103   using generator_type =
0104       detray::random_track_generator<traccc::free_track_parameters<>,
0105                                      uniform_gen_t>;
0106   generator_type::configuration gen_cfg{};
0107   gen_cfg.n_tracks(n_truth_tracks);
0108   gen_cfg.origin(std::get<1>(GetParam()));
0109   gen_cfg.origin_stddev(std::get<2>(GetParam()));
0110   gen_cfg.phi_range(std::get<5>(GetParam()));
0111   gen_cfg.eta_range(std::get<4>(GetParam()));
0112   gen_cfg.mom_range(std::get<3>(GetParam()));
0113   gen_cfg.randomize_charge(random_charge);
0114   gen_cfg.seed(42);
0115   generator_type generator(gen_cfg);
0116 
0117   // Smearing value for measurements
0118   traccc::measurement_smearer<traccc::default_algebra> meas_smearer(
0119       smearing[0], smearing[1]);
0120 
0121   using writer_type = traccc::smearing_writer<
0122       traccc::measurement_smearer<traccc::default_algebra>>;
0123 
0124   typename writer_type::config smearer_writer_cfg{meas_smearer};
0125   traccc::seed_generator<host_detector_type>::config seed_cfg{};
0126   seed_cfg.initial_sigmas = stddevs;
0127 
0128   // Run simulator
0129   std::filesystem::create_directories(path);
0130   auto sim = traccc::simulator<host_detector_type, b_field_t, generator_type,
0131                                writer_type>(
0132       ptc, n_events, host_detector.as<detector_traits>(),
0133       field.as_field<traccc::const_bfield_backend_t<traccc::scalar>>(),
0134       std::move(generator), std::move(smearer_writer_cfg), path.native());
0135   sim.get_config().propagation.navigation.search_window = search_window;
0136   sim.run();
0137 
0138   /*****************************
0139    * Do the reconstruction
0140    *****************************/
0141 
0142   // Seed generator
0143   seed_generator<host_detector_type> sg(host_detector.as<detector_traits>(),
0144                                         seed_cfg);
0145 
0146   // Finding algorithm configuration
0147   traccc::sycl::combinatorial_kalman_filter_algorithm::config_type cfg;
0148   cfg.ptc_hypothesis = ptc;
0149   cfg.max_num_branches_per_seed = 500;
0150   cfg.max_num_branches_per_surface = 2;
0151   cfg.chi2_max = 10.f;
0152   cfg.propagation.navigation.search_window = search_window;
0153   cfg.run_smoother = smoother_type::e_none;
0154 
0155   // Finding algorithm object
0156   traccc::host::combinatorial_kalman_filter_algorithm host_finding(cfg,
0157                                                                    host_mr);
0158 
0159   // Finding algorithm object
0160   traccc::sycl::combinatorial_kalman_filter_algorithm device_finding{
0161       cfg, mr, copy, traccc_queue};
0162 
0163   // Iterate over events
0164   for (std::size_t i_evt = 0; i_evt < n_events; i_evt++) {
0165     // Truth Track Candidates
0166     traccc::event_data evt_data(path.native(), i_evt, host_mr);
0167 
0168     traccc::edm::measurement_collection::host truth_measurements{host_mr};
0169     traccc::edm::track_container<traccc::default_algebra>::host
0170         truth_track_candidates{host_mr};
0171     evt_data.generate_truth_candidates(truth_track_candidates,
0172                                        truth_measurements, sg, host_mr);
0173     truth_track_candidates.measurements = vecmem::get_data(truth_measurements);
0174 
0175     ASSERT_EQ(truth_track_candidates.tracks.size(), n_truth_tracks);
0176 
0177     // Prepare truth seeds
0178     traccc::bound_track_parameters_collection_types::host seeds(&host_mr);
0179     for (unsigned int i_trk = 0; i_trk < n_truth_tracks; i_trk++) {
0180       seeds.push_back(truth_track_candidates.tracks.at(i_trk).params());
0181     }
0182     ASSERT_EQ(seeds.size(), n_truth_tracks);
0183 
0184     traccc::bound_track_parameters_collection_types::buffer seeds_buffer{
0185         static_cast<unsigned int>(seeds.size()), mr.main};
0186     copy.setup(seeds_buffer)->wait();
0187     copy(vecmem::get_data(seeds), seeds_buffer,
0188          vecmem::copy::type::host_to_device)
0189         ->wait();
0190 
0191     // Read measurements
0192     traccc::edm::measurement_collection::host measurements_per_event{host_mr};
0193     traccc::io::read_measurements(measurements_per_event, i_evt, path.native());
0194 
0195     traccc::edm::measurement_collection::buffer measurements_buffer(
0196         static_cast<unsigned int>(measurements_per_event.size()), mr.main);
0197     copy.setup(measurements_buffer)->wait();
0198     copy(vecmem::get_data(measurements_per_event), measurements_buffer)->wait();
0199 
0200     // Run host finding
0201     auto track_candidates = host_finding(
0202         host_detector, field, vecmem::get_data(measurements_per_event),
0203         vecmem::get_data(seeds));
0204 
0205     // Run device finding
0206     auto track_candidates_sycl_buffer = device_finding(
0207         detector_buffer, field, measurements_buffer, seeds_buffer);
0208 
0209     traccc::edm::track_collection<traccc::default_algebra>::host
0210         track_candidates_sycl{host_mr};
0211     copy(track_candidates_sycl_buffer.tracks, track_candidates_sycl)->wait();
0212 
0213     // Simple check
0214     ASSERT_LE(static_cast<double>(
0215                   std::llabs(static_cast<long>(track_candidates.tracks.size()) -
0216                              static_cast<long>(track_candidates_sycl.size()))) /
0217                   static_cast<double>(track_candidates.tracks.size()),
0218               0.001f)
0219         << "No. tracks (host): " << track_candidates.tracks.size() << "/"
0220         << n_truth_tracks
0221         << "\nNo. tracks (device): " << track_candidates_sycl.size() << "/"
0222         << n_truth_tracks;
0223     ASSERT_GE(track_candidates.tracks.size(), n_truth_tracks);
0224 
0225     // Make sure that the outputs from cpu and cuda CKF are equivalent
0226     unsigned int n_matches = 0u;
0227     for (unsigned int i = 0u; i < track_candidates.tracks.size(); i++) {
0228       traccc::details::is_same_object<traccc::edm::track_collection<
0229           traccc::default_algebra>::host::const_proxy_type>
0230           iso{track_candidates.measurements, track_candidates.measurements,
0231               vecmem::get_data(track_candidates.states),
0232               vecmem::get_data(track_candidates.states),
0233               track_candidates.tracks.at(i)};
0234 
0235       for (unsigned int j = 0u; j < track_candidates_sycl.size(); j++) {
0236         if (iso(track_candidates_sycl.at(j))) {
0237           n_matches++;
0238           break;
0239         }
0240       }
0241     }
0242 
0243     float matching_rate = float(n_matches) / static_cast<float>(std::max(
0244                                                  track_candidates.tracks.size(),
0245                                                  track_candidates_sycl.size()));
0246     EXPECT_GE(matching_rate, 0.998f);
0247   }
0248 }
0249 
0250 INSTANTIATE_TEST_SUITE_P(
0251     SYCLCkfToyDetectorValidation, CkfToyDetectorTests,
0252     ::testing::Values(
0253         std::make_tuple("toy_n_particles_1",
0254                         std::array<scalar, 3u>{0.f, 0.f, 0.f},
0255                         std::array<scalar, 3u>{0.f, 0.f, 0.f},
0256                         std::array<scalar, 2u>{1.f, 100.f},
0257                         std::array<scalar, 2u>{-4.f, 4.f},
0258                         std::array<scalar, 2u>{-traccc::constant<scalar>::pi,
0259                                                traccc::constant<scalar>::pi},
0260                         traccc::muon<scalar>(), 1, 1, false),
0261         std::make_tuple("toy_n_particles_10000",
0262                         std::array<scalar, 3u>{0.f, 0.f, 0.f},
0263                         std::array<scalar, 3u>{0.f, 0.f, 0.f},
0264                         std::array<scalar, 2u>{1.f, 100.f},
0265                         std::array<scalar, 2u>{-4.f, 4.f},
0266                         std::array<scalar, 2u>{-traccc::constant<scalar>::pi,
0267                                                traccc::constant<scalar>::pi},
0268                         traccc::muon<scalar>(), 10000, 1, false),
0269         std::make_tuple("toy_n_particles_10000_random_charge",
0270                         std::array<scalar, 3u>{0.f, 0.f, 0.f},
0271                         std::array<scalar, 3u>{0.f, 0.f, 0.f},
0272                         std::array<scalar, 2u>{1.f, 100.f},
0273                         std::array<scalar, 2u>{-4.f, 4.f},
0274                         std::array<scalar, 2u>{-traccc::constant<scalar>::pi,
0275                                                traccc::constant<scalar>::pi},
0276                         traccc::muon<scalar>(), 10000, 1, true)));
0277 
0278 }  // namespace traccc