Back to home page

EIC code displayed by LXR

 
 

    


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

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 // Project include(s).
0009 #include "traccc/bfield/construct_const_bfield.hpp"
0010 #include "traccc/fitting/kalman_fitting_algorithm.hpp"
0011 #include "traccc/io/read_detector.hpp"
0012 #include "traccc/io/utils.hpp"
0013 #include "traccc/resolution/fitting_performance_writer.hpp"
0014 #include "traccc/simulation/event_generators.hpp"
0015 #include "traccc/simulation/measurement_smearer.hpp"
0016 #include "traccc/simulation/simulator.hpp"
0017 #include "traccc/simulation/smearing_writer.hpp"
0018 #include "traccc/utils/ranges.hpp"
0019 #include "traccc/utils/seed_generator.hpp"
0020 
0021 // Test include(s).
0022 #include "tests/kalman_fitting_wire_chamber_test.hpp"
0023 
0024 // VecMem include(s).
0025 #include <vecmem/memory/host_memory_resource.hpp>
0026 #include <vecmem/utils/copy.hpp>
0027 
0028 // GTest include(s).
0029 #include <gtest/gtest.h>
0030 
0031 // System include(s).
0032 #include <filesystem>
0033 #include <string>
0034 
0035 using namespace traccc;
0036 
0037 TEST_P(KalmanFittingWireChamberTests, Run) {
0038   // Get the parameters
0039   const std::string name = std::get<0>(GetParam());
0040   const traccc::pdg_particle<scalar> ptc = std::get<6>(GetParam());
0041   const unsigned int n_truth_tracks = std::get<7>(GetParam());
0042   const unsigned int n_events = std::get<8>(GetParam());
0043   const bool random_charge = std::get<9>(GetParam());
0044 
0045   // Performance writer
0046   traccc::fitting_performance_writer::config fit_writer_cfg;
0047   fit_writer_cfg.file_path = "performance_track_fitting_" + name + ".root";
0048   traccc::fitting_performance_writer fit_performance_writer(
0049       fit_writer_cfg, traccc::getDefaultLogger("FittingPerformanceWriter",
0050                                                traccc::Logging::Level::INFO));
0051 
0052   /*****************************
0053    * Build a drift chamber
0054    *****************************/
0055   // Memory resources used by the application.
0056   vecmem::host_memory_resource host_mr;
0057   // Copy obejct
0058   vecmem::copy copy;
0059 
0060   // Read back detector file
0061   const std::string path = name + "/";
0062   traccc::host_detector detector;
0063   traccc::io::read_detector(
0064       detector, host_mr,
0065       std::filesystem::absolute(
0066           std::filesystem::path(path + "wire_chamber_geometry.json"))
0067           .native(),
0068       std::filesystem::absolute(
0069           std::filesystem::path(path +
0070                                 "wire_chamber_homogeneous_material.json"))
0071           .native(),
0072       std::filesystem::absolute(
0073           std::filesystem::path(path + "wire_chamber_surface_grids.json"))
0074           .native());
0075   const auto field = traccc::construct_const_bfield(B);
0076 
0077   /***************************
0078    * Generate simulation data
0079    ***************************/
0080 
0081   // Track generator
0082   using generator_type =
0083       detray::random_track_generator<traccc::free_track_parameters<>,
0084                                      uniform_gen_t>;
0085   generator_type::configuration gen_cfg{};
0086   gen_cfg.n_tracks(n_truth_tracks);
0087   gen_cfg.origin(std::get<1>(GetParam()));
0088   gen_cfg.origin_stddev(std::get<2>(GetParam()));
0089   gen_cfg.phi_range(std::get<5>(GetParam()));
0090   gen_cfg.eta_range(std::get<4>(GetParam()));
0091   gen_cfg.mom_range(std::get<3>(GetParam()));
0092   gen_cfg.randomize_charge(random_charge);
0093   gen_cfg.seed(42);
0094   generator_type generator(gen_cfg);
0095 
0096   // Smearing value for measurements
0097   traccc::measurement_smearer<traccc::default_algebra> meas_smearer(
0098       smearing[0], smearing[1]);
0099 
0100   using writer_type = traccc::smearing_writer<
0101       traccc::measurement_smearer<traccc::default_algebra>>;
0102 
0103   typename writer_type::config smearer_writer_cfg{meas_smearer};
0104   traccc::seed_generator<host_detector_type>::config seed_cfg{};
0105   seed_cfg.initial_sigmas = stddevs;
0106 
0107   // Run simulator
0108   const std::string full_path = io::data_directory() + path;
0109   std::filesystem::create_directories(full_path);
0110   auto sim = traccc::simulator<host_detector_type, b_field_t, generator_type,
0111                                writer_type>(
0112       ptc, n_events, detector.as<detector_traits>(),
0113       field.as_field<traccc::const_bfield_backend_t<traccc::scalar>>(),
0114       std::move(generator), std::move(smearer_writer_cfg), full_path);
0115 
0116   sim.get_config().propagation.navigation.search_window = search_window;
0117 
0118   sim.run();
0119 
0120   /***************
0121    * Run fitting
0122    ***************/
0123 
0124   // Seed generator
0125   seed_generator<host_detector_type> sg(detector.as<detector_traits>(),
0126                                         seed_cfg);
0127 
0128   // Fitting algorithm object
0129   traccc::fitting_config fit_cfg;
0130   fit_cfg.propagation.navigation.intersection.min_mask_tolerance =
0131       static_cast<float>(mask_tolerance);
0132   fit_cfg.propagation.navigation.search_window = search_window;
0133   // TODO: Disable until overlaps are handled correctly
0134   fit_cfg.propagation.navigation.estimate_scattering_noise = false;
0135   fit_cfg.ptc_hypothesis = ptc;
0136   fit_cfg.min_pT = 100.f * traccc::unit<float>::MeV;
0137   traccc::host::kalman_fitting_algorithm fitting(fit_cfg, host_mr, copy);
0138 
0139   // Iterate over events
0140   for (std::size_t i_evt = 0; i_evt < n_events; i_evt++) {
0141     // Event map
0142     traccc::event_data evt_data(path, i_evt, host_mr);
0143     // Truth Track Candidates
0144     traccc::edm::measurement_collection::host measurements(host_mr);
0145     traccc::edm::track_container<traccc::default_algebra>::host
0146         track_candidates{host_mr};
0147     evt_data.generate_truth_candidates(track_candidates, measurements, sg,
0148                                        host_mr);
0149     track_candidates.measurements = vecmem::get_data(measurements);
0150 
0151     // n_trakcs = 100
0152     ASSERT_EQ(track_candidates.tracks.size(), n_truth_tracks);
0153 
0154     // Run fitting
0155     auto track_states = fitting(
0156         detector, field,
0157         traccc::edm::track_container<traccc::default_algebra>::const_data(
0158             track_candidates));
0159 
0160     // Iterator over tracks
0161     const std::size_t n_tracks = track_states.tracks.size();
0162 
0163     ASSERT_GE(static_cast<float>(n_tracks),
0164               0.98 * static_cast<float>(n_truth_tracks));
0165 
0166     const std::size_t n_fitted_tracks =
0167         count_successfully_fitted_tracks(track_states.tracks);
0168     ASSERT_GE(static_cast<float>(n_fitted_tracks),
0169               0.92f * static_cast<float>(n_truth_tracks));
0170 
0171     for (std::size_t i_trk = 0; i_trk < n_tracks; i_trk++) {
0172       // Some fits fail. The results of those cannot be reasonably tested.
0173       if (track_states.tracks.at(i_trk).fit_outcome() !=
0174           traccc::track_fit_outcome::SUCCESS) {
0175         continue;
0176       }
0177 
0178       consistency_tests(track_states.tracks.at(i_trk), track_states.states);
0179 
0180       ndf_tests(track_states.tracks.at(i_trk), track_states.states,
0181                 measurements);
0182 
0183       fit_performance_writer.write(track_states.tracks.at(i_trk),
0184                                    track_states.states, measurements,
0185                                    detector.as<detector_traits>(), evt_data);
0186     }
0187   }
0188 
0189   fit_performance_writer.finalize();
0190 
0191   /********************
0192    * Pull value test
0193    ********************/
0194 
0195   static const std::vector<std::string> pull_names{
0196       "pull_d0", "pull_z0", "pull_phi", "pull_theta", "pull_qop"};
0197   pull_value_tests(fit_writer_cfg.file_path, pull_names);
0198 
0199   /********************
0200    * P-value test
0201    ********************/
0202 
0203   //@TODO: Develop an extension of KF-based fitter (e.g. Deterministic
0204   // Annealing Filter) to resolve left-right ambiguity and pass the p-value
0205   // test
0206   // p_value_tests(fit_writer_cfg.file_path);
0207 
0208   /********************
0209    * Success rate test
0210    ********************/
0211 
0212   scalar success_rate = static_cast<scalar>(n_success) /
0213                         static_cast<scalar>(n_truth_tracks * n_events);
0214 
0215   // TODO: Raise back to 95%
0216   ASSERT_GE(success_rate, 0.93f);
0217   ASSERT_LE(success_rate, 1.00f);
0218 }
0219 
0220 INSTANTIATE_TEST_SUITE_P(
0221     KalmanFitWireChamberValidation0, KalmanFittingWireChamberTests,
0222     ::testing::Values(std::make_tuple(
0223         "wire_2_GeV_muon", std::array<scalar, 3u>{0.f, 0.f, 0.f},
0224         std::array<scalar, 3u>{0.f, 0.f, 0.f}, std::array<scalar, 2u>{2.f, 2.f},
0225         std::array<scalar, 2u>{-1.f, 1.f},
0226         std::array<scalar, 2u>{-traccc::constant<scalar>::pi,
0227                                traccc::constant<scalar>::pi},
0228         traccc::muon<scalar>(), 100, 100, false)));
0229 
0230 // @TODO: Make full eta range work
0231 INSTANTIATE_TEST_SUITE_P(
0232     KalmanFitWireChamberValidation1, KalmanFittingWireChamberTests,
0233     ::testing::Values(std::make_tuple(
0234         "wire_10_GeV_muon", std::array<scalar, 3u>{0.f, 0.f, 0.f},
0235         std::array<scalar, 3u>{0.f, 0.f, 0.f},
0236         std::array<scalar, 2u>{10.f, 10.f}, std::array<scalar, 2u>{-0.3f, 0.3f},
0237         std::array<scalar, 2u>{-traccc::constant<scalar>::pi,
0238                                traccc::constant<scalar>::pi},
0239         traccc::muon<scalar>(), 100, 100, false)));
0240 
0241 // @TODO: Make full eta range work
0242 INSTANTIATE_TEST_SUITE_P(
0243     KalmanFitWireChamberValidation2, KalmanFittingWireChamberTests,
0244     ::testing::Values(std::make_tuple(
0245         "wire_100_GeV_muon", std::array<scalar, 3u>{0.f, 0.f, 0.f},
0246         std::array<scalar, 3u>{0.f, 0.f, 0.f},
0247         std::array<scalar, 2u>{100.f, 100.f},
0248         std::array<scalar, 2u>{-0.4f, 0.4f},
0249         std::array<scalar, 2u>{-traccc::constant<scalar>::pi,
0250                                traccc::constant<scalar>::pi},
0251         traccc::muon<scalar>(), 100, 100, false)));
0252 
0253 INSTANTIATE_TEST_SUITE_P(
0254     KalmanFitWireChamberValidation3, KalmanFittingWireChamberTests,
0255     ::testing::Values(std::make_tuple(
0256         "wire_2_GeV_anti_muon", std::array<scalar, 3u>{0.f, 0.f, 0.f},
0257         std::array<scalar, 3u>{0.f, 0.f, 0.f}, std::array<scalar, 2u>{2.f, 2.f},
0258         std::array<scalar, 2u>{-1.f, 1.f},
0259         std::array<scalar, 2u>{-traccc::constant<scalar>::pi,
0260                                traccc::constant<scalar>::pi},
0261         traccc::antimuon<scalar>(), 100, 100, false)));
0262 
0263 INSTANTIATE_TEST_SUITE_P(
0264     KalmanFitWireChamberValidation4, KalmanFittingWireChamberTests,
0265     ::testing::Values(std::make_tuple(
0266         "wire_2_GeV_random_charge", std::array<scalar, 3u>{0.f, 0.f, 0.f},
0267         std::array<scalar, 3u>{0.f, 0.f, 0.f}, std::array<scalar, 2u>{2.f, 2.f},
0268         std::array<scalar, 2u>{-1.f, 1.f},
0269         std::array<scalar, 2u>{-traccc::constant<scalar>::pi,
0270                                traccc::constant<scalar>::pi},
0271         traccc::antimuon<scalar>(), 100, 100, true)));