ATLAS Offline Software
Loading...
Searching...
No Matches
TracccTritonTool.cxx
Go to the documentation of this file.
1/*
2 Copyright (C) 2002-2026 CERN for the benefit of the ATLAS collaboration
3*/
4
5#include "TracccTritonTool.h"
6
8
9// Framework include(s).
10#include <cmath>
11
12#include <fstream>
13#include <chrono>
14
16
17TracccTritonTool::TracccTritonTool(const std::string& type, const std::string& name, const IInterface* parent)
18 : base_class(type, name, parent) {
19
20 declareInterface<ITracccTritonTool>(this);
21}
22
24
25 ATH_CHECK(m_TracccTritonTool.retrieve());
26
27 return StatusCode::SUCCESS;
28}
29
31 std::vector<uint8_t>& cellBytes,
32 std::vector<TracccTrackParameters>& TracccTrackParams,
33 std::vector<LocalMeasurementInfoInTracks>& TracccMeasurementInfoInTracks
34) const {
35
36 auto start_prep = std::chrono::high_resolution_clock::now();
37
38 // Forward the serialized traccc silicon_cell_collection as a single UINT8
39 AthInfer::InputDataMap inputData;
40 inputData["CELLS"] = std::make_pair(
41 std::vector<int64_t>{static_cast<int64_t>(cellBytes.size())},
42 std::move(cellBytes));
43
44 AthInfer::OutputDataMap outputData;
45 outputData["TRK_PARAMS"] = std::make_pair(
46 std::vector<int64_t>{-1, 8}, std::vector<float>{});
47 outputData["MEASUREMENTS"] = std::make_pair(
48 std::vector<int64_t>{-1, 6}, std::vector<float>{});
49 outputData["COVARIANCES"] = std::make_pair(
50 std::vector<int64_t>{-1, 25}, std::vector<float>{});
51 outputData["GEOMETRY_IDS"] = std::make_pair(
52 std::vector<int64_t>{-1, 1}, std::vector<int64_t>{});
53
54 auto end_prep = std::chrono::high_resolution_clock::now();
55 std::chrono::duration<double, std::milli> prep_time = end_prep - start_prep;
56 ATH_MSG_INFO("Traccc Triton input preparation time: " << prep_time.count() << " ms");
57
58 auto start_inference = std::chrono::high_resolution_clock::now();
59 ATH_CHECK(m_TracccTritonTool->inference(inputData, outputData));
60 auto end_inference = std::chrono::high_resolution_clock::now();
61 std::chrono::duration<double, std::milli> inference_time = end_inference - start_inference;
62 ATH_MSG_INFO("Traccc Triton inference time: " << inference_time.count() << " ms");
63
64 auto start_parse = std::chrono::high_resolution_clock::now();
65
66 // Parse outputs
67 auto& trkParamsVec = std::get<std::vector<float>>(outputData["TRK_PARAMS"].second);
68 auto& measurementsVec = std::get<std::vector<float>>(outputData["MEASUREMENTS"].second);
69 auto& covariancesVec = std::get<std::vector<float>>(outputData["COVARIANCES"].second);
70 auto& outputGeometryIds = std::get<std::vector<int64_t>>(outputData["GEOMETRY_IDS"].second);
71
72 if (trkParamsVec.empty()) {
73 ATH_MSG_DEBUG("No tracks found in the event.");
74 return StatusCode::SUCCESS;
75 }
76
77 TracccTrackParams.clear();
78 TracccMeasurementInfoInTracks.clear();
79
80 int numTrkFeatures = 8;
81 int numMeasurementFeatures = 6;
82
83 const size_t nMeasEntries = measurementsVec.size() / numMeasurementFeatures;
84 const size_t nCovEntries = covariancesVec.size() / 25;
85
87
88 size_t track = 0;
89 size_t meas_pos = 0;
90 size_t cov_pos = 0;
91 for (size_t geo_idx = 0; geo_idx < outputGeometryIds.size(); ++geo_idx)
92 {
93 // tracks are flattened by a zero
94 if (outputGeometryIds.at(geo_idx) == 0)
95 {
96 TracccMeasurementInfoInTracks.push_back(measurement);
97
98 measurement.local_x.clear();
99 measurement.local_y.clear();
100 measurement.phi.clear();
101 measurement.theta.clear();
102 measurement.qop.clear();
103 measurement.time.clear();
104 measurement.covariances.clear();
105 measurement.athena_id.clear();
106
108 params.chi2 = trkParamsVec.at(track * numTrkFeatures + 0);
109 params.ndf = trkParamsVec.at(track * numTrkFeatures + 1);
110 params.l0 = trkParamsVec.at(track * numTrkFeatures + 2);
111 params.l1 = trkParamsVec.at(track * numTrkFeatures + 3);
112 params.phi = trkParamsVec.at(track * numTrkFeatures + 4);
113 params.theta = trkParamsVec.at(track * numTrkFeatures + 5);
114 params.qop = trkParamsVec.at(track * numTrkFeatures + 6);
115 params.time = trkParamsVec.at(track * numTrkFeatures + 7);
116
117 TracccTrackParams.push_back(params);
118 track++;
119
120 continue;
121 }
122
123 if (meas_pos >= nMeasEntries || cov_pos >= nCovEntries) {
124 ATH_MSG_ERROR("Out-of-range while parsing outputs: meas_pos="
125 << meas_pos << "/" << nMeasEntries
126 << " cov_pos=" << cov_pos << "/" << nCovEntries
127 << " geo_idx=" << geo_idx);
128 break;
129 }
130
131 measurement.local_x.push_back(measurementsVec.at(meas_pos * numMeasurementFeatures + 0));
132 measurement.local_y.push_back(measurementsVec.at(meas_pos * numMeasurementFeatures + 1));
133 measurement.phi.push_back(measurementsVec.at(meas_pos * numMeasurementFeatures + 2));
134 measurement.theta.push_back(measurementsVec.at(meas_pos * numMeasurementFeatures + 3));
135 measurement.qop.push_back(measurementsVec.at(meas_pos * numMeasurementFeatures + 4));
136 measurement.time.push_back(measurementsVec.at(meas_pos * numMeasurementFeatures + 5));
137
138 for (size_t cov_idx = 0; cov_idx < 25; ++cov_idx)
139 {
140 measurement.covariances.push_back(covariancesVec.at(cov_pos * 25 + cov_idx));
141 }
142
143 measurement.athena_id.push_back(outputGeometryIds.at(geo_idx));
144
145 meas_pos++;
146 cov_pos++;
147 }
148
149 // Push the last track (no trailing separator in GEOMETRY_IDS)
150 if (!measurement.athena_id.empty())
151 {
152 TracccMeasurementInfoInTracks.push_back(measurement);
153
155 params.chi2 = trkParamsVec.at(track * numTrkFeatures + 0);
156 params.ndf = trkParamsVec.at(track * numTrkFeatures + 1);
157 params.l0 = trkParamsVec.at(track * numTrkFeatures + 2);
158 params.l1 = trkParamsVec.at(track * numTrkFeatures + 3);
159 params.phi = trkParamsVec.at(track * numTrkFeatures + 4);
160 params.theta = trkParamsVec.at(track * numTrkFeatures + 5);
161 params.qop = trkParamsVec.at(track * numTrkFeatures + 6);
162 params.time = trkParamsVec.at(track * numTrkFeatures + 7);
163 TracccTrackParams.push_back(params);
164 }
165
166 ATH_MSG_INFO("Number of tracks found: " << TracccTrackParams.size());
167
168 if (TracccTrackParams.size() != TracccMeasurementInfoInTracks.size()) {
169 ATH_MSG_WARNING("Mismatch: tracks=" << TracccTrackParams.size()
170 << " measurements=" << TracccMeasurementInfoInTracks.size());
171 }
172
173 int n_tracks_to_print = std::min(3, (int)TracccTrackParams.size());
174 for (int i = 0; i < n_tracks_to_print; ++i) {
175 ATH_MSG_DEBUG("Track " << i << " parameters: "
176 << "chi2=" << TracccTrackParams[i].chi2
177 << ", ndf=" << TracccTrackParams[i].ndf
178 << ", l0=" << TracccTrackParams[i].l0
179 << ", l1=" << TracccTrackParams[i].l1
180 << ", phi=" << TracccTrackParams[i].phi
181 << ", theta=" << TracccTrackParams[i].theta
182 << ", qop=" << TracccTrackParams[i].qop
183 << ", time=" << TracccTrackParams[i].time);
184
185 if (i < (int)TracccMeasurementInfoInTracks.size()) {
186 const auto& measurements = TracccMeasurementInfoInTracks[i];
187 for (size_t j = 0; j < measurements.athena_id.size(); ++j) {
188 ATH_MSG_DEBUG(" Measurement " << j << ": "
189 << "local_x=" << measurements.local_x[j]
190 << ", local_y=" << measurements.local_y[j]
191 << ", phi=" << measurements.phi[j]
192 << ", theta=" << measurements.theta[j]
193 << ", qop=" << measurements.qop[j]
194 << ", time=" << measurements.time[j]
195 << ", athena_id=" << measurements.athena_id[j]);
196
197 // Print full covariance matrix (assumes 5x5 = 25 elements per measurement)
198 const size_t cov_per_meas = 25;
199 size_t cov_start = j * cov_per_meas;
200 if (measurements.covariances.size() >= cov_start + cov_per_meas) {
201 std::ostringstream oss;
202 oss << " Covariance matrix:";
203 for (size_t r = 0; r < 5; ++r) {
204 oss << "\n [ ";
205 for (size_t c = 0; c < 5; ++c) {
206 oss << measurements.covariances[cov_start + r * 5 + c] << " ";
207 }
208 oss << "]";
209 }
210 ATH_MSG_DEBUG(oss.str());
211 } else {
212 ATH_MSG_WARNING(" Covariance data unavailable for measurement " << j);
213 }
214 }
215 }
216 }
217
218 auto end_parse = std::chrono::high_resolution_clock::now();
219 std::chrono::duration<double, std::milli> parse_time = end_parse - start_parse;
220 ATH_MSG_INFO("Traccc Triton output parsing time: " << parse_time.count() << " ms");
221
222 return StatusCode::SUCCESS;
223}
Scalar phi() const
phi method
Scalar theta() const
theta method
#define ATH_CHECK
Evaluate an expression and check for errors.
#define ATH_MSG_DEBUG(x,...)
#define ATH_MSG_ERROR(x,...)
#define ATH_MSG_WARNING(x,...)
#define ATH_MSG_INFO(x,...)
virtual StatusCode initialize() override
virtual StatusCode getTracks(std::vector< uint8_t > &cellBytes, std::vector< TracccTrackParameters > &TracccTrackParameters, std::vector< LocalMeasurementInfoInTracks > &TracccMeasurementsInfoInTracks) const override
Get track candidates from serialized traccc cells.
ToolHandle< AthInfer::IAthInferenceTool > m_TracccTritonTool
TracccTritonTool(const std::string &type, const std::string &name, const IInterface *parent)
double chi2(TH1 *h0, TH1 *h1)
int r
Definition globals.cxx:22
std::map< std::string, InferenceData > OutputDataMap
std::map< std::string, InferenceData > InputDataMap
std::vector< int64_t > athena_id
std::vector< float > covariances