ATLAS Offline Software
Loading...
Searching...
No Matches
TauTrackRNNClassifier.cxx
Go to the documentation of this file.
1/*
2 Copyright (C) 2002-2026 CERN for the benefit of the ATLAS collaboration
3*/
4
7
10
14
15#include <fstream>
16
17using namespace tauRecTools;
18
19//==============================================================================
20// class TauTrackRNNClassifier
21//==============================================================================
22
23//______________________________________________________________________________
25 : TauRecToolBase(name) {
26}
27
28//______________________________________________________________________________
32
33//______________________________________________________________________________
35{
36 for (const auto& classifier : m_vClassifier){
37 ATH_MSG_INFO("Intialize TauTrackRNNClassifier tool : " << classifier );
38 ATH_CHECK(classifier.retrieve());
39 }
40
41 ATH_CHECK( m_vertexContainerKey.initialize() );
42
43 return StatusCode::SUCCESS;
44}
45
46
47//______________________________________________________________________________
49
51 if (!vertexInHandle.isValid()) {
52 ATH_MSG_ERROR ("Could not retrieve HiveDataObj with key " << vertexInHandle.key());
53 return StatusCode::FAILURE;
54 }
55 const xAOD::VertexContainer* vertexContainer = vertexInHandle.cptr();
56
57 std::vector<xAOD::TauTrack*> vTracks = xAOD::TauHelpers::allTauTracksNonConst(&xTau, &tauTrackCon);
58
59 for (xAOD::TauTrack* xTrack : vTracks) {
60 // reset all track flags and set status to unclassified
61 xTrack->setFlag(xAOD::TauJetParameters::classifiedCharged, false);
63 xTrack->setFlag(xAOD::TauJetParameters::classifiedIsolation, false);
64 xTrack->setFlag(xAOD::TauJetParameters::classifiedFake, false);
65 xTrack->setFlag(xAOD::TauJetParameters::unclassified, true);
66 }
67
68 // Collect the associated tracks from TauTrackFinder and either classify
69 // with dedicated TC or not at all
70 if (!m_classifyLRT) {
71 std::vector<xAOD::TauTrack*> vLRTs;
72 std::vector<xAOD::TauTrack*>::iterator it = vTracks.begin();
73 while(it != vTracks.end()) {
75 vLRTs.push_back(*it);
76 it = vTracks.erase(it);
77 }
78 else {
79 ++it;
80 }
81 }
82
83 // decorate LRTs with default RNN scores
84 for (auto classifier : m_vClassifier) {
85 ATH_CHECK(classifier->classifyTracks(vLRTs, xTau, vertexContainer, tauTrackCon, true));
86 }
87 }
88
89 // With this options, RNN track classifier will only be applied to tracks passing the track quality
90 // requirements currently settled in the TauTrack association step. This option is currently applied only
91 // for Run4 sample production, since towards we might skip completely the track association step and just
92 // use the tracks passing the quality requirements without caring about if the track is in the core
93 // region or in the isolation region
95 std::vector<xAOD::TauTrack*> excludedBadTracks;
96 std::vector<xAOD::TauTrack*>::iterator it = vTracks.begin();
97 while(it != vTracks.end()) {
99 excludedBadTracks.push_back(*it);
100 it = vTracks.erase(it);
101 } else {
102 ++it;
103 }
104 }
105 // decorate excludedTracks with default RNN scores
106 for (auto classifier : m_vClassifier) {
107 ATH_CHECK(classifier->classifyTracks(excludedBadTracks, xTau, vertexContainer, tauTrackCon, true));
108 }
109 }
110
111 // With this options, RNN track classifier will only be applied to a unique set of tracks
112 // without any duplication between different taus (see the dedicated protection in the TauTrackFinder).
113 // This option is currently NOT applied as default, since this will change AODs in R22+ reconstruction
114 // and also can lead to reconstruction inefficiency when compared to current reconstruction.
115 // Put here as this can used for Run4 studies
117 std::vector<xAOD::TauTrack*> excludedTracks;
118 std::vector<xAOD::TauTrack*>::iterator it = vTracks.begin();
119 while(it != vTracks.end()) {
121
122 excludedTracks.push_back(*it);
123 it = vTracks.erase(it);
124 }
125 else {
126 ++it;
127 }
128 }
129 // decorate excludedTracks with default RNN scores
130 for (auto classifier : m_vClassifier) {
131 ATH_CHECK(classifier->classifyTracks(excludedTracks, xTau, vertexContainer, tauTrackCon, true));
132 }
133 }
134
135 // classify tracks
137 ATH_CHECK(classifyLRTTracks(vTracks, xTau));
138 } else {
139 for (auto classifier : m_vClassifier) {
140 ATH_CHECK(classifier->classifyTracks(vTracks, xTau, vertexContainer, tauTrackCon));
141 }
142 }
143
144 std::vector< ElementLink< xAOD::TauTrackContainer > >& tauTrackLinks(xTau.allTauTrackLinksNonConst());
145 std::sort(tauTrackLinks.begin(), tauTrackLinks.end(), sortTracks);
146 float charge=0.0;
148 charge += trk->track()->charge();
149 }
150 xTau.setCharge(charge);
151 xTau.setDetail(xAOD::TauJetParameters::nChargedTracks, static_cast<int>(xTau.nTracks()));
153
154 // decorations for now, may be turned into Aux
155 static const SG::Accessor<int> nTrkConv("nConversionTracks");
156 static const SG::Accessor<int> nTrkFake("nFakeTracks");
157 nTrkConv(xTau) = static_cast<int>(xTau.nTracks(xAOD::TauJetParameters::classifiedConversion));
158 nTrkFake(xTau) = static_cast<int>(xTau.nTracks(xAOD::TauJetParameters::classifiedFake));
159
160 //set modifiedIsolationTrack
161 for (xAOD::TauTrack* xTrack : vTracks) {
164 }
165 else {
167 }
168 }
170
171 return StatusCode::SUCCESS;
172}
173
174//______________________________________________________________________________
175
176StatusCode TauTrackRNNClassifier::classifyLRTTracks(std::vector<xAOD::TauTrack*>& vTracks, xAOD::TauJet& xTau) const {
177 static const SG::Accessor<float> idScoreCharged("rnn_chargedScore");
178 static const SG::Accessor<float> idScoreIso("rnn_isolationScore");
179 static const SG::Accessor<float> idScoreConv("rnn_conversionScore");
180 static const SG::Accessor<float> idScoreFake("rnn_fakeScore");
181 for(xAOD::TauTrack* xTrack : vTracks) {
182 idScoreCharged(*xTrack) = 0.;
183 idScoreConv(*xTrack) = 0.;
184 idScoreIso(*xTrack) = 0.;
185 idScoreFake(*xTrack) = 0.;
186
187 double d0_weight = (xTrack->d0TJVA() ? xTrack->d0SigTJVA() / xTrack->d0TJVA(): 0);
188 double log10_pt_ratio = std::log10(xTrack->pt() / xTau.pt());
189 double abs_d0_sig = std::abs(xTrack->d0SigTJVA());
190 double dR = xTau.p4().DeltaR(xTrack->p4());
191 double log10_rConv = std::log10(xTrack->rConv());
192
193 // Cut values taken from a trained decision tree classifier
194 bool passed = false;
195 if (dR <= 0.20) {
196 if (d0_weight <= 32.19) {
197 // Captures high d0 tracks
198 if (log10_pt_ratio >= 0.08 && abs_d0_sig >= 5.95) {
199 passed = true;
200 }
201 } else {
202 if (log10_rConv >= 1.41) {
203 passed = true;
204 } else {
205 if (dR <= 0.03) {
206 passed = true;
207 }
208 }
209 }
210 }
211
212 if (passed) {
213 xTrack->setFlag(xAOD::TauJetParameters::classifiedCharged, true);
214 xTrack->setFlag(xAOD::TauJetParameters::classifiedFake, false);
215 } else {
216 xTrack->setFlag(xAOD::TauJetParameters::classifiedCharged, false);
217 xTrack->setFlag(xAOD::TauJetParameters::classifiedFake, true);
218 }
219 xTrack->setFlag(xAOD::TauJetParameters::classifiedConversion, false);
220 xTrack->setFlag(xAOD::TauJetParameters::classifiedIsolation, false);
221 xTrack->setFlag(xAOD::TauJetParameters::unclassified, false);
222 }
223 return StatusCode::SUCCESS;
224}
225
226//==============================================================================
227// class TrackRNN
228//==============================================================================
229
230//______________________________________________________________________________
231TrackRNN::TrackRNN(const std::string& name)
232 : TauRecToolBase(name)
233{
234}
235
236//______________________________________________________________________________
240
241//______________________________________________________________________________
243{
244 std::string inputWeightsPath = find_file(m_inputWeightsPath);
245 ATH_MSG_INFO("Using calibration file: " << inputWeightsPath);
246
247 std::ifstream nn_config_istream(inputWeightsPath);
248
249 lwtDev::GraphConfig NNconfig = lwtDev::parse_json_graph(nn_config_istream);
250
251 m_RNNClassifier = std::make_unique<lwtDev::LightweightGraph>(NNconfig, NNconfig.outputs.begin()->first);
252 if(!m_RNNClassifier) {
253 ATH_MSG_FATAL("Couldn't configure neural network!");
254 return StatusCode::FAILURE;
255 }
256
257 return StatusCode::SUCCESS;
258}
259
260//______________________________________________________________________________
261StatusCode TrackRNN::classifyTracks(std::vector<xAOD::TauTrack*>& vTracks,
262 xAOD::TauJet& xTau,
263 const xAOD::VertexContainer* vertexContainer,
264 const xAOD::TauTrackContainer& tauTrackCon,
265 bool skipTracks) const
266{
267 if(vTracks.empty()) {
268 return StatusCode::SUCCESS;
269 }
270
271 static const SG::Accessor<float> idScoreCharged("rnn_chargedScore");
272 static const SG::Accessor<float> idScoreIso("rnn_isolationScore");
273 static const SG::Accessor<float> idScoreConv("rnn_conversionScore");
274 static const SG::Accessor<float> idScoreFake("rnn_fakeScore");
275
276 // don't classify tracks, set default decorations
277 if(skipTracks) {
278 for(xAOD::TauTrack* track : vTracks) {
279 idScoreCharged(*track) = 0.;
280 idScoreConv(*track) = 0.;
281 idScoreIso(*track) = 0.;
282 idScoreFake(*track) = 0.;
283 }
284 return StatusCode::SUCCESS;
285 }
286
287 std::sort(vTracks.begin(), vTracks.end(), [](const xAOD::TauTrack * a, const xAOD::TauTrack * b) {return a->pt() > b->pt();});
288
289 VectorMap valueMap;
290 ATH_CHECK(calculateVars(vTracks, xTau, vertexContainer, valueMap));
291
292 SeqNodeMap seqInput;
293 NodeMap nodeInput;
294
295 seqInput["input_1"] = std::move(valueMap);
296
297 VectorMap mValue = m_RNNClassifier->scan(nodeInput, seqInput, "time_distributed_2");
298
299 std::vector<double> vClassProb(5);
300
301 for (unsigned int i = 0; i < vTracks.size(); ++i){
302
303 if(i >= m_nMaxNtracks && m_nMaxNtracks > 0){
304 vClassProb[0] = 0.0;
305 vClassProb[1] = 0.0;
306 vClassProb[2] = 0.0;
307 vClassProb[3] = 0.0;
308 vClassProb[4] = 1.0;
309 }else{
310 vClassProb[0] = mValue["type_0"][i];
311 vClassProb[1] = mValue["type_1"][i];
312 vClassProb[2] = mValue["type_2"][i];
313 vClassProb[3] = mValue["type_3"][i];
314 vClassProb[4] = 0.0;
315 }
316
317 idScoreCharged(*vTracks[i]) = vClassProb[0];
318 idScoreConv(*vTracks[i]) = vClassProb[1];
319 idScoreIso(*vTracks[i]) = vClassProb[2];
320 idScoreFake(*vTracks[i]) = vClassProb[3];
321
322 int iMaxIndex = 0;
323 for (unsigned int j = 1; j < vClassProb.size(); ++j){
324 if(vClassProb[j] > vClassProb[iMaxIndex]) iMaxIndex = j;
325 }
326
327 if(iMaxIndex < 4) {
328 vTracks[i]->setFlag(xAOD::TauJetParameters::unclassified, false);
329 }
330
331 if(iMaxIndex == 3){
332 vTracks[i]->setFlag(xAOD::TauJetParameters::classifiedFake, true);
333 }else if(iMaxIndex == 0){
334 vTracks[i]->setFlag(xAOD::TauJetParameters::classifiedCharged, true);
335 }else if(iMaxIndex == 1){
336 vTracks[i]->setFlag(xAOD::TauJetParameters::classifiedConversion, true);
337 }else if(iMaxIndex == 2){
338 vTracks[i]->setFlag(xAOD::TauJetParameters::classifiedIsolation, true);
339 }
340 }
341
343 bool alreadyUsed = false;
344 for (unsigned int i = 0; i < vTracks.size(); ++i){
345 alreadyUsed = false;
346 //loop over all up-to-now charged tracks
347 for( const xAOD::TauTrack* tau_trk : tauTrackCon ) {
348 if(!(vTracks[i]->flag(xAOD::TauJetParameters::TauTrackFlag::classifiedCharged))) continue;
349 if( vTracks[i]->track() == tau_trk->track()) alreadyUsed = true;
350 }
351 //if this track has already been used by another tau, don't consider
352 if (alreadyUsed) {
353 ATH_MSG_INFO( "Found Already Used charged track new, now putting it as unclassified" );
354 vTracks[i]->setFlag(xAOD::TauJetParameters::classifiedCharged, false);
356 } else {
357 ++i;
358 }
359 }
360 }
361
362
363 return StatusCode::SUCCESS;
364}
365
366
367//______________________________________________________________________________
368StatusCode TrackRNN::calculateVars(const std::vector<xAOD::TauTrack*>& vTracks,
369 const xAOD::TauJet& xTau,
370 const xAOD::VertexContainer* vertexContainer,
371 tauRecTools::VectorMap& valueMap) const
372{
373 // initialize map with values
374 valueMap.clear();
375 unsigned int n_timeSteps = vTracks.size();
376 if(m_nMaxNtracks > 0 && n_timeSteps > m_nMaxNtracks) {
377 n_timeSteps = m_nMaxNtracks;
378 }
379
380 valueMap["log(trackPt)"] = std::vector<double>(n_timeSteps);
381 valueMap["log(jetSeedPt)"] = std::vector<double>(n_timeSteps);
382 valueMap["trackPt/tauPtIntermediateAxis"] = std::vector<double>(n_timeSteps);
383 valueMap["trackEta"] = std::vector<double>(n_timeSteps);
384 valueMap["z0sinthetaTJVA"] = std::vector<double>(n_timeSteps);
385 valueMap["z0sinthetaSigTJVA"] = std::vector<double>(n_timeSteps);
386 valueMap["log(rConv)"] = std::vector<double>(n_timeSteps);
387 valueMap["tanh(rConvII/500)"] = std::vector<double>(n_timeSteps);
388 valueMap["dRJetSeedAxis"] = std::vector<double>(n_timeSteps);
389 valueMap["dRIntermediateAxis"] = std::vector<double>(n_timeSteps);
390 valueMap["tanh(d0SigTJVA/10)"] = std::vector<double>(n_timeSteps);
391 valueMap["tanh(d0TJVA/10)"] = std::vector<double>(n_timeSteps);
392 valueMap["qOverP*1000"] = std::vector<double>(n_timeSteps);
393 valueMap["numberOfInnermostPixelLayerHits"] = std::vector<double>(n_timeSteps);
394 valueMap["numberOfPixelSharedHits"] = std::vector<double>(n_timeSteps);
395 valueMap["numberOfSCTSharedHits"] = std::vector<double>(n_timeSteps);
396 valueMap["numberOfTRTHits"] = std::vector<double>(n_timeSteps);
397 valueMap["eProbabilityHT"] = std::vector<double>(n_timeSteps);
398 valueMap["nPixHits"] = std::vector<double>(n_timeSteps);
399 valueMap["nSCTHits"] = std::vector<double>(n_timeSteps);
400 valueMap["dz0_TV_PV0"] = std::vector<double>(n_timeSteps);
401 valueMap["log_sumpt_TV"] = std::vector<double>(n_timeSteps);
402 valueMap["log_sumpt2_TV"] = std::vector<double>(n_timeSteps);
403 valueMap["log_sumpt_PV0"] = std::vector<double>(n_timeSteps);
404 valueMap["log_sumpt2_PV0"] = std::vector<double>(n_timeSteps);
405 valueMap["charge"] = std::vector<double>(n_timeSteps);
406 // used by RNN track classifier for upgrade
407 valueMap["(trackPt/jetSeedPt)"] = std::vector<double>(n_timeSteps);
408 valueMap["numberOfInnermostPixelLayerEndcapHits"] = std::vector<double>(n_timeSteps);
409 valueMap["nSiHits"] = std::vector<double>(n_timeSteps);
410
411 // tau variable
412 double log_ptJetSeed = std::log( xTau.ptJetSeed() );
413
414 // vertex variables
415 double dz0_TV_PV0 = 0., sumpt_TV = 0., sumpt2_TV = 0., sumpt_PV0 = 0., sumpt2_PV0 = 0.;
416 if(vertexContainer != nullptr && !vertexContainer->empty() && xTau.vertex()!=nullptr) {
417 dz0_TV_PV0 = xTau.vertex()->z() - vertexContainer->at(0)->z();
418
419 // Some AODs reconstructed before 22.0.48 contain rare cases of tracks with only dead sensors instead of hits
420 // due to an edge case in the Si Hit definitions see e.g https://its.cern.ch/jira/browse/ATLIDTRKCP-395
421 // These could be used in the primary vertexing but then were thinned away by the TRT Standalone thinning
422 // this only checked for nHits < 4 rather than the TRT Standalone bit pattern specifically, removing these only dead sensor tracks
423 // we guard against this unresolved track link issue by counting and skipping invalid links
424 // should be extremely rare, but prevents crashes
425 unsigned int nUnresolved = 0;
426 for (const ElementLink<xAOD::TrackParticleContainer>& trk : vertexContainer->at(0)->trackParticleLinks()) {
427 if (!trk.isValid()) { ++nUnresolved; continue; }
428 sumpt_PV0 += (*trk)->pt();
429 sumpt2_PV0 += pow((*trk)->pt(), 2.);
430 }
432 if (!trk.isValid()) { ++nUnresolved; continue; }
433 sumpt_TV += (*trk)->pt();
434 sumpt2_TV += pow((*trk)->pt(), 2.);
435 }
436 if (nUnresolved > 0) {
437 ATH_MSG_WARNING(nUnresolved << " unresolvable track link(s) on the primary or tau "
438 << "vertex skipped: log_sumpt_PV0 / log_sumpt_TV inputs of "
439 << "the track RNN computed from remaining tracks");
440 }
441 }
442 //these are false positives
443 //cppcheck-suppress invalidFunctionArg
444 double log_sumpt_TV = (sumpt_TV>0.) ? std::log(sumpt_TV) : 0.;
445 //cppcheck-suppress invalidFunctionArg
446 double log_sumpt2_TV = (sumpt2_TV>0.) ? std::log(sumpt2_TV) : 0.;
447 //cppcheck-suppress invalidFunctionArg
448 double log_sumpt_PV0 = (sumpt_PV0>0.) ? std::log(sumpt_PV0) : 0.;
449 //cppcheck-suppress invalidFunctionArg
450 double log_sumpt2_PV0 = (sumpt2_PV0>0.) ? std::log(sumpt2_PV0) : 0.;
451
452 // track variables
453 unsigned int i = 0;
454
455 for(xAOD::TauTrack* xTrack : vTracks)
456 {
457 const xAOD::TrackParticle* xTrackParticle = xTrack->track();
458
459 uint8_t numberOfInnermostPixelLayerHits = 0; ATH_CHECK( xTrackParticle->summaryValue(numberOfInnermostPixelLayerHits, xAOD::numberOfInnermostPixelLayerHits) );
460 uint8_t nPixelHits = 0; ATH_CHECK( xTrackParticle->summaryValue(nPixelHits, xAOD::numberOfPixelHits) );
461 uint8_t nPixelSharedHits = 0; ATH_CHECK( xTrackParticle->summaryValue(nPixelSharedHits, xAOD::numberOfPixelSharedHits) );
462 uint8_t nPixelDeadSensors = 0; ATH_CHECK( xTrackParticle->summaryValue(nPixelDeadSensors, xAOD::numberOfPixelDeadSensors) );
463 uint8_t nSCTHits = 0; ATH_CHECK( xTrackParticle->summaryValue(nSCTHits, xAOD::numberOfSCTHits) );
464 uint8_t nSCTSharedHits = 0; ATH_CHECK( xTrackParticle->summaryValue(nSCTSharedHits, xAOD::numberOfSCTSharedHits) );
465 uint8_t nSCTDeadSensors = 0; ATH_CHECK( xTrackParticle->summaryValue(nSCTDeadSensors, xAOD::numberOfSCTDeadSensors) );
466 uint8_t nTRTHits = 0; ATH_CHECK( xTrackParticle->summaryValue(nTRTHits, xAOD::numberOfTRTHits) );
467 float eProbabilityHT; ATH_CHECK( xTrackParticle->summaryValue( eProbabilityHT, xAOD::eProbabilityHT) );
468
469 // used by RNN track classifier for upgrade
470 uint8_t numberOfInnermostPixelLayerEndcapHits = 0;
471 uint8_t tmp_var = 0;
472 if(xTrackParticle->summaryValue(tmp_var, xAOD::numberOfInnermostPixelLayerEndcapHits) ){
473 numberOfInnermostPixelLayerEndcapHits = tmp_var;
474 }
475 uint8_t nSiHits = nPixelHits + nPixelDeadSensors + nSCTHits + nSCTDeadSensors;
476
477 valueMap["log(trackPt)"][i] = std::log( xTrackParticle->pt() );
478 valueMap["log(jetSeedPt)"][i] = log_ptJetSeed;
479 valueMap["trackPt/tauPtIntermediateAxis"][i] = xTrackParticle->pt()/xTau.ptIntermediateAxis();
480 valueMap["trackEta"][i] = xTrackParticle->eta();
481 valueMap["z0sinthetaTJVA"][i] = xTrack->z0sinthetaTJVA();
482 valueMap["z0sinthetaSigTJVA"][i] = xTrack->z0sinthetaSigTJVA();
483 valueMap["log(rConv)"][i] = std::log( xTrack->rConv() );
484 valueMap["tanh(rConvII/500)"][i] = std::tanh( xTrack->rConvII()/500. );
485 // there is no seed jets in AOD so dRJetSeedAxis wont work
486 valueMap["dRJetSeedAxis"][i] = xTrack->p4().DeltaR(xTau.p4(xAOD::TauJetParameters::JetSeed));
487 valueMap["dRIntermediateAxis"][i] = xTrack->p4().DeltaR( xTau.p4(xAOD::TauJetParameters::IntermediateAxis) );
488 valueMap["tanh(d0SigTJVA/10)"][i] = std::tanh( xTrack->d0SigTJVA()/10. );
489 valueMap["tanh(d0TJVA/10)"][i] = std::tanh( xTrack->d0TJVA()/10. );
490 valueMap["qOverP*1000"][i] = xTrackParticle->qOverP()*1000.;
491 valueMap["numberOfInnermostPixelLayerHits"][i] = numberOfInnermostPixelLayerHits;
492 valueMap["numberOfPixelSharedHits"][i] = nPixelSharedHits;
493 valueMap["numberOfSCTSharedHits"][i] = nSCTSharedHits;
494 valueMap["numberOfTRTHits"][i] = nTRTHits;
495 valueMap["eProbabilityHT"][i] = eProbabilityHT;
496 valueMap["nPixHits"][i] = nPixelHits + nPixelDeadSensors;
497 valueMap["nSCTHits"][i] = nSCTHits + nSCTDeadSensors;
498 valueMap["dz0_TV_PV0"][i] = dz0_TV_PV0;
499 valueMap["log_sumpt_TV"][i] = log_sumpt_TV;
500 valueMap["log_sumpt2_TV"][i] = log_sumpt2_TV;
501 valueMap["log_sumpt_PV0"][i] = log_sumpt_PV0;
502 valueMap["log_sumpt2_PV0"][i] = log_sumpt2_PV0;
503 valueMap["charge"][i] = xTrackParticle->charge();
504 // used by RNN track classifier for upgrade
505 valueMap["(trackPt/jetSeedPt)"][i] = xTrackParticle->pt()/xTau.ptJetSeed();
506 valueMap["numberOfInnermostPixelLayerEndcapHits"][i] = numberOfInnermostPixelLayerEndcapHits;
507 valueMap["nSiHits"][i] = nSiHits;
508
509 ++i;
510 if(m_nMaxNtracks > 0 && i >= m_nMaxNtracks) {
511 break;
512 }
513 }
514
515 return StatusCode::SUCCESS;
516}
#define ATH_CHECK
Evaluate an expression and check for errors.
#define ATH_MSG_ERROR(x,...)
#define ATH_MSG_WARNING(x,...)
#define ATH_MSG_INFO(x,...)
#define ATH_MSG_FATAL(x,...)
Handle class for reading from StoreGate.
double charge(const T &p)
Definition AtlasPID.h:1003
bool passed(DecisionID id, const DecisionIDContainer &)
checks if required decision ID is in the set of IDs in the container
static Double_t a
const T * at(size_type n) const
Access an element, as an rvalue.
bool empty() const noexcept
Returns true if the collection is empty.
Helper class to provide type-safe access to aux data.
virtual bool isValid() override final
Can the handle be successfully dereferenced?
const_pointer_type cptr()
Dereference the pointer.
virtual const std::string & key() const override final
Return the StoreGate ID for the referenced object.
The base class for all tau tools.
TauRecToolBase(const std::string &name)
std::string find_file(const std::string &fname) const
ToolHandleArray< TrackRNN > m_vClassifier
SG::ReadHandleKey< xAOD::VertexContainer > m_vertexContainerKey
TauTrackRNNClassifier(const std::string &name="TauTrackRNNClassifier")
Gaudi::Property< bool > m_classifyLRTWithDedicated
virtual StatusCode executeTrackClassifier(xAOD::TauJet &pTau, xAOD::TauTrackContainer &tauTrackContainer) const override
virtual StatusCode initialize() override
Tool initializer.
StatusCode classifyLRTTracks(std::vector< xAOD::TauTrack * > &vTracks, xAOD::TauJet &xTau) const
Gaudi::Property< bool > m_classifyOnlyCoreTracks
Gaudi::Property< unsigned int > m_nMaxNtracks
Gaudi::Property< bool > m_removeDuplicateChargedTracks
virtual StatusCode initialize() override
Tool initializer.
Gaudi::Property< std::string > m_inputWeightsPath
StatusCode calculateVars(const std::vector< xAOD::TauTrack * > &vTracks, const xAOD::TauJet &xTau, const xAOD::VertexContainer *vertexContainer, VectorMap &valueMap) const
ASG_TOOL_CLASS2(TrackRNN, TauRecToolBase, ITauToolBase) public ~TrackRNN()
Create a proper constructor for Athena.
std::unique_ptr< lwtDev::LightweightGraph > m_RNNClassifier
StatusCode classifyTracks(std::vector< xAOD::TauTrack * > &vTracks, xAOD::TauJet &xTau, const xAOD::VertexContainer *vertexContainer, const xAOD::TauTrackContainer &tauTrackContainer, bool skipTracks=false) const
TauTrackLinks_t & allTauTrackLinksNonConst()
In order to sort track links.
virtual FourMom_t p4() const
The full 4-momentum of the particle.
Definition TauJet_v3.cxx:96
virtual double pt() const
The transverse momentum ( ) of the particle.
void setCharge(float)
double ptIntermediateAxis() const
void setDetail(TauJetParameters::Detail detail, int value)
const Vertex * vertex() const
double ptJetSeed() const
size_t nTracks(TauJetParameters::TauTrackFlag flag=TauJetParameters::TauTrackFlag::classifiedCharged) const
std::vector< const TauTrack * > tracks(TauJetParameters::TauTrackFlag flag=TauJetParameters::TauTrackFlag::classifiedCharged) const
Get the v<const pointer> to a given tauTrack collection associated with this tau.
bool summaryValue(uint8_t &value, const SummaryType &information) const
Accessor for TrackSummary values.
float qOverP() const
Returns the parameter.
virtual double pt() const override final
The transverse momentum ( ) of the particle.
virtual double eta() const override final
The pseudorapidity ( ) of the particle.
float charge() const
Returns the charge.
float z() const
Returns the z position.
const TrackParticleLinks_t & trackParticleLinks() const
Get all the particles associated with the vertex.
GraphConfig parse_json_graph(std::istream &json)
void sort(typename DataModel_detail::iterator< DVL > beg, typename DataModel_detail::iterator< DVL > end)
Specialization of sort for DataVector/List.
Implementation of a TrackClassifier based on an RNN.
Definition BDTHelper.cxx:12
std::map< std::string, std::vector< double > > VectorMap
bool sortTracks(const ElementLink< xAOD::TauTrackContainer > &l1, const ElementLink< xAOD::TauTrackContainer > &l2)
std::map< std::string, ValueMap > NodeMap
std::map< std::string, VectorMap > SeqNodeMap
std::vector< xAOD::TauTrack * > allTauTracksNonConst(const xAOD::TauJet *tau, xAOD::TauTrackContainer *trackCont)
TrackParticle_v1 TrackParticle
Reference the current persistent version:
VertexContainer_v1 VertexContainer
Definition of the current "Vertex container version".
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
TauTrackContainer_v1 TauTrackContainer
Definition of the current TauTrack container version.
@ numberOfInnermostPixelLayerEndcapHits
these are the hits in the 0th pixel layer endcap [unit8_t].
@ numberOfTRTHits
number of TRT hits [unit8_t].
@ numberOfSCTDeadSensors
number of dead SCT sensors crossed [unit8_t].
@ eProbabilityHT
Electron probability from High Threshold (HT) information [float].
@ numberOfSCTHits
number of hits in SCT [unit8_t].
@ numberOfInnermostPixelLayerHits
these are the hits in the 0th pixel barrel layer
@ numberOfPixelHits
these are the pixel hits, including the b-layer [unit8_t].
@ numberOfPixelSharedHits
number of Pixel all-layer hits shared by several tracks [unit8_t].
@ numberOfSCTSharedHits
number of SCT hits shared by several tracks [unit8_t].
@ numberOfPixelDeadSensors
number of dead pixel sensors crossed [unit8_t].
std::map< std::string, OutputNodeConfig > outputs