File indexing completed on 2026-07-26 08:22:23
0001
0002
0003
0004
0005
0006
0007
0008
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
0015 #include <vecmem/memory/host_memory_resource.hpp>
0016
0017
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 }
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
0067
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
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
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
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
0171
0172
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
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
0190
0191
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
0217
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
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
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
0300 ASSERT_EQ(last, pattern.end());
0301 pattern.erase(last, pattern.end());
0302
0303
0304 ASSERT_EQ(pattern.size(), track_length);
0305
0306
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
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 }