@@ -317,12 +317,18 @@ struct HfTreeCreatorDplusToPiKPi {
317317
318318 std::vector<float > outputMl = {-999 ., -999 .};
319319 if constexpr (DoMl) {
320- for (unsigned int iclass = 0 ; iclass < classMlIndexes->size (); iclass++) {
321- outputMl[iclass] = candidate.mlProbDplusToPiKPi ()[classMlIndexes->at (iclass)];
320+ if (candidate.mlProbDplusToPiKPi ().empty ()) {
321+ rowCandidateMl (
322+ -999 .,
323+ -999 .);
324+ } else {
325+ for (unsigned int iclass = 0 ; iclass < classMlIndexes->size (); iclass++) {
326+ outputMl[iclass] = candidate.mlProbDplusToPiKPi ()[classMlIndexes->at (iclass)];
327+ }
328+ rowCandidateMl (
329+ outputMl[0 ],
330+ outputMl[1 ]);
322331 }
323- rowCandidateMl (
324- outputMl[0 ],
325- outputMl[1 ]);
326332 }
327333
328334 float cent{-1 .};
@@ -494,6 +500,35 @@ struct HfTreeCreatorDplusToPiKPi {
494500
495501 PROCESS_SWITCH (HfTreeCreatorDplusToPiKPi, processData, " Process data" , true );
496502
503+ void processDataWMl (aod::Collisions const & collisions,
504+ soa::Filtered<soa::Join<aod::HfCand3ProngWPidPiKa, aod::HfSelDplusToPiKPi, aod::HfMlDplusToPiKPi>> const & candidates,
505+ TracksWPid const &)
506+ {
507+ // Filling event properties
508+ rowCandidateFullEvents.reserve (collisions.size ());
509+ for (const auto & collision : collisions) {
510+ fillEvent (collision, 0 , 1 );
511+ }
512+
513+ // Filling candidate properties
514+ if (fillCandidateLiteTable) {
515+ rowCandidateLite.reserve (candidates.size ());
516+ } else {
517+ rowCandidateFull.reserve (candidates.size ());
518+ }
519+ for (const auto & candidate : candidates) {
520+ if (downSampleBkgFactor < 1 .) {
521+ float const pseudoRndm = candidate.ptProng0 () * 1000 . - static_cast <int64_t >(candidate.ptProng0 () * 1000 );
522+ if (candidate.pt () < ptMaxForDownSample && pseudoRndm >= downSampleBkgFactor) {
523+ continue ;
524+ }
525+ }
526+ fillCandidateTable<aod::Collisions, false , true >(candidate);
527+ }
528+ }
529+
530+ PROCESS_SWITCH (HfTreeCreatorDplusToPiKPi, processDataWMl, " Process data with ML" , false );
531+
497532 void processDataWCent (CollisionsCent const & collisions,
498533 soa::Filtered<soa::Join<aod::HfCand3ProngWPidPiKa, aod::HfSelDplusToPiKPi>> const & candidates,
499534 TracksWPid const &)
@@ -523,6 +558,35 @@ struct HfTreeCreatorDplusToPiKPi {
523558
524559 PROCESS_SWITCH (HfTreeCreatorDplusToPiKPi, processDataWCent, " Process data with cent" , false );
525560
561+ void processDataWCentMl (CollisionsCent const & collisions,
562+ soa::Filtered<soa::Join<aod::HfCand3ProngWPidPiKa, aod::HfSelDplusToPiKPi, aod::HfMlDplusToPiKPi>> const & candidates,
563+ TracksWPid const &)
564+ {
565+ // Filling event properties
566+ rowCandidateFullEvents.reserve (collisions.size ());
567+ for (const auto & collision : collisions) {
568+ fillEvent (collision, 0 , 1 );
569+ }
570+
571+ // Filling candidate properties
572+ if (fillCandidateLiteTable) {
573+ rowCandidateLite.reserve (candidates.size ());
574+ } else {
575+ rowCandidateFull.reserve (candidates.size ());
576+ }
577+ for (const auto & candidate : candidates) {
578+ if (downSampleBkgFactor < 1 .) {
579+ float const pseudoRndm = candidate.ptProng0 () * 1000 . - static_cast <int64_t >(candidate.ptProng0 () * 1000 );
580+ if (candidate.pt () < ptMaxForDownSample && pseudoRndm >= downSampleBkgFactor) {
581+ continue ;
582+ }
583+ }
584+ fillCandidateTable<CollisionsCent, false , true >(candidate);
585+ }
586+ }
587+
588+ PROCESS_SWITCH (HfTreeCreatorDplusToPiKPi, processDataWCentMl, " Process data with cent and ML" , false );
589+
526590 template <bool ApplyMl = false , typename CandTypeMcRec, typename CandTypeMcGen, typename CollType>
527591 void fillMcTables (CollType const & collisions,
528592 aod::McCollisions const &,
0 commit comments