ATLAS Offline Software
Loading...
Searching...
No Matches
CaloClusterMLCalibToolLite Class Reference

#include <CaloClusterMLCalibToolLite.h>

Inheritance diagram for CaloClusterMLCalibToolLite:
Collaboration diagram for CaloClusterMLCalibToolLite:

Public Member Functions

 CaloClusterMLCalibToolLite (const std::string &type, const std::string &name, const IInterface *parent)
 ~CaloClusterMLCalibToolLite ()
virtual StatusCode initialize () override
virtual StatusCode finalize () override
virtual StatusCode inference (const xAOD::CaloClusterContainer &clusters, int nPrimVtx, float avgMu, std::vector< double > &clusterE_ML_vec, std::vector< double > &clusterE_ML_Unc_vec) const override

Private Attributes

Gaudi::Property< std::vector< std::string > > m_preprocessingTransformNames {this, "PreprocessingTransformNames", {}, "Names of preprocessing transforms"}
Gaudi::Property< std::vector< std::vector< double > > > m_preprocessingTransformParams {this, "PreprocessingTransformParams", {}, "Parameters for preprocessing transforms"}
int m_numFeatures = 0
std::vector< PreprocessTransform > m_featurePreprocessingTransforms
ToolHandle< AthInfer::IAthInferenceTool > m_onnxTool

Detailed Description

Definition at line 24 of file CaloClusterMLCalibToolLite.h.

Constructor & Destructor Documentation

◆ CaloClusterMLCalibToolLite()

CaloClusterMLCalibToolLite::CaloClusterMLCalibToolLite ( const std::string & type,
const std::string & name,
const IInterface * parent )

Definition at line 12 of file CaloClusterMLCalibToolLite.cxx.

12: base_class(type, name, parent) {}

◆ ~CaloClusterMLCalibToolLite()

CaloClusterMLCalibToolLite::~CaloClusterMLCalibToolLite ( )

Definition at line 14 of file CaloClusterMLCalibToolLite.cxx.

14{}

Member Function Documentation

◆ finalize()

StatusCode CaloClusterMLCalibToolLite::finalize ( )
overridevirtual

Definition at line 232 of file CaloClusterMLCalibToolLite.cxx.

233{
234 ATH_MSG_DEBUG("Finalizing " << name() << "...");
235 return StatusCode::SUCCESS;
236}
#define ATH_MSG_DEBUG(x,...)

◆ inference()

StatusCode CaloClusterMLCalibToolLite::inference ( const xAOD::CaloClusterContainer & clusters,
int nPrimVtx,
float avgMu,
std::vector< double > & clusterE_ML_vec,
std::vector< double > & clusterE_ML_Unc_vec ) const
overridevirtual

Definition at line 46 of file CaloClusterMLCalibToolLite.cxx.

