ATLAS Offline Software
Loading...
Searching...
No Matches
NNSharingOnnxSvc.cxx
Go to the documentation of this file.
1/*
2 Copyright (C) 2002-2025 CERN for the benefit of the ATLAS collaboration
3*/
4
8
9namespace FlavorTagInference {
10
11 std::shared_ptr<const GNN> NNSharingOnnxSvc::get(
12 const std::string& nn_name,
13 const GNNOptions& opts) {
14 NNHashing::NNKey key{nn_name, opts};
15 if (m_gnns.count(key)) {
16 ATH_MSG_INFO("getting " << nn_name << " from cached NNs");
17 return m_gnns.at(key);
18 } else if (m_base_gnns.count(nn_name) ) {
19 ATH_MSG_INFO("adapting " << nn_name << " from cached NNs, new opts");
20 auto nn = std::make_shared<const GNN>(*m_base_gnns.at(nn_name), opts);
21 m_gnns[key] = nn;
22 return nn;
23 }
24 std::shared_ptr<const GNN> nn;
25 ATH_MSG_INFO("building " << nn_name << " from onnx file on the "
26 << m_executionProvider.value() << " execution provider");
27 SaltModelOptions salt_opts {
28 m_executionProvider.value(), m_deviceId.value(), m_useTF32.value()};
29 ISaltModelPtr salt = std::make_shared<const SaltModel>(
30 PathResolverFindCalibFile(nn_name), salt_opts);
31 nn = std::make_shared<const GNN>(salt, opts);
32 m_base_gnns[nn_name] = nn;
33 m_gnns[key] = nn;
34 return nn;
35 }
36}
#define ATH_MSG_INFO(x)
std::string PathResolverFindCalibFile(const std::string &logical_file_name)
std::unordered_map< NNHashing::NNKey, val_t, NNHashing::NNHasher > m_gnns
Gaudi::Property< std::string > m_executionProvider
std::unordered_map< std::string, val_t > m_base_gnns
virtual std::shared_ptr< const GNN > get(const std::string &nn_name, const GNNOptions &opts) override
This file contains "getter" functions used for accessing tagger inputs from the EDM.
std::shared_ptr< const ISaltModel > ISaltModelPtr
Definition ISaltModel.h:56