11#include <nlohmann/json.hpp>
12#include <onnxruntime_cxx_api.h>
19#include <unordered_map>
24constexpr char TAU_ID_SCORE[] =
"TausRUsTauIDScore";
25constexpr char ELE_REJ_SCORE[] =
"TausRUsEleRejScore";
26constexpr char DECAY_MODE[] =
"TausRUsDecayMode";
27constexpr char DECAY_MODE_SCORE_PREFIX[] =
"TausRUsDecayModeScore";
28constexpr char TAU_P4[] =
"TausRUsTauP4";
29constexpr char CHARGED_PION_P4[] =
"TausRUsChargedPionP4";
30constexpr char NEUTRAL_PION_P4[] =
"TausRUsNeutralPionP4";
31constexpr char VERTEX[] =
"TausRUsVertex";
32constexpr char TAU_CHARGE[] =
"TausRUsTauCharge";
34constexpr char TRACK_CLASS[] =
"TausRUsTrackClass";
35constexpr char TRACK_SCORE_PREFIX[] =
"TausRUsTrackScore";
37constexpr char METADATA_KEY[] =
"metadata";
39constexpr char P4_SUFFIXES[4][5] = {
"_pt",
"_eta",
"_phi",
"_m"};
40constexpr char POSITION_SUFFIXES[3][3] = {
"_x",
"_y",
"_z"};
44float twoClassScore(
float signal,
float background) {
45 return 1.f / (1.f + std::exp(background - signal));
49int argMax(std::span<const float>
scores) {
52 return static_cast<int>(std::distance(
scores.begin(), largest));
55std::string
toString(
const std::vector<int64_t>& dims) {
56 std::string
out =
"(";
57 for (
size_t i = 0;
i < dims.size(); ++
i) {
58 out += (
i ?
", " :
"") + std::to_string(dims[i]);
65std::map<std::string, std::vector<int64_t>> nodeShapes(
const Ort::Session& session,
67 Ort::AllocatorWithDefaultOptions allocator;
68 std::map<std::string, std::vector<int64_t>> shapes;
69 const size_t nNodes = isInput ? session.GetInputCount() : session.GetOutputCount();
70 for (
size_t i = 0;
i < nNodes; ++
i) {
71 const std::string
name = isInput
72 ? session.GetInputNameAllocated(i, allocator).get()
73 : session.GetOutputNameAllocated(i, allocator).get();
74 const Ort::TypeInfo typeInfo = isInput ? session.GetInputTypeInfo(i)
75 : session.GetOutputTypeInfo(i);
76 std::vector<int64_t> dims = typeInfo.GetTensorTypeAndShapeInfo().GetShape();
77 if (!dims.empty() && dims[0] < 0) dims[0] = 1;
78 shapes[
name] = std::move(dims);
101 Ort::Env env(ORT_LOGGING_LEVEL_FATAL,
"");
102 Ort::SessionOptions sessionOptions;
103 sessionOptions.SetIntraOpNumThreads(1);
104 sessionOptions.SetLogSeverityLevel(4);
105 sessionOptions.DisableCpuMemArena();
106 const Ort::Session session(env, path.c_str(), sessionOptions);
108 Ort::AllocatorWithDefaultOptions allocator;
109 const Ort::ModelMetadata modelMetadata = session.GetModelMetadata();
110 const Ort::AllocatedStringPtr metadataString =
111 modelMetadata.LookupCustomMetadataMapAllocated(METADATA_KEY, allocator);
112 if (!metadataString) {
113 ATH_MSG_ERROR(
"Model " << path <<
" has no '" << METADATA_KEY <<
"' metadata");
114 return StatusCode::FAILURE;
117 const auto inputShapes = nodeShapes(session,
true);
118 const auto outputShapes = nodeShapes(session,
false);
120 const nlohmann::json metadata = nlohmann::json::parse(metadataString.get());
122 std::vector<TausRUsDataLoader::InputConfig> inputs;
123 for (
const nlohmann::json&
node : metadata.at(
"inputs")) {
125 input.name =
node.at(
"name").get<std::string>();
126 input.collection =
node.at(
"collection").get<std::string>();
127 input.maxObjects =
node.value(
"max_objects",
size_t{1});
128 for (
const nlohmann::json& variable :
node.at(
"variables")) {
129 input.variables.push_back({variable.at(
"name").get<std::string>(),
130 variable.at(
"offset").get<
float>(),
131 variable.at(
"scale").get<
float>()});
133 if (inputShapes.find(input.name) == inputShapes.end()) {
134 ATH_MSG_ERROR(
"The model has no input node '" << input.name <<
"'");
135 return StatusCode::FAILURE;
137 inputs.push_back(std::move(input));
142 std::map<std::string, nlohmann::json> outputMetadata;
143 for (
const nlohmann::json&
node : metadata.at(
"outputs")) {
145 output.name =
node.at(
"name").get<std::string>();
146 output.type =
node.value(
"type",
"");
147 const auto shape = outputShapes.find(output.name);
148 if (shape == outputShapes.end()) {
149 ATH_MSG_ERROR(
"The model has no output node '" << output.name <<
"'");
150 return StatusCode::FAILURE;
152 output.dims = shape->second;
153 output.size = std::accumulate(output.dims.begin(), output.dims.end(),
size_t{1},
154 std::multiplies<>());
155 outputMetadata[output.name] =
node;
160 m_nDecayModes = outputMetadata.at(
"decay_mode").at(
"num_classes").get<
size_t>();
161 m_nTrackClasses = outputMetadata.at(
"tautrack_class").at(
"num_classes").get<
size_t>();
164 const auto quantiles = outputMetadata.at(
"tes").at(
"quantiles").get<std::vector<double>>();
165 for (
const char* name : {
"charged_pion_pt",
"neutral_pion_pt"}) {
166 if (outputMetadata.at(name).at(
"quantiles").get<std::vector<double>>() != quantiles) {
167 ATH_MSG_ERROR(
"Output '" << name <<
"' regresses other quantiles than 'tes'");
168 return StatusCode::FAILURE;
171 const auto median = std::find_if(quantiles.begin(), quantiles.end(),
172 [](
double q) { return std::abs(q - 0.5) < 1e-6; });
173 if (median == quantiles.end()) {
175 return StatusCode::FAILURE;
182 ATH_MSG_ERROR(
"The model needs both a track and a vertex input");
183 return StatusCode::FAILURE;
186 return StatusCode::SUCCESS;
194 if (modelPath.empty()) {
196 return StatusCode::FAILURE;
198 m_loader = std::make_unique<TausRUsDataLoader>(name() +
"_DataLoader");
204 auto checkNode = [
this](
const std::string& name,
const std::string&
type,
205 const std::vector<int64_t>& dims) -> StatusCode {
207 [&name](
const Output& o) { return o.name == name; });
209 ATH_MSG_ERROR(
"The model metadata has no output '" << name <<
"'");
210 return StatusCode::FAILURE;
212 if (output->type !=
type) {
213 ATH_MSG_ERROR(
"Output '" << name <<
"' is a '" << output->type
214 <<
"' head but the tool decodes it as '" <<
type <<
"'");
215 return StatusCode::FAILURE;
217 if (output->dims != dims) {
218 ATH_MSG_ERROR(
"Output '" << name <<
"' has shape " << toString(output->dims)
219 <<
" but the tool decodes " << toString(dims));
220 return StatusCode::FAILURE;
222 return StatusCode::SUCCESS;
225 ATH_CHECK(checkNode(
"primary",
"classification", {1, 3}));
226 ATH_CHECK(checkNode(
"decay_mode",
"classification",
228 for (
const char* name : {
"tes",
"charged_pion_pt",
"neutral_pion_pt"}) {
229 ATH_CHECK(checkNode(name,
"quantile_regression",
232 for (
const char* name : {
"tau_eta",
"charged_pion_eta",
"neutral_pion_eta"}) {
233 ATH_CHECK(checkNode(name,
"regression", {1, 1}));
235 for (
const char* name : {
"tau_phi",
"charged_pion_phi",
"neutral_pion_phi"}) {
236 ATH_CHECK(checkNode(name,
"periodic_regression", {1, 2}));
238 ATH_CHECK(checkNode(
"vertex_classification",
"classification",
239 {1,
static_cast<int64_t
>(
m_loader->maxVertices())}));
240 ATH_CHECK(checkNode(
"tautrack_class",
"classification",
241 {1,
static_cast<int64_t
>(
m_loader->maxTracks()),
247 + std::to_string(i));
249 for (
const auto& suffix : P4_SUFFIXES) {
250 m_tauP4.emplace_back(std::string(TAU_P4) + suffix);
254 for (
const auto& suffix : POSITION_SUFFIXES) {
258 m_trackScores.emplace_back(std::string(TRACK_SCORE_PREFIX) + std::to_string(i));
263 <<
" on each of its up to " <<
m_loader->maxTracks() <<
" leading tracks");
277 ATH_MSG_ERROR(
"TauTrackContainerName is not set, but the per-track "
279 return StatusCode::FAILURE;
288 return StatusCode::SUCCESS;
298 std::vector<std::string> names{TAU_ID_SCORE, ELE_REJ_SCORE, DECAY_MODE, TAU_CHARGE};
300 names.emplace_back(std::string(DECAY_MODE_SCORE_PREFIX) + std::to_string(iMode));
302 for (
const char*
base : {TAU_P4, CHARGED_PION_P4, NEUTRAL_PION_P4}) {
303 for (
const auto& suffix : P4_SUFFIXES) {
304 names.emplace_back(std::string(
base) + suffix);
307 for (
const auto& suffix : POSITION_SUFFIXES) {
308 names.emplace_back(std::string(VERTEX) + suffix);
314 std::vector<std::string> names{TRACK_CLASS};
316 names.emplace_back(std::string(TRACK_SCORE_PREFIX) + std::to_string(iClass));
341 if (!track)
continue;
351 std::span<const float> ptQuantiles,
352 std::span<const float> etaValues,
353 std::span<const float> phiValues,
357 const float ptJetSeed = tau.ptJetSeed();
358 if (ptJetSeed <= 0.f) {
359 ATH_MSG_DEBUG(
"Seed jet pt is " << ptJetSeed <<
", leaving the regressed "
360 "four-momenta at their defaults");
365 decorators[0](tau) = std::exp(std::log(ptJetSeed) -
response);
366 decorators[1](tau) = etaValues[0];
367 decorators[2](tau) = std::atan2(phiValues[
PHI_SIN], phiValues[
PHI_COS]);
368 decorators[3](tau) = mass;
372 std::span<const float>
scores,
373 const std::vector<const xAOD::Vertex*>& vertices)
const {
377 const size_t nSlots = std::min(vertices.size(),
scores.size());
379 ATH_MSG_DEBUG(
"No vertex to choose from, leaving the vertex position at its default");
383 const int selected = argMax(
scores.first(nSlots));
391 std::span<const float>
scores)
const {
393 const std::vector<const xAOD::TauTrack*> tracks =
m_loader->selectTracks(tau);
396 for (
size_t iTrack = 0; iTrack < tracks.size(); ++iTrack) {
397 const std::span<const float> slot =
399 const int trackClass = argMax(slot);
406 tauCharge +=
static_cast<int>(tracks[iTrack]->track()->
charge());
411 if (tauCharge == -1 || tauCharge == 1) {
423 return StatusCode::SUCCESS;
427 if (!vertexInHandle.
isValid()) {
430 return StatusCode::FAILURE;
437 outputData[output.name] = std::make_pair(output.dims, std::vector<float>{});
443 std::unordered_map<std::string, std::span<const float>> raw;
446 const std::vector<float>& values =
447 std::get<std::vector<float>>(outputData.at(output.name).second);
448 if (values.size() != output.size) {
449 ATH_MSG_ERROR(
"Output '" << output.name <<
"' returned " << values.size()
450 <<
" values but the tool expects " << output.size);
451 return StatusCode::FAILURE;
453 raw[output.name] = values;
456 const std::span<const float> primary = raw.at(
"primary");
460 const std::span<const float> decayMode = raw.at(
"decay_mode");
461 m_decayMode(tau) =
static_cast<float>(argMax(decayMode));
472 raw.at(
"charged_pion_eta"), raw.at(
"charged_pion_phi"),
475 raw.at(
"neutral_pion_eta"), raw.at(
"neutral_pion_phi"),
479 m_loader->selectVertices(*vertexInHandle));
482 return StatusCode::SUCCESS;
#define ATH_CHECK
Evaluate an expression and check for errors.
#define ATH_MSG_DEBUG(x,...)
#define ATH_MSG_ERROR(x,...)
#define ATH_MSG_INFO(x,...)
double charge(const T &p)
std::vector< std::vector< float > > scores
A number of constexpr particle constants to avoid hardcoding them directly in various places.
size_t size() const
Number of registered mappings.
Helper class to provide type-safe access to aux data.
virtual bool isValid() override final
Can the handle be successfully dereferenced?
Property holding a SG store/key/clid/attr name from which a WriteDecorHandle is made.
static constexpr int TAU_TRACK_CLASS
Class for true-tau track from tau_track_class head.
void decorateTracks(xAOD::TauJet &tau, std::span< const float > scores) const
Decorate each track of tau with its slot of the per-track head.
static constexpr float DEFAULT_VALUE
initialized with these default values
virtual StatusCode initialize() override
Tool initializer.
Gaudi::Property< std::string > m_tauContainerName
std::unique_ptr< TausRUsDataLoader > m_loader
virtual ~TausRUsEvaluator()
std::vector< SG::Accessor< float > > FourMomDecorators
FourMomDecorators m_chargedPionP4
Gaudi::Property< std::string > m_modelFile
std::vector< SG::WriteDecorHandleKey< xAOD::TauTrackContainer > > m_trackDecorKeys
void decorateFourMomentum(xAOD::TauJet &tau, const FourMomDecorators &decorators, std::span< const float > ptQuantiles, std::span< const float > etaValues, std::span< const float > phiValues, float mass) const
Decorate the four floats of decorators from a pt, an eta and a phi head, taking the mass from mass.
size_t m_nDecayModes
Read off the metadata at initialize().
virtual StatusCode execute(xAOD::TauJet &tau) const override
Execute - called for each tau candidate.
FourMomDecorators m_neutralPionP4
Gaudi::Property< std::string > m_tauTrackContainerName
std::vector< SG::WriteDecorHandleKey< xAOD::TauJetContainer > > m_decorKeys
StatusCode readModel(const std::string &path)
Configure the data loader and build m_outputs from the metadata and node shapes of the model.
Gaudi::Property< float > m_minTauPt
SG::Accessor< float > m_decayMode
void decorateVertex(xAOD::TauJet &tau, std::span< const float > scores, const std::vector< const xAOD::Vertex * > &vertices) const
Decorate the position of the vertex that scores picks out of vertices, which must be the list the inp...
SG::Accessor< float > m_eleRejScore
SG::Accessor< float > m_tauIDScore
The decorators.
SG::Accessor< float > m_tauCharge
ToolHandle< AthInfer::IAthInferenceTool > m_inferenceTool
TausRUsEvaluator(const std::string &name="TausRUsEvaluator")
FourMomDecorators m_tauP4
void setDefaults(xAOD::TauJet &tau) const
std::vector< SG::Accessor< float > > m_vertexPosition
std::vector< std::string > trackDecorationNames() const
std::vector< std::string > tauDecorationNames() const
Every decoration written on the tau, and every one written on its tracks, which is what the output da...
std::vector< Output > m_outputs
The model's output nodes, from its metadata.
std::vector< SG::Accessor< float > > m_decayModeScores
SG::ReadHandleKey< xAOD::VertexContainer > m_vertexInputContainer
SG::Decorator< int > m_trackClass
static constexpr int DEFAULT_CLASS
std::vector< SG::Decorator< float > > m_trackScores
size_t m_ptMedianQuantile
The pt heads ('tes', 'charged_pion_pt', 'neutral_pion_pt') regress quantiles of the log response,...
std::string toString(const Translation3D &translation, int precision=4)
GeoPrimitvesToStringConverter.
std::map< std::string, InferenceData > OutputDataMap
std::map< std::string, InferenceData > InputDataMap
constexpr double tauMassInMeV
the mass of the tau (in MeV)
constexpr double piZeroMassInMeV
the mass of the pi zero (in MeV)
constexpr double chargedPionMassInMeV
the mass of the charged pion (in MeV)
SG::Decorator< T, ALLOC > Decorator
Helper class to provide type-safe access to aux data, specialized for JaggedVecElt.
Vertex_v1 Vertex
Define the latest version of the vertex class.
TauTrack_v1 TauTrack
Definition of the current version.
TauJet_v3 TauJet
Definition of the current "tau version".