51{
52 ATH_MSG_DEBUG("Executing " << name() << "...");
53
54 double clusterE = 0;
55 double clusterEta = 0;
56 double cluster_SIGNIFICANCE = 0;
57 double cluster_time = 0;
58 double cluster_SECOND_TIME = 0;
59 double cluster_CENTER_LAMBDA = 0;
60 double cluster_CENTER_MAG = 0;
61 double cluster_ENG_FRAC_EM_INCL = 0;
62 double cluster_FIRST_ENG_DENS = 0;
63 double cluster_LONGITUDINAL = 0;
64 double cluster_LATERAL = 0;
65 double cluster_PTD = 0;
66 double cluster_ISOLATION = 0;
67
68 std::vector<float> transformedFeatures;
69 std::vector<bool> clusterInputsValid;
70 clusterInputsValid.reserve(clusters.size());
71 bool ok{}; // for checking return value of cluster->retrieveMoment
72
73 for (const xAOD::CaloCluster *cluster : clusters)
74 {
75 clusterE = cluster->e(xAOD::CaloCluster::UNCALIBRATED) / Gaudi::Units::GeV;
76 clusterEta = cluster->eta(xAOD::CaloCluster::UNCALIBRATED);
77 ok = cluster->retrieveMoment(xAOD::CaloCluster::MomentType::SIGNIFICANCE, cluster_SIGNIFICANCE);
78 ok &= cluster->retrieveMoment(xAOD::CaloCluster::MomentType::SECOND_TIME, cluster_SECOND_TIME);
79 cluster_SECOND_TIME /= (Gaudi::Units::nanosecond * Gaudi::Units::nanosecond);
80 ok &= cluster->retrieveMoment(xAOD::CaloCluster::MomentType::CENTER_LAMBDA, cluster_CENTER_LAMBDA);
81 cluster_CENTER_LAMBDA /= Gaudi::Units::millimeter;
82 ok &= cluster->retrieveMoment(xAOD::CaloCluster::MomentType::CENTER_MAG, cluster_CENTER_MAG);
83 ok &= cluster->retrieveMoment(xAOD::CaloCluster::MomentType::FIRST_ENG_DENS, cluster_FIRST_ENG_DENS);
84 cluster_FIRST_ENG_DENS /= (Gaudi::Units::GeV / Gaudi::Units::millimeter3);
85 ok &= cluster->retrieveMoment(xAOD::CaloCluster::MomentType::LONGITUDINAL, cluster_LONGITUDINAL);
86 ok &= cluster->retrieveMoment(xAOD::CaloCluster::MomentType::LATERAL, cluster_LATERAL);
87 ok &= cluster->retrieveMoment(xAOD::CaloCluster::MomentType::PTD, cluster_PTD);
88 ok &= cluster->retrieveMoment(xAOD::CaloCluster::MomentType::ISOLATION, cluster_ISOLATION);
89 cluster_time = cluster->time() / Gaudi::Units::nanosecond;
90
91 if (!ok) {
92 ATH_MSG_ERROR("retrieveMoment() failed for " << cluster);
93 return StatusCode::FAILURE;
94 }
95
96 float e_EM = 0.0;
97 for (size_t s = CaloSampling::PreSamplerB; s < CaloSampling::Unknown; s++)
98 {
99 if (s == CaloSampling::EMB1 || s == CaloSampling::EMB2 || s == CaloSampling::EMB3 || s == CaloSampling::EME1 || s == CaloSampling::EME2 || s == CaloSampling::EME3 || s == CaloSampling::FCAL0)
100 {
101 e_EM += cluster->eSample(static_cast<xAOD::CaloCluster::CaloSample>(s));
102 }
103 }
104 cluster_ENG_FRAC_EM_INCL = e_EM / cluster->rawE();
105
106 std::vector<float> rawValues = {
107 static_cast<float>(clusterE),
108 static_cast<float>(clusterEta),
109 static_cast<float>(cluster_SIGNIFICANCE),
110 static_cast<float>(cluster_time),
111 static_cast<float>(cluster_SECOND_TIME),
112 static_cast<float>(cluster_CENTER_LAMBDA),
113 static_cast<float>(cluster_CENTER_MAG),
114 static_cast<float>(cluster_ENG_FRAC_EM_INCL),
115 static_cast<float>(cluster_FIRST_ENG_DENS),
116 static_cast<float>(cluster_LONGITUDINAL),
117 static_cast<float>(cluster_LATERAL),
118 static_cast<float>(cluster_PTD),
119 static_cast<float>(cluster_ISOLATION),
120 static_cast<float>(nPrimVtx),
121 avgMu
122 };
123
124 bool inputsValid = true;
125 for (int i = 0; i < m_numFeatures; i++)
126 {
127 const PreprocessTransform &transform = m_featurePreprocessingTransforms.at(i);
128 const float raw = rawValues.at(i);
129 float transformed = transform.processor(raw, transform.parameters);
130 if (!std::isfinite(transformed)) {
131 inputsValid = false;
132 // Keep the batched ONNX input finite. Its output is ignored for
133 // this cluster below.
134 transformed = 0.0F;
135 }
136 transformedFeatures.push_back(transformed);
137 }
138 clusterInputsValid.push_back(inputsValid);
139 }
140
141 int numClusters = clusters.size();
142 std::vector<int64_t> inputShape = {numClusters, 15};
143
144 AthInfer::InputDataMap inputData;
145 inputData["features"] = std::make_pair(
146 inputShape, std::move(transformedFeatures));
147
148 AthInfer::OutputDataMap outputData;
149
150 outputData["mus"] = std::make_pair(
151 std::vector<int64_t>{numClusters, 3}, std::vector<float>{});
152 outputData["sigmas"] = std::make_pair(
153 std::vector<int64_t>{numClusters, 3}, std::vector<float>{});
154 outputData["alphas"] = std::make_pair(
155 std::vector<int64_t>{numClusters, 3}, std::vector<float>{});
156
157 ATH_CHECK(m_onnxTool->inference(inputData, outputData));
158
159 std::vector<float> &onnx_mus = std::get<std::vector<float>>(outputData["mus"].second);
160 std::vector<float> &onnx_sigma2s = std::get<std::vector<float>>(outputData["sigmas"].second);
161 std::vector<float> &onnx_alphas = std::get<std::vector<float>>(outputData["alphas"].second);
162
163 if (msgLvl(MSG::DEBUG)) {
164 int nan_in_mus = 0;
165 int nan_in_sigma2s = 0;
166 int nan_in_alphas = 0;
167
168 for (float val : onnx_mus)
169 if (std::isnan(val))
170 nan_in_mus++;
171
172 for (float val : onnx_sigma2s)
173 if (std::isnan(val))
174 nan_in_sigma2s++;
175
176 for (float val : onnx_alphas)
177 if (std::isnan(val))
178 nan_in_alphas++;
179
180 if (nan_in_mus > 0)
181 ATH_MSG_DEBUG(nan_in_mus << " NaN value found in `mus` output layer during ONNX inference");
182
183 if (nan_in_sigma2s)
184 ATH_MSG_DEBUG(nan_in_sigma2s << " NaN value found in `sigmas` output layer during ONNX inference");
185
186 if (nan_in_alphas)
187 ATH_MSG_DEBUG(nan_in_alphas << " NaN value found in `alphas` output layer during ONNX inference");
188 }
189
190 clusterE_ML_vec.clear();
191 clusterE_ML_Unc_vec.clear();
192 clusterE_ML_vec.reserve(numClusters);
193 clusterE_ML_Unc_vec.reserve(numClusters);
194
195 for (int i = 0; i < numClusters; ++i)
196 {
197 bool calibrateCluster = clusterInputsValid.at(i);
198 for (size_t j=0; j<3; ++j) {
199 if (std::isnan(onnx_mus[i*3+j]) || std::isnan(onnx_sigma2s[i*3+j]) || std::isnan(onnx_alphas[i*3+j])) {
200 calibrateCluster = false;
201 break;
202 }
203 }
204 float r = 1.;
205 float s = 0.;
206 if (calibrateCluster) {
207 std::vector<float> current_mus = {onnx_mus[i * 3], onnx_mus[i * 3 + 1], onnx_mus[i * 3 + 2]};
208 std::vector<float> current_sigma2s = {onnx_sigma2s[i * 3], onnx_sigma2s[i * 3 + 1], onnx_sigma2s[i * 3 + 2]};
209 std::vector<float> current_alphas = {onnx_alphas[i * 3], onnx_alphas[i * 3 + 1], onnx_alphas[i * 3 + 2]};
210
211 float mode = CaloClusterMLCalib::modes(current_mus, current_sigma2s, current_alphas);
212 r = std::pow(10, mode);
213 float onnx_s = CaloClusterMLCalib::sigma_stoch(current_mus, current_sigma2s, current_alphas);
214 s = std::abs(std::log(10) * r) * onnx_s;
215
216 if (!std::isfinite(r) || std::abs(r) < 1e-6) {
217 ATH_MSG_WARNING("ML-correction factor to cluster energy (used as denominator) is " << r << "; The ML-correction factor is reset to 1. Uncertainty is set to 0.");
218 r = 1.0;
219 s = 0.0;
220 }
221 }
222
223 const double cluster_energy = clusters[i]->e(xAOD::CaloCluster::UNCALIBRATED) / static_cast<double>(r);
224
225 clusterE_ML_vec.push_back(cluster_energy);
226 clusterE_ML_Unc_vec.push_back(static_cast<double>(s));
227 }
228
229 return StatusCode::SUCCESS;
230}
#define ATH_CHECK
Evaluate an expression and check for errors.
#define ATH_MSG_ERROR(x,...)
#define ATH_MSG_WARNING(x,...)
ToolHandle< AthInfer::IAthInferenceTool > m_onnxTool
std::vector< PreprocessTransform > m_featurePreprocessingTransforms
@ PTD
relative spread of pT of constiuent cells = sqrt(n)*RMS/Mean
@ SECOND_TIME
Second moment of cell time distribution in cluster.
@ LATERAL
Normalized lateral moment.
@ LONGITUDINAL
Normalized longitudinal moment.
@ FIRST_ENG_DENS
First Moment in E/V.
@ CENTER_LAMBDA
Shower depth at Cluster Centroid.
@ SIGNIFICANCE
Cluster significance.
@ CENTER_MAG
Cluster Centroid ( ).
@ ISOLATION
Energy weighted fraction of non-clustered perimeter cells.
CaloSampling::CaloSample CaloSample
int r
Definition globals.cxx:22
Amg::Vector3D transform(Amg::Vector3D &v, Amg::Transform3D &tr)
Transform a point from a Trasformation3D.
std::map< std::string, InferenceData > OutputDataMap
std::map< std::string, InferenceData > InputDataMap
float modes(const std::vector< float > &mus, const std::vector< float > &log_sigma2s, const std::vector< float > &alphas)
float sigma_stoch(const std::vector< float > &mus, const std::vector< float > &log_sigma2s, const std::vector< float > &alphas)
float j(const xAOD::IParticle &, const xAOD::TrackMeasurementValidation &hit, const Eigen::Matrix3d &jab_inv)
CaloCluster_v1 CaloCluster
Define the latest version of the calorimeter cluster class.

