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 );
33 std::string serviceName;
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)");
39 ANA_MSG_WARNING(
" from AnaAlgorithm.DualUseConfig import createService");
40 ANA_MSG_WARNING(
" onnxSvc = createService('AthOnnx::OnnxRuntimeSvc', 'OnnxRuntimeSvc', myAlgSequence)");
41 serviceName =
"OnnxRuntimeSvc";
43 serviceName =
"AthOnnx::OnnxRuntimeSvc";
48 return std::make_unique<Ort::Session>(
svc->env(),
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++ ) {
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();
66 input_node_dims = tensor_info.GetShape();
68 return std::make_tuple(input_node_dims, input_node_names);
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;
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;
82 Ort::TypeInfo type_info = session->GetOutputTypeInfo(i);
83 auto tensor_info = type_info.GetTensorTypeAndShapeInfo();
85 output_node_dims = tensor_info.GetShape();
87 return std::make_tuple(output_node_dims, output_node_names);
105 VarAccessorRetriever(
const std::string &n): m_acc(n) {}
108 return m_acc(
jet) * eScale;
111 SG::AuxElement::ConstAccessor<float> m_acc;
117 RatioAccessorRetriever(): m_accTau1(
"Tau1_wta"),
118 m_accTau2(
"Tau2_wta"),
119 m_accTau3(
"Tau3_wta"),
124 virtual float value(
const xAOD::Jet& jet, JetEventInfo&,
double eScale) = 0;
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;
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 ; } }
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();} },
179 if( it ==
knownVar.end() )
return new VarAccessorRetriever(name);
214 if ( !
m_config ) {
ATH_MSG_FATAL(
"Config file not specified. Aborting.");
return StatusCode::FAILURE; }
221 for (
long unsigned int i=0;i<
m_NNInputs.size();i++) {
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;
252 std::string modelPath =
"";
263 m_session = CreateORTSession(fullModelPath);
269 std::tuple<std::vector<int64_t>, std::vector<const char*> > inputInfo = GetInputNodeInfo(
m_session);
288 std::tuple<std::vector<int64_t>, std::vector<const char*> > outputInfo = GetOutputNodeInfo(
m_session);
313 ATH_MSG_FATAL(
"DNN input features not the same size as in config, will exit");
314 return StatusCode::FAILURE;
320 return StatusCode::SUCCESS;
330 jetStartP4 =
jet.jetP4();
333 if(
jet.m()<=0 ||
jet.numConstituents()==1){
335 return StatusCode::SUCCESS;
342 for (
long unsigned int i=0;i<input_tensor_values.size();i++)
ATH_MSG_DEBUG(
" " << input_tensor_values[i]);
347 int nNan = std::count_if(input_tensor_values.begin(), input_tensor_values.end(), [](
float f){return std::isnan(f) || std::isinf(f);});
349 ATH_MSG_WARNING(
"Encountered Nan or inf value in input features, will not apply calibration");
351 return StatusCode::SUCCESS;
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(),
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");
367 return StatusCode::SUCCESS;
372 auto output_tensor = session.Run( Ort::RunOptions{
nullptr},
379 ATH_MSG_WARNING(
"Output tensor does not have the same size as output layer, will not apply calibration");
381 return StatusCode::SUCCESS;
385 float* outputE = output_tensor.at(0).GetTensorMutableData<
float>();
388 const bool energyOnly = (output_tensor.size() == 1);
389 float* outputM = energyOnly ? outputE : output_tensor.at(1).GetTensorMutableData<
float>();
392 float predRespE = outputE[0];
393 float predRespM = outputM[0];
399 if (predRespE==0 || predRespM==0) {
400 ATH_MSG_WARNING(
"Predictions give 0 values, will not apply calibration");
402 return StatusCode::SUCCESS;
410 jet.setJetP4( calibP4 );
411 return StatusCode::SUCCESS;
415 float calibE = jetStartP4.e() / predRespE;
418 float calibM = jetStartP4.mass();
424 float calibpT = std::sqrt( calibE*calibE - calibM*calibM )/std::cosh( jetStartP4.eta() );
427 TLorentzVector TLVjet;
428 TLVjet.SetPtEtaPhiM( calibpT, jetStartP4.eta(), jetStartP4.phi(), calibM );
430 calibP4.SetPxPyPzE( TLVjet.Px(), TLVjet.Py(), TLVjet.Pz(), TLVjet.E() );
434 jet.setJetP4( calibP4 );
436 return StatusCode::SUCCESS;
444 std::vector<float> input_tensor_values(
m_NNInputs.size());
447 for(
size_t i=0;i<input_tensor_values.size();i++){
453 return input_tensor_values;
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,...)
std::vector< size_t > vec
#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< double > m_NormScales
bool m_noMassCalibBelow40
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
GlobalLargeRDNNCalibration()
The constructor.
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
std::string m_modelFileName
virtual StatusCode initialize() override
Returns the charged fraction of a jet.
std::vector< double > m_eScales
std::vector< double > m_NormOffsets
virtual ~GlobalLargeRDNNCalibration()
The destructor.
std::vector< const char * > m_input_node_names
std::vector< TString > m_NNInputs
std::string m_jetStartScale
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.
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
virtual ~VarRetriever()=default