ATLAS Offline Software
Toggle main menu visibility
Loading...
Searching...
No Matches
Reconstruction
tauRecTools
Root
TauGNN.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 "
tauRecTools/TauGNN.h
"
6
7
8
TauGNN::TauGNN
(
const
TauGNNDataLoader::Config
&config):
9
asg
::
AsgMessaging
(
"TauGNN"
),
10
m_saltModel
(
std
::make_shared<
FlavorTagInference
::
SaltModel
>(config.nnFile)),
11
m_dataloader
(
TauGNNDataLoader
(
m_saltModel
, config))
12
{
13
ATH_MSG_INFO
(
"TauGNN object initialized successfully!"
);
14
}
15
16
std::tuple<
17
std::map<std::string, float>,
18
std::map<std::string, std::vector<char>>,
19
std::map<std::string, std::vector<float>> >
20
TauGNN::compute
(
const
xAOD::TauJet
&tau)
const
{
21
ATH_MSG_DEBUG
(
"Computing TauGNN features..."
);
22
auto
salt_model_input_data =
m_dataloader
.loadInputs(&tau);
23
FlavorTagInference::InputMap
input_with_aliases = salt_model_input_data.gnn_inputs;
24
// Support both legacy loader keys and ONNX-hard-coded input tensor names.
25
const
std::map<std::string, std::string> alias_map = {
26
{
"jet_var"
,
"jet_features"
},
27
{
"tracks_r22default_sd0sort"
,
"track_features"
},
28
{
"cells_var"
,
"cell_features"
}
29
};
30
for
(
const
auto
& [legacy_name, onnx_name] : alias_map) {
31
auto
it = salt_model_input_data.gnn_inputs.find(legacy_name);
32
if
(it != salt_model_input_data.gnn_inputs.end()) {
33
input_with_aliases[onnx_name] = it->second;
34
}
35
}
36
ATH_MSG_DEBUG
(
"Running inference..."
);
37
auto
[out_f, out_vc, out_vf] =
m_saltModel
->runInference(input_with_aliases);
38
ATH_MSG_DEBUG
(
"Inference done."
);
39
return
std::make_tuple(out_f, out_vc, out_vf);
40
}
ATH_MSG_DEBUG
#define ATH_MSG_DEBUG(x,...)
Definition
AthMsgStreamMacros.h:43
ATH_MSG_INFO
#define ATH_MSG_INFO(x,...)
Definition
AthMsgStreamMacros.h:45
TauGNN.h
SaltModel
Definition
JetTagPerformanceCalibration/xAODBTaggingEfficiency/xAODBTaggingEfficiency/SaltModel.h:14
TauGNNDataLoader
Definition
TauGNNDataLoader.h:66
TauGNN::m_dataloader
TauGNNDataLoader m_dataloader
Definition
TauGNN.h:39
TauGNN::compute
std::tuple< std::map< std::string, float >, std::map< std::string, std::vector< char > >, std::map< std::string, std::vector< float > > > compute(const xAOD::TauJet &tau) const
Definition
TauGNN.cxx:20
TauGNN::TauGNN
TauGNN(const TauGNNDataLoader::Config &config)
Definition
TauGNN.cxx:8
TauGNN::m_saltModel
std::shared_ptr< const FlavorTagInference::SaltModel > m_saltModel
Definition
TauGNN.h:38
asg::AsgMessaging::AsgMessaging
AsgMessaging(const std::string &name)
Constructor with a name.
Definition
AsgMessaging.cxx:17
FlavorTagInference
This file contains "getter" functions used for accessing tagger inputs from the EDM.
Definition
CaloClusterLoader.h:27
FlavorTagInference::InputMap
std::map< std::string, Inputs, std::less<> > InputMap
Definition
ISaltModel.h:37
asg
Definition
DataHandleTestTool.h:28
std
STL namespace.
xAOD::TauJet
TauJet_v3 TauJet
Definition of the current "tau version".
Definition
TauJet.h:17
TauGNNDataLoader::Config
Definition
TauGNNDataLoader.h:68
Generated on
for ATLAS Offline Software by
1.17.0