ATLAS Offline Software
Loading...
Searching...
No Matches
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
6
7
9 asg::AsgMessaging("TauGNN"),
10 m_saltModel(std::make_shared<FlavorTagInference::SaltModel>(config.nnFile)),
12{
13 ATH_MSG_INFO("TauGNN object initialized successfully!");
14}
15
16std::tuple<
17 std::map<std::string, float>,
18 std::map<std::string, std::vector<char>>,
19 std::map<std::string, std::vector<float>> >
20TauGNN::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}
#define ATH_MSG_DEBUG(x,...)
#define ATH_MSG_INFO(x,...)
TauGNNDataLoader m_dataloader
Definition TauGNN.h:39
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(const TauGNNDataLoader::Config &config)
Definition TauGNN.cxx:8
std::shared_ptr< const FlavorTagInference::SaltModel > m_saltModel
Definition TauGNN.h:38
AsgMessaging(const std::string &name)
Constructor with a name.
This file contains "getter" functions used for accessing tagger inputs from the EDM.
std::map< std::string, Inputs, std::less<> > InputMap
Definition ISaltModel.h:37
STL namespace.
TauJet_v3 TauJet
Definition of the current "tau version".
Definition TauJet.h:17