◆ initialize()

StatusCode CaloClusterMLCalibToolLite::initialize ( )
overridevirtual

Definition at line 16 of file CaloClusterMLCalibToolLite.cxx.

17{
18 ATH_MSG_DEBUG("Initializing " << name() << "...");
20 for (int i = 0; i < m_numFeatures; i++)
21 {
22 PreprocessTransform transform;
25 {
26 transform.processor = funcIt->second;
27 std::vector<float> floatParams;
28 for (double param : m_preprocessingTransformParams[i])
29 {
30 floatParams.push_back(static_cast<float>(param));
31 }
32 transform.parameters = std::move(floatParams);
33 m_featurePreprocessingTransforms.push_back(std::move(transform));
34 }
35 else
36 {
37 ATH_MSG_WARNING("Undefined transformation " << m_preprocessingTransformNames[i]);
38 return StatusCode::FAILURE;
39 }
40 }
41
42 ATH_CHECK(m_onnxTool.retrieve());
43 return StatusCode::SUCCESS;
44}
Gaudi::Property< std::vector< std::vector< double > > > m_preprocessingTransformParams
Gaudi::Property< std::vector< std::string > > m_preprocessingTransformNames
const std::map< std::string, TransformFunc > TRANSFORMATIONS

Member Data Documentation

