ATLAS Offline Software
Loading...
Searching...
No Matches
TausRUsDataLoader.cxx
Go to the documentation of this file.
1/*
2 Copyright (C) 2002-2026 CERN for the benefit of the ATLAS collaboration
3*/
4
6
10
11#include <algorithm>
12#include <cmath>
13
14namespace {
15
16float logNonzero(double raw) {
17 return raw == 0. ? 0.f : static_cast<float>(std::log(std::max(raw, 1e-8)));
18}
19
23template <class Sequence, class Constituent>
24std::vector<float> fillTensor(const Sequence& sequence, const xAOD::TauJet& tau,
25 const std::vector<const Constituent*>& constituents) {
26 const size_t nVariables = sequence.funcs.size();
27 std::vector<float> tensor(sequence.maxObjects * nVariables, 0.f);
28 for (size_t iObject = 0; iObject < constituents.size(); ++iObject) {
29 for (size_t iVar = 0; iVar < nVariables; ++iVar) {
30 float value = 0.f;
31 if (sequence.funcs[iVar](tau, *constituents[iObject], value)) {
32 tensor[iObject * nVariables + iVar] =
33 (value + sequence.offsets[iVar]) * sequence.scales[iVar];
34 }
35 }
36 }
37 return tensor;
38}
39
40} // anonymous namespace
41
42// --------------------------------------------------------------------------
43// MARK: TausRUs-only variables
44// --------------------------------------------------------------------------
45
46// uncorrected cluster variables
48
49bool dEtaRaw(const xAOD::TauJet& tau, const xAOD::CaloVertexedTopoCluster& cluster, float& out) {
50 out = cluster.clust().eta() - tau.eta();
51 return true;
52}
53
54bool dPhiRaw(const xAOD::TauJet& tau, const xAOD::CaloVertexedTopoCluster& cluster, float& out) {
55 out = xAOD::P4Helpers::deltaPhi(cluster.clust().phi(), tau.phi());
56 return true;
57}
58
59bool etaRaw(const xAOD::TauJet&, const xAOD::CaloVertexedTopoCluster& cluster, float& out) {
60 out = cluster.clust().eta();
61 return true;
62}
63
64bool phiRaw(const xAOD::TauJet&, const xAOD::CaloVertexedTopoCluster& cluster, float& out) {
65 out = cluster.clust().phi();
66 return true;
67}
68
69bool log_etRaw(const xAOD::TauJet&, const xAOD::CaloVertexedTopoCluster& cluster, float& out) {
70 out = logNonzero(cluster.clust().et());
71 return true;
72}
73
74bool log_eRaw(const xAOD::TauJet&, const xAOD::CaloVertexedTopoCluster& cluster, float& out) {
75 out = logNonzero(cluster.clust().e());
76 return true;
77}
78
79} // namespace TausRUsClusterVars
80
82
83bool log_pt(const xAOD::TauJet&, const xAOD::TauTrack& track, float& out) {
84 out = logNonzero(track.track()->pt());
85 return true;
86}
87
88bool log_e(const xAOD::TauJet&, const xAOD::TauTrack& track, float& out) {
89 out = logNonzero(track.track()->e());
90 return true;
91}
92
93bool z0(const xAOD::TauJet&, const xAOD::TauTrack& track, float& out) {
94 out = track.track()->z0();
95 return true;
96}
97
98bool eProbabilityNN_trackParticle(const xAOD::TauJet&, const xAOD::TauTrack& track, float& out) {
99 static const SG::ConstAccessor<float> acc("eProbabilityNN");
100 if (!acc.isAvailable(*track.track())) return false;
101 out = acc(*track.track());
102 return true;
103}
104
105} // namespace TausRUsTrackVars
106
108
109// A padded vertex slot stays all zero; the model counts a slot as a real
110// vertex if any of its variables is nonzero.
111bool sumPt2(const xAOD::TauJet&, const xAOD::Vertex& vertex, float& out) {
112 static const SG::ConstAccessor<float> acc("sumPt2");
113 if (!acc.isAvailable(vertex)) return false;
114 out = acc(vertex);
115 return true;
116}
117
118bool x(const xAOD::TauJet&, const xAOD::Vertex& vertex, float& out) {
119 out = vertex.x();
120 return true;
121}
122
123bool y(const xAOD::TauJet&, const xAOD::Vertex& vertex, float& out) {
124 out = vertex.y();
125 return true;
126}
127
128bool z(const xAOD::TauJet&, const xAOD::Vertex& vertex, float& out) {
129 out = vertex.z();
130 return true;
131}
132
133} // namespace TausRUsVertexVars
134
135// --------------------------------------------------------------------------
136// MARK: TausRUsDataLoader
137// --------------------------------------------------------------------------
138
140 : asg::AsgMessaging(name) {}
141
142template <class Func> StatusCode TausRUsDataLoader::resolve(const InputConfig& input,
143 const std::unordered_map<std::string, Func>& funcMap,
144 Sequence<Func>& sequence) const {
145 if (!sequence.name.empty()) {
146 ATH_MSG_ERROR("Inputs '" << sequence.name << "' and '" << input.name
147 << "' are both built from " << input.collection);
148 return StatusCode::FAILURE;
149 }
150 sequence.name = input.name;
151 sequence.maxObjects = input.maxObjects;
152 for (const VariableConfig& variable : input.variables) {
153 const auto func = funcMap.find(variable.name);
154 if (func == funcMap.end()) {
155 ATH_MSG_ERROR("Variable '" << variable.name << "' of input '" << input.name
156 << "' is not defined for " << input.collection);
157 return StatusCode::FAILURE;
158 }
159 sequence.funcs.push_back(func->second);
160 sequence.offsets.push_back(variable.offset);
161 sequence.scales.push_back(variable.scale);
162 }
163 ATH_MSG_INFO("TausRUs input '" << input.name << "' (" << input.collection << "): "
164 << input.maxObjects << " objects x " << input.variables.size() << " variables");
165 return StatusCode::SUCCESS;
166}
167
168StatusCode TausRUsDataLoader::initialize(const std::vector<InputConfig>& inputs) {
169 for (const InputConfig& input : inputs) {
170 if (input.collection == "clusters") {
172 } else if (input.collection == "tracks") {
174 } else if (input.collection == "vertices") {
176 } else if (input.collection == "seedjet") {
178 } else {
179 ATH_MSG_ERROR("Input '" << input.name << "' is built from unknown collection '"
180 << input.collection << "'");
181 return StatusCode::FAILURE;
182 }
183 }
184 return StatusCode::SUCCESS;
185}
186
187std::vector<const xAOD::CaloVertexedTopoCluster*> TausRUsDataLoader::selectClusters(
188 const xAOD::TauJet& tau,
189 std::vector<xAOD::CaloVertexedTopoCluster>& storage) const {
190 storage = tau.vertexedClusters();
191
192 std::vector<const xAOD::CaloVertexedTopoCluster*> selected;
193 selected.reserve(storage.size());
194 for (const xAOD::CaloVertexedTopoCluster& cluster : storage) {
195 selected.push_back(&cluster);
196 }
197 std::sort(selected.begin(), selected.end(),
198 [](const xAOD::CaloVertexedTopoCluster* lhs,
200 return lhs->clust().e() > rhs->clust().e();
201 });
202 if (selected.size() > m_clusters.maxObjects) selected.resize(m_clusters.maxObjects);
203 return selected;
204}
205
206std::vector<const xAOD::TauTrack*> TausRUsDataLoader::selectTracks(const xAOD::TauJet& tau) const {
207 std::vector<const xAOD::TauTrack*> tracks = tau.allTracks();
208 std::erase_if(tracks, [](const xAOD::TauTrack* track) {
209 return track == nullptr || track->track() == nullptr;
210 });
211 std::sort(tracks.begin(), tracks.end(),
212 [](const xAOD::TauTrack* lhs, const xAOD::TauTrack* rhs) {
213 return lhs->track()->pt() > rhs->track()->pt();
214 });
215 if (tracks.size() > m_tracks.maxObjects) tracks.resize(m_tracks.maxObjects);
216 return tracks;
217}
218
219std::vector<const xAOD::Vertex*> TausRUsDataLoader::selectVertices(
220 const xAOD::VertexContainer& vertices) const {
221 std::vector<const xAOD::Vertex*> selected;
222 selected.reserve(std::min(vertices.size(), m_vertices.maxObjects));
223 for (const xAOD::Vertex* vertex : vertices) {
224 if (selected.size() == m_vertices.maxObjects) break;
225 if (!vertex) continue;
226 const int type = vertex->vertexType();
227 if (type == 0 || type == -99) continue;
228 selected.push_back(vertex);
229 }
230 return selected;
231}
232
234 const xAOD::VertexContainer& vertices) const {
235 AthInfer::InputDataMap inputData;
236
237 auto add = [&inputData](const auto& sequence, std::vector<float> tensor, bool isScalar) {
238 const auto nVariables = static_cast<int64_t>(sequence.funcs.size());
239 std::vector<int64_t> shape = isScalar
240 ? std::vector<int64_t>{1, nVariables}
241 : std::vector<int64_t>{1, static_cast<int64_t>(sequence.maxObjects), nVariables};
242 inputData[sequence.name] = std::make_pair(std::move(shape), std::move(tensor));
243 };
244
245 if (!m_clusters.name.empty()) {
246 std::vector<xAOD::CaloVertexedTopoCluster> storage;
247 add(m_clusters, fillTensor(m_clusters, tau, selectClusters(tau, storage)), false);
248 }
249 if (!m_tracks.name.empty()) {
250 add(m_tracks, fillTensor(m_tracks, tau, selectTracks(tau)), false);
251 }
252 if (!m_vertices.name.empty()) {
253 add(m_vertices, fillTensor(m_vertices, tau, selectVertices(vertices)), false);
254 }
255 if (!m_scalars.name.empty()) {
256 std::vector<float> tensor(m_scalars.funcs.size(), 0.f);
257 for (size_t iVar = 0; iVar < m_scalars.funcs.size(); ++iVar) {
258 float value = 0.f;
259 if (m_scalars.funcs[iVar](tau, value)) {
260 tensor[iVar] = (value + m_scalars.offsets[iVar]) * m_scalars.scales[iVar];
261 }
262 }
263 add(m_scalars, std::move(tensor), true);
264 }
265 return inputData;
266}
#define ATH_CHECK
Evaluate an expression and check for errors.
#define ATH_MSG_ERROR(x,...)
#define ATH_MSG_INFO(x,...)
Helper class to provide constant type-safe access to aux data.
#define y
#define x
#define z
size_type size() const noexcept
Returns the number of elements in the collection.
Helper class to provide constant type-safe access to aux data.
static const std::unordered_map< std::string, VertexFunc_t > m_vertexFuncs
StatusCode resolve(const InputConfig &input, const std::unordered_map< std::string, Func > &funcMap, Sequence< Func > &sequence) const
StatusCode initialize(const std::vector< InputConfig > &inputs)
Resolve the variables of inputs to their functions.
std::vector< const xAOD::TauTrack * > selectTracks(const xAOD::TauJet &tau) const
The tracks and vertices the input tensors are built from, in slot order, for decoding the per-slot he...
static const std::unordered_map< std::string, ScalarFunc_t > m_scalarFuncs
std::vector< const xAOD::CaloVertexedTopoCluster * > selectClusters(const xAOD::TauJet &tau, std::vector< xAOD::CaloVertexedTopoCluster > &storage) const
Sequence< VertexFunc_t > m_vertices
AthInfer::InputDataMap loadInputs(const xAOD::TauJet &tau, const xAOD::VertexContainer &vertices) const
TausRUsDataLoader(const std::string &name)
static const std::unordered_map< std::string, TrackFunc_t > m_trackFuncs
static const std::unordered_map< std::string, ClusterFunc_t > m_clusterFuncs
Sequence< ClusterFunc_t > m_clusters
Sequence< ScalarFunc_t > m_scalars
Sequence< TrackFunc_t > m_tracks
std::vector< const xAOD::Vertex * > selectVertices(const xAOD::VertexContainer &vertices) const
AsgMessaging(const std::string &name)
Constructor with a name.
virtual double eta() const
The pseudorapidity ( ) of the particle.
virtual double e() const
The total energy of the particle.
virtual double phi() const
The azimuthal angle ( ) of the particle.
const CaloCluster & clust() const
Return the cluster being proxied,.
Evaluate cluster kinematics with a different vertex / signal state.
bool add(const std::string &hname, TKey *tobj)
Definition fastadd.cxx:55
std::map< std::string, InferenceData > InputDataMap
TausRUs input variables that no other tau network uses.
bool dPhiRaw(const xAOD::TauJet &tau, const xAOD::CaloVertexedTopoCluster &cluster, float &out)
bool dEtaRaw(const xAOD::TauJet &tau, const xAOD::CaloVertexedTopoCluster &cluster, float &out)
bool log_etRaw(const xAOD::TauJet &, const xAOD::CaloVertexedTopoCluster &cluster, float &out)
bool log_eRaw(const xAOD::TauJet &, const xAOD::CaloVertexedTopoCluster &cluster, float &out)
bool phiRaw(const xAOD::TauJet &, const xAOD::CaloVertexedTopoCluster &cluster, float &out)
bool etaRaw(const xAOD::TauJet &, const xAOD::CaloVertexedTopoCluster &cluster, float &out)
bool log_e(const xAOD::TauJet &, const xAOD::TauTrack &track, float &out)
bool z0(const xAOD::TauJet &, const xAOD::TauTrack &track, float &out)
bool log_pt(const xAOD::TauJet &, const xAOD::TauTrack &track, float &out)
bool eProbabilityNN_trackParticle(const xAOD::TauJet &, const xAOD::TauTrack &track, float &out)
bool sumPt2(const xAOD::TauJet &, const xAOD::Vertex &vertex, float &out)
void sort(typename DataModel_detail::iterator< DVL > beg, typename DataModel_detail::iterator< DVL > end)
Specialization of sort for DataVector/List.
std::size_t erase_if(T_container &container, T_Func pred)
double deltaPhi(double phiA, double phiB)
delta Phi in range [-pi,pi[
VertexContainer_v1 VertexContainer
Definition of the current "Vertex container version".
Vertex_v1 Vertex
Define the latest version of the vertex class.
TauTrack_v1 TauTrack
Definition of the current version.
Definition TauTrack.h:16
TauJet_v3 TauJet
Definition of the current "tau version".
Definition TauJet.h:17
One input node, its variables resolved to their functions.