16float logNonzero(
double raw) {
17 return raw == 0. ? 0.f :
static_cast<float>(std::log(std::max(raw, 1e-8)));
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) {
31 if (sequence.funcs[iVar](tau, *constituents[iObject], value)) {
32 tensor[iObject * nVariables + iVar] =
33 (
value + sequence.offsets[iVar]) * sequence.scales[iVar];
50 out = cluster.
clust().
eta() - tau.eta();
70 out = logNonzero(cluster.
clust().
et());
75 out = logNonzero(cluster.
clust().
e());
84 out = logNonzero(track.track()->pt());
89 out = logNonzero(track.track()->e());
94 out = track.track()->z0();
100 if (!acc.isAvailable(*track.track()))
return false;
101 out = acc(*track.track());
113 if (!acc.isAvailable(vertex))
return false;
143 const std::unordered_map<std::string, Func>& funcMap,
145 if (!sequence.
name.empty()) {
147 <<
"' are both built from " << input.collection);
148 return StatusCode::FAILURE;
150 sequence.
name = input.name;
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;
159 sequence.
funcs.push_back(func->second);
160 sequence.
offsets.push_back(variable.offset);
161 sequence.
scales.push_back(variable.scale);
163 ATH_MSG_INFO(
"TausRUs input '" << input.name <<
"' (" << input.collection <<
"): "
164 << input.maxObjects <<
" objects x " << input.variables.size() <<
" variables");
165 return StatusCode::SUCCESS;
170 if (input.collection ==
"clusters") {
172 }
else if (input.collection ==
"tracks") {
174 }
else if (input.collection ==
"vertices") {
176 }
else if (input.collection ==
"seedjet") {
179 ATH_MSG_ERROR(
"Input '" << input.name <<
"' is built from unknown collection '"
180 << input.collection <<
"'");
181 return StatusCode::FAILURE;
184 return StatusCode::SUCCESS;
189 std::vector<xAOD::CaloVertexedTopoCluster>& storage)
const {
190 storage = tau.vertexedClusters();
192 std::vector<const xAOD::CaloVertexedTopoCluster*> selected;
193 selected.reserve(storage.size());
195 selected.push_back(&cluster);
197 std::sort(selected.begin(), selected.end(),
200 return lhs->clust().e() > rhs->clust().e();
207 std::vector<const xAOD::TauTrack*> tracks = tau.allTracks();
209 return track ==
nullptr || track->track() ==
nullptr;
213 return lhs->track()->pt() > rhs->track()->pt();
215 if (tracks.size() >
m_tracks.maxObjects) tracks.resize(
m_tracks.maxObjects);
221 std::vector<const xAOD::Vertex*> selected;
222 selected.reserve(std::min(vertices.
size(),
m_vertices.maxObjects));
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);
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));
246 std::vector<xAOD::CaloVertexedTopoCluster> storage;
256 std::vector<float> tensor(
m_scalars.funcs.size(), 0.f);
257 for (
size_t iVar = 0; iVar <
m_scalars.funcs.size(); ++iVar) {
#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.
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)
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.
TauJet_v3 TauJet
Definition of the current "tau version".
One input node, its variables resolved to their functions.
std::vector< float > scales
std::vector< Func > funcs
std::vector< float > offsets