ATLAS Offline Software
Loading...
Searching...
No Matches
GlobalLargeRDNNCalibration.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// System includes
7
8
9#ifdef XAOD_STANDALONE
11#endif
12
16
17#include <TEnv.h>
18#include <tuple>
19#include <cmath>
20#include <map>
21#include <algorithm> //count_if
22
23
24namespace{
25 // Redefine some functions from the package OnnxRuntimeUtils which is not (yet) available in AnalysisBase
26 // Set up the ONNX Runtime session
27 std::unique_ptr< Ort::Session > CreateORTSession(const std::string& modelFile){
28 Ort::SessionOptions sessionOptions;
29 sessionOptions.SetIntraOpNumThreads( 1 );
30 sessionOptions.SetGraphOptimizationLevel( ORT_ENABLE_BASIC );
31
32 // Set the ONNX service name depending on the actual analysis release
33 std::string serviceName;
34#ifdef XAOD_STANDALONE
35 using namespace asg::msgUserCode;
36 ANA_MSG_WARNING("If running DNN calibration in AnalysisBase: necessary to instantiate the ONNX service AthOnnx::OnnxRuntimeSvc with name OnnxRuntimeSvc");
37 ANA_MSG_WARNING("Either in C++ config (see exemple in JetCalibTools_Example.cxx)");
38 ANA_MSG_WARNING("Or in python config with");
39 ANA_MSG_WARNING(" from AnaAlgorithm.DualUseConfig import createService");
40 ANA_MSG_WARNING(" onnxSvc = createService('AthOnnx::OnnxRuntimeSvc', 'OnnxRuntimeSvc', myAlgSequence)");
41 serviceName = "OnnxRuntimeSvc";
42#else
43 serviceName = "AthOnnx::OnnxRuntimeSvc";
44#endif
45
46 ServiceHandle< AthOnnx::IOnnxRuntimeSvc > svc(serviceName, "AthOnnx::OnnxRuntimeSvc");
47
48 return std::make_unique<Ort::Session>( svc->env(),
49 modelFile.c_str(),
50 sessionOptions );
51 }
52
53 // Get dimensions and names of the input nodes
54 std::tuple<std::vector<int64_t>, std::vector<const char*> > GetInputNodeInfo(const std::unique_ptr< Ort::Session >& session){
55 std::vector<int64_t> input_node_dims;
56 size_t num_input_nodes = session->GetInputCount();
57 std::vector<const char*> input_node_names(num_input_nodes);
58 Ort::AllocatorWithDefaultOptions allocator;
59 for( std::size_t i = 0; i < num_input_nodes; i++ ) {
60
61 char* input_name = session->GetInputNameAllocated(i, allocator).release();
62 input_node_names[i] = input_name;
63 Ort::TypeInfo type_info = session->GetInputTypeInfo(i);
64 auto tensor_info = type_info.GetTensorTypeAndShapeInfo();
65
66 input_node_dims = tensor_info.GetShape();
67 }
68 return std::make_tuple(input_node_dims, input_node_names);
69 }
70
71 // Get dimensions and names of the output nodes
72 std::tuple<std::vector<int64_t>, std::vector<const char*> > GetOutputNodeInfo(const std::unique_ptr< Ort::Session >& session){
73 std::vector<int64_t> output_node_dims;
74 size_t num_output_nodes = session->GetOutputCount();
75 std::vector<const char*> output_node_names(num_output_nodes);
76 Ort::AllocatorWithDefaultOptions allocator;
77
78 for( std::size_t i = 0; i < num_output_nodes; i++ ) {
79 char* output_name = session->GetOutputNameAllocated(i, allocator).release();
80 output_node_names[i] = output_name;
81
82 Ort::TypeInfo type_info = session->GetOutputTypeInfo(i);
83 auto tensor_info = type_info.GetTensorTypeAndShapeInfo();
84
85 output_node_dims = tensor_info.GetShape();
86 }
87 return std::make_tuple(output_node_dims, output_node_names);
88 }
89}
90
91
95
97 virtual float value(const xAOD::Jet& jet, JetEventInfo& jetInfo, double eScale) = 0;
98 virtual ~VarRetriever()= default;
99};
100
101namespace {
102
104 struct VarAccessorRetriever : public GlobalLargeRDNNCalibration::VarRetriever {
105 VarAccessorRetriever(const std::string &n): m_acc(n) {}
106
107 virtual float value(const xAOD::Jet& jet, JetEventInfo&, double eScale) {
108 return m_acc(jet) * eScale;
109 }
110
111 SG::AuxElement::ConstAccessor<float> m_acc;
112 };
113
116 struct RatioAccessorRetriever : public GlobalLargeRDNNCalibration::VarRetriever {
117 RatioAccessorRetriever(): m_accTau1("Tau1_wta"),
118 m_accTau2("Tau2_wta"),
119 m_accTau3("Tau3_wta"),
120 m_accECF1("ECF1"),
121 m_accECF2("ECF2"),
122 m_accECF3("ECF3") {}
123
124 virtual float value(const xAOD::Jet& jet, JetEventInfo&, double eScale) = 0;
125
126 SG::AuxElement::ConstAccessor<float> m_accTau1;
127 SG::AuxElement::ConstAccessor<float> m_accTau2;
128 SG::AuxElement::ConstAccessor<float> m_accTau3;
129 SG::AuxElement::ConstAccessor<float> m_accECF1;
130 SG::AuxElement::ConstAccessor<float> m_accECF2;
131 SG::AuxElement::ConstAccessor<float> m_accECF3;
132 };
133
135 #define DEF_RETRIEVER0(cname, expr ) struct Var_##cname : public GlobalLargeRDNNCalibration::VarRetriever { float value(const xAOD::Jet& jet, JetEventInfo& , double eScale ) { return expr ; } }
136 #define DEF_RETRIEVER1(cname, expr ) struct Var_##cname : public GlobalLargeRDNNCalibration::VarRetriever { float value(const xAOD::Jet& , JetEventInfo& jetInfo, double eScale ) { return expr ; } }
137 #define DEF_RATIO_RETRIEVER(cname, expr ) struct Ratio_##cname : public RatioAccessorRetriever { float value(const xAOD::Jet& jet, JetEventInfo& , double eScale ) { return expr ; } }
138
139 // Std jet variables
140 DEF_RETRIEVER0( eta, jet.eta()*eScale ) ;
141 DEF_RETRIEVER0( rapidity, jet.rapidity()*eScale ) ;
142 DEF_RETRIEVER0( log_e, log(jet.e()*eScale) ) ;
143 DEF_RETRIEVER0( log_m, log(jet.m()*eScale) ) ;
144 DEF_RETRIEVER0( m, jet.m()*eScale ) ;
145 DEF_RETRIEVER0( log_m_40, jet.m()<40000 ? log(40000*eScale) : log(jet.m()*eScale) ) ;
146
147 // Ratio variables -- default values consistent with DNN training
148 DEF_RATIO_RETRIEVER( Tau21_wta, m_accTau1(jet) > 1e-8 ? eScale * m_accTau2(jet) / m_accTau1(jet) : -0.1);
149 DEF_RATIO_RETRIEVER( Tau32_wta, m_accTau2(jet) > 1e-8 ? eScale * m_accTau3(jet) / m_accTau2(jet) : -0.1);
150 DEF_RATIO_RETRIEVER( C2, m_accECF2(jet) > 1e-8 ? eScale * m_accECF3(jet) * m_accECF1(jet) / pow(m_accECF2(jet), 2.0) : -0.1);
151 DEF_RATIO_RETRIEVER( D2, m_accECF2(jet) > 1e-8 ? eScale * m_accECF3(jet) * pow(m_accECF1(jet), 3.0) / pow(m_accECF2(jet), 3.0) : -0.1);
152
153 // Std pile-up info
154 DEF_RETRIEVER1( mu, jetInfo.mu()*eScale );
155 DEF_RETRIEVER1( NPV, jetInfo.NPV()*eScale );
156
157 #undef DEF_RETRIEVER
158
160 GlobalLargeRDNNCalibration::VarRetriever* buildVarRetriever(const std::string & name){
161 // create a map of known specialized VarRetriever.
162 // it's just a map "name" <-> function returning a Var_xyz()
163 static const std::map<std::string, std::function<GlobalLargeRDNNCalibration::VarRetriever*()> > knownVar{
164 {"eta", [](){return new Var_eta();} },
165 {"rapidity", [](){return new Var_rapidity();} },
166 {"log_e", [](){return new Var_log_e();} },
167 {"log_m", [](){return new Var_log_m();} },
168 {"log_m_40", [](){return new Var_log_m_40();} },
169 {"Tau21_wta", [](){return new Ratio_Tau21_wta();} },
170 {"Tau32_wta", [](){return new Ratio_Tau32_wta();} },
171 {"C2", [](){return new Ratio_C2();} },
172 {"D2", [](){return new Ratio_D2();} },
173 {"mu", [](){return new Var_mu();} },
174 {"NPV", [](){return new Var_NPV();} },
175 };
176
177 auto it = knownVar.find(name);
178 // if name is not a known variable, assume it's a jet attribute, so return a generic VarAccessorRetriever
179 if( it == knownVar.end() ) return new VarAccessorRetriever(name);
180 // else we just return an instance of a known VarRetriever class
181 // (it->second is the function : we call it to obtain a new pointer)
182 return it->second();
183 }
184
185}
186
189 : JetCalibrationStep::JetCalibrationStep("GlobalLargeRDNNCalibration/GlobalLargeRDNNCalibration"),
190 m_config(nullptr), m_calibArea("")
191{
192}
193
195 : JetCalibrationStep::JetCalibrationStep(name.c_str()),
196 m_config(nullptr), m_calibArea("")
197{
198}
199
200GlobalLargeRDNNCalibration::GlobalLargeRDNNCalibration(const std::string& name, TEnv * config, const TString& calibArea, bool dev)
201 : JetCalibrationStep::JetCalibrationStep( name.c_str() ),
202 m_config(config), m_calibArea(calibArea), m_devMode(dev)
203{
204}
205
210
211// Initialize
213 ATH_MSG_DEBUG("Initializing tool");
214 if ( !m_config ) { ATH_MSG_FATAL("Config file not specified. Aborting."); return StatusCode::FAILURE; }
215
216 // Get list of input features
217 m_NNInputs = JetCalibUtils::Vectorize( m_config->GetValue("DNNC.Inputs","") );
218 // Now build a VarRetriever for each of the input features
219 m_varretrievers.resize(m_NNInputs.size());
220 ATH_MSG_DEBUG("DNN inputs");
221 for (long unsigned int i=0;i<m_NNInputs.size();i++) {
222 m_varretrievers[i] = buildVarRetriever( m_NNInputs[i].Data() );
223 ATH_MSG_DEBUG(" " << m_NNInputs[i]);
224 }
225
226 // Get normalization constants for input features
227 m_eScales = JetCalibUtils::VectorizeD( m_config->GetValue("DNNC.EScales","") );
228 m_NormOffsets = JetCalibUtils::VectorizeD( m_config->GetValue("DNNC.NormOffsets","") );
229 m_NormScales = JetCalibUtils::VectorizeD( m_config->GetValue("DNNC.NormScales","") );
230
231 if (m_eScales.size()!=m_NNInputs.size() || m_NormOffsets.size()!=m_NNInputs.size() || m_NormScales.size()!=m_NNInputs.size()) {
232 ATH_MSG_FATAL("Misconfiguration of config file : not same number of offset/scale parameters and number of features. Will exit");
233 return StatusCode::FAILURE;
234 }
235
236 if( msgLvl(MSG::DEBUG) ){
237 ATH_MSG_DEBUG("m_NormOffsets size : " << m_NormOffsets.size());
238 ATH_MSG_DEBUG("m_NormOffsets");
239 for (long unsigned int i=0;i<m_NormOffsets.size();i++) {
240 ATH_MSG_DEBUG(" " << m_NormOffsets[i]);
241 }
242 ATH_MSG_DEBUG("m_NormScales size : " << m_NormScales.size());
243 ATH_MSG_DEBUG("m_NormScales");
244 for (long unsigned int i=0;i<m_NormScales.size();i++) {
245 ATH_MSG_DEBUG(" " << m_NormScales[i]);
246 }
247 }
248
249 // Get DNN config file
250 m_modelFileName = m_config->GetValue("DNNC.ONNXInput","");
251 m_noMassCalibBelow40 = m_config->GetValue("DNNC.NoMassCalibBelow40", 1);
252 std::string modelPath = "";
253 if (m_devMode) {
254 modelPath="JetCalibTools/"+m_modelFileName;
255 } else {
256 modelPath="JetCalibTools/"+m_calibArea+"CalibrationConfigs/"+m_modelFileName;
257 }
258 const std::string fullModelPath = PathResolverFindCalibFile( modelPath ); // Full path
259 ATH_MSG_INFO("Using ONNX model : " << m_modelFileName );
260 ATH_MSG_INFO("resolved in: " << fullModelPath);
261
262 // Set up the ONNX Runtime session.
263 m_session = CreateORTSession(fullModelPath);
264 ATH_MSG_DEBUG( "ONNX Runtime session succesfully created" );
265
266
267 /************************** Input Nodes *****************************/
268 /*********************************************************************/
269 std::tuple<std::vector<int64_t>, std::vector<const char*> > inputInfo = GetInputNodeInfo(m_session);
270 m_input_node_dims = std::get<0>(inputInfo);
271 m_input_node_names = std::get<1>(inputInfo);
272
273 if( msgLvl(MSG::DEBUG) ){
274 for( std::size_t i = 0; i < m_input_node_names.size(); i++ ) {
275 // print input node names
276 ATH_MSG_DEBUG("Input "<<i<<" : "<<" name= "<<m_input_node_names[i]);
277
278 // print input shapes/dims
279 ATH_MSG_DEBUG("Input "<<i<<" : num_dims= "<<m_input_node_dims.size());
280 for (std::size_t j = 0; j < m_input_node_dims.size(); j++){
281 ATH_MSG_DEBUG("Input "<<i<<" : dim "<<j<<"= "<<m_input_node_dims[j]);
282 }
283 }
284 }
285
286 /************************** Output Nodes *****************************/
287 /*********************************************************************/
288 std::tuple<std::vector<int64_t>, std::vector<const char*> > outputInfo = GetOutputNodeInfo(m_session);
289 m_output_node_dims = std::get<0>(outputInfo);
290 m_output_node_names = std::get<1>(outputInfo);
291
292 if( msgLvl(MSG::DEBUG) ){
293 for( std::size_t i = 0; i < m_output_node_names.size(); i++ ) {
294 // print input node names
295 ATH_MSG_DEBUG("Output "<<i<<" : "<<" name= "<<m_output_node_names[i]);
296
297 // print input shapes/dims
298 ATH_MSG_DEBUG("Output "<<i<<" : num_dims= "<<m_output_node_dims.size());
299 for (std::size_t j = 0; j < m_output_node_dims.size(); j++){
300 ATH_MSG_DEBUG("Output "<<i<<" : dim "<<j<<"= "<<m_output_node_dims[j]);
301 }
302 }
303 }
304
305 /**************************************************************************************
306 * m_input_node_dims[0] = -1; -1 needs to be replaced by the batch size; for no batch --> 1
307 * m_input_node_dims[1] should be equal to m_NNInputs.size()
308 ****************************************************************************************/
309 m_input_node_dims[0] = 1;
310 m_output_node_dims[0] = 1;
311
312 if (m_NNInputs.size()!=(long unsigned int)m_input_node_dims[1]) {
313 ATH_MSG_FATAL("DNN input features not the same size as in config, will exit");
314 return StatusCode::FAILURE;
315 }
316
317 // Set jet starting scale
318 m_jetStartScale = "JetConstitScaleMomentum";
319
320 return StatusCode::SUCCESS;
321}
322
323
324
326
327 // Set jet initial scale
328 xAOD::JetFourMom_t jetStartP4;
330 jetStartP4 = jet.jetP4();
331
332 // Don't apply calibration for jets with negative or null mass or for one constituent jets
333 if(jet.m()<=0 || jet.numConstituents()==1){
334 jet.setAttribute<xAOD::JetFourMom_t>("JetDNNCScaleMomentum",jetStartP4);
335 return StatusCode::SUCCESS;
336 }
337
338 // Get input features normalized for jet
339 std::vector<float> input_tensor_values = getJetFeatures(jet, jetEventInfo);
340 if( msgLvl(MSG::DEBUG) ){
341 ATH_MSG_DEBUG("Input tensor values : ");
342 for (long unsigned int i=0;i<input_tensor_values.size();i++) ATH_MSG_DEBUG(" " << input_tensor_values[i]);
343 }
344 ATH_MSG_DEBUG(" start M : " << jetStartP4.M());
345
346 // Check for nan or +/- inf values
347 int nNan = std::count_if(input_tensor_values.begin(), input_tensor_values.end(), [](float f){return std::isnan(f) || std::isinf(f);});
348 if (nNan>0) {
349 ATH_MSG_WARNING("Encountered Nan or inf value in input features, will not apply calibration");
350 jet.setAttribute<xAOD::JetFourMom_t>("JetDNNCScaleMomentum",jetStartP4);
351 return StatusCode::SUCCESS;
352 }
353
354 // Convert input_tensor_values array to onnx-compatible tensor
355 Ort::MemoryInfo memory_info = Ort::MemoryInfo::CreateCpu(OrtArenaAllocator, OrtMemTypeCPU);
356 Ort::Value input_tensor = Ort::Value::CreateTensor<float>( memory_info,
357 input_tensor_values.data(),
358 input_tensor_values.size(),
359 m_input_node_dims.data(),
360 m_input_node_dims.size());
361
362 // Make sure we get the same input values in tensor
363 std::vector<float> vec(input_tensor.GetTensorMutableData<float>(), input_tensor.GetTensorMutableData<float>() + m_input_node_dims[1]);
364 if (vec!=input_tensor_values) {
365 ATH_MSG_WARNING("Input tensor after convertion to Ort tensor is not the same as the input vector, will not apply calibration");
366 jet.setAttribute<xAOD::JetFourMom_t>("JetDNNCScaleMomentum",jetStartP4);
367 return StatusCode::SUCCESS;
368 }
369
370 // Run inference on input_tensor
371 Ort::Session& session ATLAS_THREAD_SAFE = *m_session;
372 auto output_tensor = session.Run( Ort::RunOptions{nullptr},
373 m_input_node_names.data(),
374 &input_tensor,
375 m_input_node_names.size(),
376 m_output_node_names.data(),
377 m_output_node_names.size());
378 if (!output_tensor.front().IsTensor() || output_tensor.size() != m_output_node_names.size() || output_tensor.front().GetTensorTypeAndShapeInfo().GetShape() != m_output_node_dims) {
379 ATH_MSG_WARNING("Output tensor does not have the same size as output layer, will not apply calibration");
380 jet.setAttribute<xAOD::JetFourMom_t>("JetDNNCScaleMomentum",jetStartP4);
381 return StatusCode::SUCCESS;
382 }
383
384 // Get pointer to output tensor float values
385 float* outputE = output_tensor.at(0).GetTensorMutableData<float>();
386 // Models with a single output node (e.g. small-R JES-only DNNs) only predict the
387 // energy response: there is no mass response to read
388 const bool energyOnly = (output_tensor.size() == 1);
389 float* outputM = energyOnly ? outputE : output_tensor.at(1).GetTensorMutableData<float>();
390
391 // Get predicted calibration factors
392 float predRespE = outputE[0]; // first element is predicted response
393 float predRespM = outputM[0];
394
395 // Print the output predictions for E/M
396 ATH_MSG_DEBUG("Output E : " << predRespE);
397 ATH_MSG_DEBUG("Output M : " << predRespM);
398
399 if (predRespE==0 || predRespM==0) {
400 ATH_MSG_WARNING("Predictions give 0 values, will not apply calibration");
401 jet.setAttribute<xAOD::JetFourMom_t>("JetDNNCScaleMomentum",jetStartP4);
402 return StatusCode::SUCCESS;
403 }
404
405 // Energy-only models: scale the whole 4-vector by the energy response (so the mass
406 // scales with it, without the m > 40 GeV requirement applied to large-R jets below)
407 if (energyOnly) {
408 const xAOD::JetFourMom_t calibP4 = jetStartP4 * (1. / predRespE);
409 jet.setAttribute<xAOD::JetFourMom_t>("JetDNNCScaleMomentum",calibP4);
410 jet.setJetP4( calibP4 );
411 return StatusCode::SUCCESS;
412 }
413
414 // Apply calibration to jet p4
415 float calibE = jetStartP4.e() / predRespE;
416
417 // For mass, only apply calibration if m>40 GeV (if m_noMassCalibBelow40)
418 float calibM = jetStartP4.mass();
419 if( ! m_noMassCalibBelow40 || (calibM>40000) ){
420 calibM /= predRespM;
421 }
422
423 // Propagate energy and mass calibration to jet pT
424 float calibpT = std::sqrt( calibE*calibE - calibM*calibM )/std::cosh( jetStartP4.eta() );
425
426 // Build calibrated jet p4
427 TLorentzVector TLVjet;
428 TLVjet.SetPtEtaPhiM( calibpT, jetStartP4.eta(), jetStartP4.phi(), calibM );
429 xAOD::JetFourMom_t calibP4;
430 calibP4.SetPxPyPzE( TLVjet.Px(), TLVjet.Py(), TLVjet.Pz(), TLVjet.E() );
431
432 // Transfer calibrated jet properties to the Jet object
433 jet.setAttribute<xAOD::JetFourMom_t>("JetDNNCScaleMomentum",calibP4);
434 jet.setJetP4( calibP4 );
435
436 return StatusCode::SUCCESS;
437
438}
439
440
441
442std::vector<float> GlobalLargeRDNNCalibration::getJetFeatures( xAOD::Jet& jet_reco, JetEventInfo& jetEventInfo) const {
443 // Init input tensor
444 std::vector<float> input_tensor_values(m_NNInputs.size());
445
446 // Retrieve all input variables from the jet and/or jetEventInfo using our VarRetriever collection:
447 for(size_t i=0;i<input_tensor_values.size();i++){
448 float v = m_varretrievers[i]->value(jet_reco, jetEventInfo, m_eScales[i]);
449 // and perform normalisation :
450 input_tensor_values[i] = v*m_NormScales[i] + m_NormOffsets[i];
451 }
452
453 return input_tensor_values;
454}
Scalar eta() const
pseudorapidity method
#define ATH_CHECK
Evaluate an expression and check for errors.
#define ATH_MSG_DEBUG(x,...)
#define ATH_MSG_WARNING(x,...)
#define ATH_MSG_INFO(x,...)
#define ATH_MSG_FATAL(x,...)
@ Data
Definition BaseObject.h:11
std::vector< size_t > vec
macros for messaging and checking status codes
#define ANA_MSG_WARNING(xmsg,...)
Macro printing warning messages.
#define DEF_RETRIEVER0(cname, expr)
Define shortcuts macro to declare specialized VarRetriever class in one line.
#define DEF_RETRIEVER1(cname, expr)
#define DEF_RATIO_RETRIEVER(cname, expr)
std::string PathResolverFindCalibFile(const std::string &logical_file_name)
std::atomic_flag m_initialized ATLAS_THREAD_SAFE
Messaging initialized (initMessaging).
virtual StatusCode calibrate(xAOD::Jet &jet, JetEventInfo &) const override
std::vector< float > getJetFeatures(xAOD::Jet &jet_reco, JetEventInfo &jetEventInfo) const
Returns a vector of input features for the NN.
std::vector< const char * > m_output_node_names
std::vector< int64_t > m_output_node_dims
std::vector< VarRetriever * > m_varretrievers
std::unique_ptr< Ort::Session > m_session
std::vector< int64_t > m_input_node_dims
virtual StatusCode initialize() override
Returns the charged fraction of a jet.
virtual ~GlobalLargeRDNNCalibration()
The destructor.
std::vector< const char * > m_input_node_names
virtual StatusCode setStartP4(xAOD::Jet &jet) const
JetCalibrationStep(const char *name="JetCalibrationStep")
bool msgLvl(const MSG::Level lvl) const
Test the output level of the object.
AthROOTErrorHandlerSvc * svc
StrV Vectorize(const TString &str, const TString &sep=" ")
VecD VectorizeD(const TString &str, const TString &sep=" ")
Jet_v1 Jet
Definition of the current "jet version".
ROOT::Math::LorentzVector< ROOT::Math::PtEtaPhiM4D< double > > JetFourMom_t
Base 4 Momentum type for Jet.
Definition JetTypes.h:17
VarRetriever is a generic class to access Jet and/or JetEventInfo variables.
virtual float value(const xAOD::Jet &jet, JetEventInfo &jetInfo, double eScale)=0
the value of the variable to be retrieved from the jet and/or JetEventInfo