Back to home page

EIC code displayed by LXR

 
 

    


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

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 // Project include(s).
0009 #include "traccc/ambiguity_resolution/greedy_ambiguity_resolution_algorithm.hpp"
0010 #include "traccc/ambiguity_resolution/legacy/greedy_ambiguity_resolution_algorithm.hpp"
0011 #include "traccc/edm/track_container.hpp"
0012 #include "traccc/utils/memory_resource.hpp"
0013 
0014 // VecMem include(s).
0015 #include <vecmem/memory/host_memory_resource.hpp>
0016 
0017 // GTest include(s).
0018 #include <chrono>
0019 #include <random>
0020 
0021 #include <gtest/gtest.h>
0022 
0023 using namespace traccc;
0024 
0025 namespace {
0026 vecmem::host_memory_resource host_mr;
0027 }  // namespace
0028 
0029 void fill_measurements(edm::measurement_collection::host& measurements,
0030                        const measurement_id_type max_meas_id) {
0031   measurements.reserve(max_meas_id + 1);
0032   for (measurement_id_type i = 0; i <= max_meas_id; i++) {
0033     measurements.push_back({});
0034     measurements.at(measurements.size() - 1).identifier() = i;
0035   }
0036 }
0037 
0038 void fill_pattern(edm::track_container<default_algebra>::host& track_candidates,
0039                   const traccc::scalar pval,
0040                   const std::vector<measurement_id_type>& pattern) {
0041   track_candidates.tracks.resize(track_candidates.tracks.size() + 1u);
0042   track_candidates.tracks.pval().back() = pval;
0043 
0044   edm::measurement_collection::const_device measurements{
0045       track_candidates.measurements};
0046 
0047   for (const auto& meas_id : pattern) {
0048     const auto meas_iter =
0049         std::lower_bound(measurements.identifier().begin(),
0050                          measurements.identifier().end(), meas_id);
0051 
0052     const auto meas_idx =
0053         std::distance(measurements.identifier().begin(), meas_iter);
0054     track_candidates.tracks.constituent_links().back().push_back(
0055         {edm::track_constituent_link::measurement,
0056          static_cast<measurement_id_type>(meas_idx)});
0057   }
0058 }
0059 
0060 std::vector<std::size_t> get_pattern(
0061     const edm::track_container<default_algebra>::host& track_candidates,
0062     const std::size_t idx) {
0063   edm::measurement_collection::const_device measurements{
0064       track_candidates.measurements};
0065   std::vector<std::size_t> ret;
0066   // A const reference would be fine here. But GCC fears that that would lead
0067   // to a dangling reference...
0068   const auto meas_links = track_candidates.tracks.at(idx).constituent_links();
0069   for (const auto& [type, meas_idx] : meas_links) {
0070     assert(type == edm::track_constituent_link::measurement);
0071     ret.push_back(measurements.at(meas_idx).identifier());
0072   }
0073 
0074   return ret;
0075 }
0076 
0077 TEST(AmbiguitySolverTests, GreedyResolverTest0) {
0078   edm::measurement_collection::host measurements{host_mr};
0079   fill_measurements(measurements, 100);
0080 
0081   edm::track_container<default_algebra>::host trk_cands{
0082       host_mr, vecmem::get_data(measurements)};
0083   fill_pattern(trk_cands, 0.23f, {5, 1, 11, 3});
0084   fill_pattern(trk_cands, 0.85f, {12, 10, 9, 8, 7, 6});
0085   fill_pattern(trk_cands, 0.42f, {4, 2, 13});
0086 
0087   traccc::host::greedy_ambiguity_resolution_algorithm::config_type
0088       resolution_config;
0089 
0090   traccc::host::greedy_ambiguity_resolution_algorithm resolution_alg(
0091       resolution_config, host_mr);
0092   {
0093     resolution_alg.get_config().min_meas_per_track = 3;
0094     auto res_trk_cands = resolution_alg(
0095         traccc::edm::track_container<traccc::default_algebra>::const_data(
0096             trk_cands));
0097     // All tracks are accepted as they have more than three measurements
0098     ASSERT_EQ(res_trk_cands.tracks.size(), 3u);
0099     ASSERT_EQ(get_pattern(res_trk_cands, 0),
0100               std::vector<std::size_t>({5, 1, 11, 3}));
0101     ASSERT_EQ(get_pattern(res_trk_cands, 1),
0102               std::vector<std::size_t>({12, 10, 9, 8, 7, 6}));
0103     ASSERT_EQ(get_pattern(res_trk_cands, 2),
0104               std::vector<std::size_t>({4, 2, 13}));
0105   }
0106 
0107   {
0108     resolution_alg.get_config().min_meas_per_track = 5;
0109     auto res_trk_cands = resolution_alg(
0110         traccc::edm::track_container<traccc::default_algebra>::const_data(
0111             trk_cands));
0112     // Only the second track with six measurements is accepted
0113     ASSERT_EQ(res_trk_cands.tracks.size(), 1u);
0114     ASSERT_EQ(get_pattern(res_trk_cands, 0),
0115               std::vector<std::size_t>({12, 10, 9, 8, 7, 6}));
0116   }
0117 
0118   /*******************
0119    * Legacy algorithm
0120    * *****************/
0121 
0122   {
0123     traccc::legacy::greedy_ambiguity_resolution_algorithm::config_t legacy_cfg;
0124     traccc::legacy::greedy_ambiguity_resolution_algorithm legacy_resolution_alg(
0125         legacy_cfg, host_mr);
0126 
0127     legacy_resolution_alg.get_config().n_measurements_min = 3;
0128     auto res_trk_cands = legacy_resolution_alg(trk_cands);
0129     ASSERT_EQ(res_trk_cands.tracks.size(), 3u);
0130     ASSERT_EQ(get_pattern(res_trk_cands, 0),
0131               std::vector<std::size_t>({5, 1, 11, 3}));
0132     ASSERT_EQ(get_pattern(res_trk_cands, 1),
0133               std::vector<std::size_t>({12, 10, 9, 8, 7, 6}));
0134     ASSERT_EQ(get_pattern(res_trk_cands, 2),
0135               std::vector<std::size_t>({4, 2, 13}));
0136   }
0137 
0138   {
0139     traccc::legacy::greedy_ambiguity_resolution_algorithm::config_t legacy_cfg;
0140     traccc::legacy::greedy_ambiguity_resolution_algorithm legacy_resolution_alg(
0141         legacy_cfg, host_mr);
0142 
0143     legacy_resolution_alg.get_config().n_measurements_min = 5;
0144     auto res_trk_cands = legacy_resolution_alg(trk_cands);
0145     ASSERT_EQ(res_trk_cands.tracks.size(), 1u);
0146     ASSERT_EQ(get_pattern(res_trk_cands, 0),
0147               std::vector<std::size_t>({12, 10, 9, 8, 7, 6}));
0148   }
0149 }
0150 
0151 TEST(AmbiguitySolverTests, GreedyResolverTest1) {
0152   edm::measurement_collection::host measurements{host_mr};
0153   fill_measurements(measurements, 100);
0154 
0155   edm::track_container<default_algebra>::host trk_cands{
0156       host_mr, vecmem::get_data(measurements)};
0157   fill_pattern(trk_cands, 0.12f, {5, 14, 1, 11, 18, 16, 3});
0158   fill_pattern(trk_cands, 0.53f, {3, 6, 5, 13});
0159 
0160   traccc::host::greedy_ambiguity_resolution_algorithm::config_type
0161       resolution_config;
0162   traccc::host::greedy_ambiguity_resolution_algorithm resolution_alg(
0163       resolution_config, host_mr);
0164   {
0165     auto res_trk_cands = resolution_alg(
0166         traccc::edm::track_container<traccc::default_algebra>::const_data(
0167             trk_cands));
0168     ASSERT_EQ(res_trk_cands.tracks.size(), 1u);
0169 
0170     // The first track is selected over the second one as its relative
0171     // shared measurement (2/7) is lower than the one of the second track
0172     // (2/4)
0173     ASSERT_EQ(get_pattern(res_trk_cands, 0),
0174               std::vector<std::size_t>({5, 14, 1, 11, 18, 16, 3}));
0175   }
0176 
0177   /*******************
0178    * Legacy algorithm
0179    * *****************/
0180 
0181   {
0182     traccc::legacy::greedy_ambiguity_resolution_algorithm::config_t legacy_cfg;
0183     traccc::legacy::greedy_ambiguity_resolution_algorithm legacy_resolution_alg(
0184         legacy_cfg, host_mr);
0185 
0186     auto res_trk_cands = legacy_resolution_alg(trk_cands);
0187     ASSERT_EQ(res_trk_cands.tracks.size(), 1u);
0188 
0189     // The first track is selected over the second one as its relative
0190     // shared measurement (2/7) is lower than the one of the second track
0191     // (2/4)
0192     ASSERT_EQ(get_pattern(res_trk_cands, 0),
0193               std::vector<std::size_t>({5, 14, 1, 11, 18, 16, 3}));
0194   }
0195 }
0196 
0197 TEST(AmbiguitySolverTests, GreedyResolverTest2) {
0198   edm::measurement_collection::host measurements{host_mr};
0199   fill_measurements(measurements, 100);
0200 
0201   edm::track_container<default_algebra>::host trk_cands{
0202       host_mr, vecmem::get_data(measurements)};
0203   fill_pattern(trk_cands, 0.8f, {1, 3, 5, 11});
0204   fill_pattern(trk_cands, 0.9f, {3, 5, 6, 13});
0205 
0206   traccc::host::greedy_ambiguity_resolution_algorithm::config_type
0207       resolution_config;
0208   traccc::host::greedy_ambiguity_resolution_algorithm resolution_alg(
0209       resolution_config, host_mr);
0210   {
0211     auto res_trk_cands = resolution_alg(
0212         traccc::edm::track_container<traccc::default_algebra>::const_data(
0213             trk_cands));
0214     ASSERT_EQ(res_trk_cands.tracks.size(), 1u);
0215 
0216     // The second track is selected over the first one as their relative
0217     // shared measurement (2/4) is the same but its p-value is higher
0218     ASSERT_EQ(get_pattern(res_trk_cands, 0),
0219               std::vector<std::size_t>({3, 5, 6, 13}));
0220   }
0221 }
0222 
0223 TEST(AmbiguitySolverTests, GreedyResolverTest3) {
0224   edm::measurement_collection::host measurements{host_mr};
0225   fill_measurements(measurements, 100);
0226 
0227   edm::track_container<default_algebra>::host trk_cands{
0228       host_mr, vecmem::get_data(measurements)};
0229   fill_pattern(trk_cands, 0.2f, {5, 1, 11, 3});
0230   fill_pattern(trk_cands, 0.5f, {6, 2});
0231   fill_pattern(trk_cands, 0.4f, {3, 21, 12, 6, 19, 14});
0232   fill_pattern(trk_cands, 0.1f, {13, 16, 2, 7, 11});
0233   fill_pattern(trk_cands, 0.3f, {1, 7, 8});
0234   fill_pattern(trk_cands, 0.6f, {1, 3, 11, 22});
0235 
0236   traccc::host::greedy_ambiguity_resolution_algorithm::config_type
0237       resolution_config;
0238   traccc::host::greedy_ambiguity_resolution_algorithm resolution_alg(
0239       resolution_config, host_mr);
0240 
0241   {
0242     auto res_trk_cands = resolution_alg(
0243         traccc::edm::track_container<traccc::default_algebra>::const_data(
0244             trk_cands));
0245     ASSERT_EQ(res_trk_cands.tracks.size(), 2u);
0246 
0247     ASSERT_EQ(get_pattern(res_trk_cands, 0),
0248               std::vector<std::size_t>({3, 21, 12, 6, 19, 14}));
0249     ASSERT_EQ(get_pattern(res_trk_cands, 1),
0250               std::vector<std::size_t>({13, 16, 2, 7, 11}));
0251   }
0252 
0253   // Legacy algorithm
0254   traccc::legacy::greedy_ambiguity_resolution_algorithm::config_t legacy_cfg;
0255   traccc::legacy::greedy_ambiguity_resolution_algorithm legacy_resolution_alg(
0256       legacy_cfg, host_mr);
0257 
0258   {
0259     auto res_trk_cands = legacy_resolution_alg(trk_cands);
0260     ASSERT_EQ(res_trk_cands.tracks.size(), 2u);
0261 
0262     ASSERT_EQ(get_pattern(res_trk_cands, 0),
0263               std::vector<std::size_t>({3, 21, 12, 6, 19, 14}));
0264     ASSERT_EQ(get_pattern(res_trk_cands, 1),
0265               std::vector<std::size_t>({13, 16, 2, 7, 11}));
0266   }
0267 }
0268 
0269 // Comparison to the legacy algorithm.
0270 TEST(AmbiguitySolverTests, GreedyResolverTest4) {
0271   edm::measurement_collection::host measurements{host_mr};
0272   const measurement_id_type max_meas_id = 10000;
0273   fill_measurements(measurements, max_meas_id);
0274 
0275   edm::track_container<default_algebra>::host trk_cands{
0276       host_mr, vecmem::get_data(measurements)};
0277 
0278   std::mt19937 gen(42);
0279 
0280   for (std::size_t i = 0; i < 10000u; i++) {
0281     std::uniform_int_distribution<std::size_t> track_length_dist(1, 20);
0282     std::uniform_int_distribution<measurement_id_type> meas_id_dist(
0283         0, max_meas_id);
0284     std::uniform_real_distribution<traccc::scalar> pval_dist(0.0f, 1.0f);
0285 
0286     const std::size_t track_length = track_length_dist(gen);
0287     const traccc::scalar pval = pval_dist(gen);
0288     std::vector<measurement_id_type> pattern;
0289     while (pattern.size() < track_length) {
0290       const measurement_id_type meas_id = meas_id_dist(gen);
0291       if (std::find(pattern.begin(), pattern.end(), meas_id) == pattern.end()) {
0292         pattern.push_back(meas_id);
0293       }
0294     }
0295 
0296     std::sort(pattern.begin(), pattern.end());
0297     auto last = std::unique(pattern.begin(), pattern.end());
0298 
0299     // There should not be duplicate
0300     ASSERT_EQ(last, pattern.end());
0301     pattern.erase(last, pattern.end());
0302 
0303     // Make sure that partern size is eqaul to the track length
0304     ASSERT_EQ(pattern.size(), track_length);
0305 
0306     // Fill the pattern
0307     fill_pattern(trk_cands, pval, pattern);
0308   }
0309 
0310   traccc::host::greedy_ambiguity_resolution_algorithm::config_type
0311       resolution_config;
0312   traccc::host::greedy_ambiguity_resolution_algorithm resolution_alg(
0313       resolution_config, host_mr);
0314 
0315   auto start_new = std::chrono::high_resolution_clock::now();
0316 
0317   auto res_trk_cands = resolution_alg(
0318       traccc::edm::track_container<traccc::default_algebra>::const_data(
0319           trk_cands));
0320 
0321   auto end_new = std::chrono::high_resolution_clock::now();
0322   auto duration_new = std::chrono::duration_cast<std::chrono::milliseconds>(
0323       end_new - start_new);
0324   std::cout << " Time for the new method " << duration_new.count() << " ms"
0325             << std::endl;
0326 
0327   // Legacy algorithm
0328   traccc::legacy::greedy_ambiguity_resolution_algorithm::config_t legacy_cfg;
0329   traccc::legacy::greedy_ambiguity_resolution_algorithm legacy_resolution_alg(
0330       legacy_cfg, host_mr);
0331 
0332   auto start_legacy = std::chrono::high_resolution_clock::now();
0333 
0334   auto legacy_res_trk_cands = legacy_resolution_alg(trk_cands);
0335 
0336   auto end_legacy = std::chrono::high_resolution_clock::now();
0337   auto duration_legacy = std::chrono::duration_cast<std::chrono::milliseconds>(
0338       end_legacy - start_legacy);
0339   std::cout << " Time for the legacy method " << duration_legacy.count()
0340             << " ms" << std::endl;
0341 
0342   std::size_t n_res_tracks = res_trk_cands.tracks.size();
0343   ASSERT_EQ(n_res_tracks, legacy_res_trk_cands.tracks.size());
0344   for (std::size_t i = 0; i < n_res_tracks; i++) {
0345     ASSERT_EQ(res_trk_cands.tracks.at(i), legacy_res_trk_cands.tracks.at(i));
0346   }
0347 }
0348 
0349 TEST(AmbiguitySolverTests, GreedyResolverTest5) {
0350   edm::measurement_collection::host measurements{host_mr};
0351   fill_measurements(measurements, 100);
0352 
0353   edm::track_container<default_algebra>::host trk_cands{
0354       host_mr, vecmem::get_data(measurements)};
0355   fill_pattern(trk_cands, 0.2f, {1, 2, 1, 1});
0356   fill_pattern(trk_cands, 0.5f, {3, 2, 1});
0357   fill_pattern(trk_cands, 0.4f, {2, 4, 5, 7, 2});
0358   fill_pattern(trk_cands, 0.1f, {6, 6, 6, 6});
0359 
0360   traccc::host::greedy_ambiguity_resolution_algorithm::config_type
0361       resolution_config;
0362   traccc::host::greedy_ambiguity_resolution_algorithm resolution_alg(
0363       resolution_config, host_mr);
0364 
0365   {
0366     auto res_trk_cands = resolution_alg(
0367         traccc::edm::track_container<traccc::default_algebra>::const_data(
0368             trk_cands));
0369     ASSERT_EQ(res_trk_cands.tracks.size(), 2u);
0370 
0371     ASSERT_EQ(get_pattern(res_trk_cands, 0),
0372               std::vector<std::size_t>({3, 2, 1}));
0373     ASSERT_EQ(get_pattern(res_trk_cands, 1),
0374               std::vector<std::size_t>({6, 6, 6, 6}));
0375   }
0376 }
0377 
0378 TEST(AmbiguitySolverTests, GreedyResolverTest6) {
0379   edm::measurement_collection::host measurements{host_mr};
0380   fill_measurements(measurements, 100);
0381 
0382   edm::track_container<default_algebra>::host trk_cands{
0383       host_mr, vecmem::get_data(measurements)};
0384   fill_pattern(trk_cands, 0.2f, {7, 3, 5, 7, 7, 7, 2});
0385   fill_pattern(trk_cands, 0.5f, {2});
0386   fill_pattern(trk_cands, 0.4f, {8, 9, 7, 2, 3, 4, 3, 7});
0387   fill_pattern(trk_cands, 0.1f, {8, 9, 0, 8, 1, 4, 6});
0388   fill_pattern(trk_cands, 0.9f, {10, 3, 2});
0389 
0390   traccc::host::greedy_ambiguity_resolution_algorithm::config_type
0391       resolution_config;
0392   traccc::host::greedy_ambiguity_resolution_algorithm resolution_alg(
0393       resolution_config, host_mr);
0394 
0395   {
0396     auto res_trk_cands = resolution_alg(
0397         traccc::edm::track_container<traccc::default_algebra>::const_data(
0398             trk_cands));
0399     ASSERT_EQ(res_trk_cands.tracks.size(), 2u);
0400 
0401     ASSERT_EQ(get_pattern(res_trk_cands, 0),
0402               std::vector<std::size_t>({7, 3, 5, 7, 7, 7, 2}));
0403     ASSERT_EQ(get_pattern(res_trk_cands, 1),
0404               std::vector<std::size_t>({8, 9, 0, 8, 1, 4, 6}));
0405   }
0406 }