◆ m_featurePreprocessingTransforms

std::vector<PreprocessTransform> CaloClusterMLCalibToolLite::m_featurePreprocessingTransforms
private

Definition at line 46 of file CaloClusterMLCalibToolLite.h.

◆ m_numFeatures

int CaloClusterMLCalibToolLite::m_numFeatures = 0
private

Definition at line 45 of file CaloClusterMLCalibToolLite.h.

◆ m_onnxTool

ToolHandle<AthInfer::IAthInferenceTool> CaloClusterMLCalibToolLite::m_onnxTool
private
Initial value:
{
this, "ORTInferenceTool", "AthOnnx::OnnxRuntimeInferenceTool"}

Definition at line 48 of file CaloClusterMLCalibToolLite.h.

48 {
49 this, "ORTInferenceTool", "AthOnnx::OnnxRuntimeInferenceTool"};

◆ m_preprocessingTransformNames

Gaudi::Property<std::vector<std::string> > CaloClusterMLCalibToolLite::m_preprocessingTransformNames {this, "PreprocessingTransformNames", {}, "Names of preprocessing transforms"}
private

Definition at line 42 of file CaloClusterMLCalibToolLite.h.

42{this, "PreprocessingTransformNames", {}, "Names of preprocessing transforms"};

◆ m_preprocessingTransformParams

Gaudi::Property<std::vector<std::vector<double> > > CaloClusterMLCalibToolLite::m_preprocessingTransformParams {this, "PreprocessingTransformParams", {}, "Parameters for preprocessing transforms"}
private

Definition at line 43 of file CaloClusterMLCalibToolLite.h.

43{this, "PreprocessingTransformParams", {}, "Parameters for preprocessing transforms"};

The documentation for this class was generated from the following files: