ATLAS Offline Software
Loading...
Searching...
No Matches
TracccTritonTool.cxx
Go to the documentation of this file.
1/*
2 Copyright (C) 2002-2025 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<TracccCell>& cells,
32 std::vector<TracccTrackParameters>& TracccTrackParams,
33 std::vector<LocalMeasurementInfoInTracks>& TracccMeasurementInfoInTracks
34) const {
35
36 auto start_prep = std::chrono::high_resolution_clock::now();
37 int numCells = cells.size();
38
39 std::vector<int64_t> CellPositions;
40 std::vector<float> CellProperties;
41
42 for (const auto& cell : cells)
43 {
44 CellPositions.push_back(cell.geometry_id);
45 CellPositions.push_back(cell.measurement_id);
46 CellPositions.push_back(cell.channel0);
47 CellPositions.push_back(cell.channel1);
48
49 CellProperties.push_back(cell.timestamp);
50 CellProperties.push_back(cell.value);
51 }
52
54 // header: geometry_id, measurement_id, channel0, channel1, timestamp, value
55 std::string csv_filename = "events/event" +
56 std::to_string(m_eventCounter.fetch_add(1)) + "-cells.csv";
57 std::ofstream csv_file(csv_filename);
58 if (!csv_file.is_open()) {
59 ATH_MSG_ERROR("Failed to open CSV file: " << csv_filename);
60 return StatusCode::FAILURE;
61 }
62 csv_file << "geometry_id,measurement_id,channel0,channel1,timestamp,value\n";
63 for (const auto& cell : cells) {
64 csv_file << cell.geometry_id << ","
65 << cell.measurement_id << ","
66 << cell.channel0 << ","
67 << cell.channel1 << ","
68 << cell.timestamp << ","
69 << cell.value << "\n";
70 }
71 csv_file.close();
72 }
73
74 AthInfer::InputDataMap inputData;
75 inputData["CELL_POSITIONS"] = std::make_pair(
76 std::vector<int64_t>{numCells, 4}, std::move(CellPositions));
77 inputData["CELL_PROPERTIES"] = std::make_pair(
78 std::vector<int64_t>{numCells, 2}, std::move(CellProperties));
79
80 AthInfer::OutputDataMap outputData;
81 outputData["TRK_PARAMS"] = std::make_pair(
82 std::vector<int64_t>{-1, 8}, std::vector<float>{});
83 outputData["MEASUREMENTS"] = std::make_pair(
84 std::vector<int64_t>{-1, 6}, std::vector<float>{});
85 outputData["COVARIANCES"] = std::make_pair(
86 std::vector<int64_t>{-1, 25}, std::vector<float>{});
87 outputData["GEOMETRY_IDS"] = std::make_pair(
88 std::vector<int64_t>{-1, 1}, std::vector<int64_t>{});
89
90 auto end_prep = std::chrono::high_resolution_clock::now();
91 std::chrono::duration<double, std::milli> prep_time = end_prep - start_prep;
92 ATH_MSG_INFO("Traccc Triton input preparation time: " << prep_time.count() << " ms");
93
94 auto start_inference = std::chrono::high_resolution_clock::now();
95 ATH_CHECK(m_TracccTritonTool->inference(inputData, outputData));
96 auto end_inference = std::chrono::high_resolution_clock::now();
97 std::chrono::duration<double, std::milli> inference_time = end_inference - start_inference;
98 ATH_MSG_INFO("Traccc Triton inference time: " << inference_time.count() << " ms");
99
100 auto start_parse = std::chrono::high_resolution_clock::now();
101
102 // Parse outputs
103 auto& trkParamsVec = std::get<std::vector<float>>(outputData["TRK_PARAMS"].second);
104 auto& measurementsVec = std::get<std::vector<float>>(outputData["MEASUREMENTS"].second);
105 auto& covariancesVec = std::get<std::vector<float>>(outputData["COVARIANCES"].second);
106 auto& outputGeometryIds = std::get<std::vector<int64_t>>(outputData["GEOMETRY_IDS"].second);
107
108 if (trkParamsVec.empty()) {
109 ATH_MSG_DEBUG("No tracks found in the event.");
110 return StatusCode::SUCCESS;
111 }
112
113 TracccTrackParams.clear();
114 TracccMeasurementInfoInTracks.clear();
115
116 int numTrkFeatures = 8;
117 int numMeasurementFeatures = 6;
118
119 const size_t nMeasEntries = measurementsVec.size() / numMeasurementFeatures;
120 const size_t nCovEntries = covariancesVec.size() / 25;
121
123
124 size_t track = 0;
125 size_t meas_pos = 0;
126 size_t cov_pos = 0;
127 for (size_t geo_idx = 0; geo_idx < outputGeometryIds.size(); ++geo_idx)
128 {
129 // tracks are flattened by a zero
130 if (outputGeometryIds.at(geo_idx) == 0)
131 {
132 TracccMeasurementInfoInTracks.push_back(measurement);
133
134 measurement.local_x.clear();
135 measurement.local_y.clear();
136 measurement.phi.clear();
137 measurement.theta.clear();
138 measurement.qop.clear();
139 measurement.time.clear();
140 measurement.covariances.clear();
141 measurement.athena_id.clear();
142
144 params.chi2 = trkParamsVec.at(track * numTrkFeatures + 0);
145 params.ndf = trkParamsVec.at(track * numTrkFeatures + 1);
146 params.l0 = trkParamsVec.at(track * numTrkFeatures + 2);
147 params.l1 = trkParamsVec.at(track * numTrkFeatures + 3);
148 params.phi = trkParamsVec.at(track * numTrkFeatures + 4);
149 params.theta = trkParamsVec.at(track * numTrkFeatures + 5);
150 params.qop = trkParamsVec.at(track * numTrkFeatures + 6);
151 params.time = trkParamsVec.at(track * numTrkFeatures + 7);
152
153 TracccTrackParams.push_back(params);
154 track++;
155
156 continue;
157 }
158
159 if (meas_pos >= nMeasEntries || cov_pos >= nCovEntries) {
160 ATH_MSG_ERROR("Out-of-range while parsing outputs: meas_pos="
161 << meas_pos << "/" << nMeasEntries
162 << " cov_pos=" << cov_pos << "/" << nCovEntries
163 << " geo_idx=" << geo_idx);
164 break;
165 }
166
167 measurement.local_x.push_back(measurementsVec.at(meas_pos * numMeasurementFeatures + 0));
168 measurement.local_y.push_back(measurementsVec.at(meas_pos * numMeasurementFeatures + 1));
169 measurement.phi.push_back(measurementsVec.at(meas_pos * numMeasurementFeatures + 2));
170 measurement.theta.push_back(measurementsVec.at(meas_pos * numMeasurementFeatures + 3));
171 measurement.qop.push_back(measurementsVec.at(meas_pos * numMeasurementFeatures + 4));
172 measurement.time.push_back(measurementsVec.at(meas_pos * numMeasurementFeatures + 5));
173
174 for (size_t cov_idx = 0; cov_idx < 25; ++cov_idx)
175 {
176 measurement.covariances.push_back(covariancesVec.at(cov_pos * 25 + cov_idx));
177 }
178
179 measurement.athena_id.push_back(outputGeometryIds.at(geo_idx));
180
181 meas_pos++;
182 cov_pos++;
183 }
184
185 // Push the last track (no trailing separator in GEOMETRY_IDS)
186 if (!measurement.athena_id.empty())
187 {
188 TracccMeasurementInfoInTracks.push_back(measurement);
189
191 params.chi2 = trkParamsVec.at(track * numTrkFeatures + 0);
192 params.ndf = trkParamsVec.at(track * numTrkFeatures + 1);
193 params.l0 = trkParamsVec.at(track * numTrkFeatures + 2);
194 params.l1 = trkParamsVec.at(track * numTrkFeatures + 3);
195 params.phi = trkParamsVec.at(track * numTrkFeatures + 4);
196 params.theta = trkParamsVec.at(track * numTrkFeatures + 5);
197 params.qop = trkParamsVec.at(track * numTrkFeatures + 6);
198 params.time = trkParamsVec.at(track * numTrkFeatures + 7);
199 TracccTrackParams.push_back(params);
200 }
201
202 ATH_MSG_INFO("Number of tracks found: " << TracccTrackParams.size());
203
204 if (TracccTrackParams.size() != TracccMeasurementInfoInTracks.size()) {
205 ATH_MSG_WARNING("Mismatch: tracks=" << TracccTrackParams.size()
206 << " measurements=" << TracccMeasurementInfoInTracks.size());
207 }
208
209 int n_tracks_to_print = std::min(3, (int)TracccTrackParams.size());
210 for (int i = 0; i < n_tracks_to_print; ++i) {
211 ATH_MSG_DEBUG("Track " << i << " parameters: "
212 << "chi2=" << TracccTrackParams[i].chi2
213 << ", ndf=" << TracccTrackParams[i].ndf
214 << ", l0=" << TracccTrackParams[i].l0
215 << ", l1=" << TracccTrackParams[i].l1
216 << ", phi=" << TracccTrackParams[i].phi
217 << ", theta=" << TracccTrackParams[i].theta
218 << ", qop=" << TracccTrackParams[i].qop
219 << ", time=" << TracccTrackParams[i].time);
220
221 if (i < (int)TracccMeasurementInfoInTracks.size()) {
222 const auto& measurements = TracccMeasurementInfoInTracks[i];
223 for (size_t j = 0; j < measurements.athena_id.size(); ++j) {
224 ATH_MSG_DEBUG(" Measurement " << j << ": "
225 << "local_x=" << measurements.local_x[j]
226 << ", local_y=" << measurements.local_y[j]
227 << ", phi=" << measurements.phi[j]
228 << ", theta=" << measurements.theta[j]
229 << ", qop=" << measurements.qop[j]
230 << ", time=" << measurements.time[j]
231 << ", athena_id=" << measurements.athena_id[j]);
232
233 // Print full covariance matrix (assumes 5x5 = 25 elements per measurement)
234 const size_t cov_per_meas = 25;
235 size_t cov_start = j * cov_per_meas;
236 if (measurements.covariances.size() >= cov_start + cov_per_meas) {
237 std::ostringstream oss;
238 oss << " Covariance matrix:";
239 for (size_t r = 0; r < 5; ++r) {
240 oss << "\n [ ";
241 for (size_t c = 0; c < 5; ++c) {
242 oss << measurements.covariances[cov_start + r * 5 + c] << " ";
243 }
244 oss << "]";
245 }
246 ATH_MSG_DEBUG(oss.str());
247 } else {
248 ATH_MSG_WARNING(" Covariance data unavailable for measurement " << j);
249 }
250 }
251 }
252 }
253
254 auto end_parse = std::chrono::high_resolution_clock::now();
255 std::chrono::duration<double, std::milli> parse_time = end_parse - start_parse;
256 ATH_MSG_INFO("Traccc Triton output parsing time: " << parse_time.count() << " ms");
257
258 return StatusCode::SUCCESS;
259}
Scalar phi() const
phi method
Scalar theta() const
theta method
#define ATH_CHECK
Evaluate an expression and check for errors.
#define ATH_MSG_ERROR(x)
#define ATH_MSG_INFO(x)
#define ATH_MSG_WARNING(x)
#define ATH_MSG_DEBUG(x)
Gaudi::Property< int > m_maxEventsToSave
virtual StatusCode initialize() override
virtual StatusCode getTracks(std::vector< TracccCell > &cells, std::vector< TracccTrackParameters > &TracccTrackParameters, std::vector< LocalMeasurementInfoInTracks > &TracccMeasurementsInfoInTracks) const override
Get track candidates from a list of space points.
ToolHandle< AthInfer::IAthInferenceTool > m_TracccTritonTool
Gaudi::Property< bool > m_saveEventsToCSV
TracccTritonTool(const std::string &type, const std::string &name, const IInterface *parent)
std::atomic< int > m_eventCounter
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