Skip to content

Commit 4e7e961

Browse files
committed
Switch to MlResponse class for multi charm
1 parent 7e64819 commit 4e7e961

2 files changed

Lines changed: 414 additions & 200 deletions

File tree

ALICE3/ML/MulticharmMlResponse.h

Lines changed: 206 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,206 @@
1+
// Copyright 2019-2020 CERN and copyright holders of ALICE O2.
2+
// See https://alice-o2.web.cern.ch/copyright for details of the copyright holders.
3+
// All rights not expressly granted are reserved.
4+
//
5+
// This software is distributed under the terms of the GNU General Public
6+
// License v3 (GPL Version 3), copied verbatim in the file "COPYING".
7+
//
8+
// In applying this license CERN does not waive the privileges and immunities
9+
// granted to it by virtue of its status as an Intergovernmental Organization
10+
// or submit itself to any jurisdiction.
11+
12+
/// \file MulticharmMlResponse.h
13+
/// \brief Class to compute the ML response for multi-charm candidates
14+
/// \author Jesper Karlsson Gumprecht
15+
16+
#ifndef ALICE3_ML_MULTICHARMRESPONSE_H_
17+
#define ALICE3_ML_MULTICHARMRESPONSE_H_
18+
19+
#include "Tools/ML/MlResponse.h"
20+
21+
#include <array>
22+
#include <cstdint>
23+
#include <string>
24+
#include <vector>
25+
26+
namespace multi_charm_ml
27+
{
28+
29+
static constexpr int NBinsPt = 10;
30+
static constexpr int NClasses = 2;
31+
static constexpr std::array<std::array<double, NClasses>, NBinsPt> Cuts = {{{0.5, 0.5},
32+
{0.5, 0.5},
33+
{0.5, 0.5},
34+
{0.5, 0.5},
35+
{0.5, 0.5},
36+
{0.5, 0.5},
37+
{0.5, 0.5},
38+
{0.5, 0.5},
39+
{0.5, 0.5},
40+
{0.5, 0.5}}};
41+
42+
static const std::vector<std::string> labelsPt = {"pT bin 0",
43+
"pT bin 1",
44+
"pT bin 2",
45+
"pT bin 3",
46+
"pT bin 4",
47+
"pT bin 5",
48+
"pT bin 6",
49+
"pT bin 7",
50+
"pT bin 8",
51+
"pT bin 9"};
52+
53+
static const std::vector<std::string> labelsCutScore = {"score class 1", "score class 2"};
54+
static const std::vector<std::string> namesInputFeatures{"xicDauDCA",
55+
"xiccDauDCA",
56+
"xiDCAxy",
57+
"xicDCAxy",
58+
"xiccDCAxy",
59+
"xiDCAz",
60+
"xicDCAz",
61+
"xiccDCAz",
62+
"pi1cDCAxy",
63+
"pi2cDCAxy",
64+
"piccDCAxy",
65+
"pi1cDCAz",
66+
"pi2cDCAz",
67+
"piccDCAz",
68+
"xicDecayRadius2D",
69+
"xiccDecayRadius2D",
70+
"xicProperLength",
71+
"xicDistanceFromPV",
72+
"xiccProperLength"};
73+
74+
} // namespace multi_charm_ml
75+
76+
namespace o2::analysis
77+
{
78+
// list of input features that can be requested via the mlConfigurations.namesInputFeatures configurable
79+
enum class InputFeaturesMulticharm : uint8_t {
80+
xicDauDCA = 0,
81+
xiccDauDCA,
82+
xiDCAxy,
83+
xicDCAxy,
84+
xiccDCAxy,
85+
xiDCAz,
86+
xicDCAz,
87+
xiccDCAz,
88+
pi1cDCAxy,
89+
pi2cDCAxy,
90+
piccDCAxy,
91+
pi1cDCAz,
92+
pi2cDCAz,
93+
piccDCAz,
94+
xicDecayRadius2D,
95+
xiccDecayRadius2D,
96+
xicProperLength,
97+
xicDistanceFromPV,
98+
xiccProperLength
99+
};
100+
101+
template <typename TypeOutputScore = float>
102+
class MulticharmMlResponse : public MlResponse<TypeOutputScore>
103+
{
104+
public:
105+
MulticharmMlResponse() = default;
106+
~MulticharmMlResponse() override = default;
107+
108+
template <typename TMulticharm, typename TCollision>
109+
std::vector<float> getInputFeatures(TMulticharm const& multicharm)
110+
{
111+
std::vector<float> inputFeatures;
112+
inputFeatures.reserve(MlResponse<TypeOutputScore>::mCachedIndices.size());
113+
114+
for (const auto& idx : MlResponse<TypeOutputScore>::mCachedIndices) {
115+
switch (idx) {
116+
case InputFeaturesMulticharm::xicDauDCA:
117+
inputFeatures.emplace_back(multicharm.xicDauDCA());
118+
break;
119+
case InputFeaturesMulticharm::xiccDauDCA:
120+
inputFeatures.emplace_back(multicharm.xiccDauDCA());
121+
break;
122+
case InputFeaturesMulticharm::xiDCAxy:
123+
inputFeatures.emplace_back(multicharm.xiDCAxy());
124+
break;
125+
case InputFeaturesMulticharm::xicDCAxy:
126+
inputFeatures.emplace_back(multicharm.xicDCAxy());
127+
break;
128+
case InputFeaturesMulticharm::xiccDCAxy:
129+
inputFeatures.emplace_back(multicharm.xiccDCAxy());
130+
break;
131+
case InputFeaturesMulticharm::xiDCAz:
132+
inputFeatures.emplace_back(multicharm.xiDCAz());
133+
break;
134+
case InputFeaturesMulticharm::xicDCAz:
135+
inputFeatures.emplace_back(multicharm.xicDCAz());
136+
break;
137+
case InputFeaturesMulticharm::xiccDCAz:
138+
inputFeatures.emplace_back(multicharm.xiccDCAz());
139+
break;
140+
case InputFeaturesMulticharm::pi1cDCAxy:
141+
inputFeatures.emplace_back(multicharm.pi1cDCAxy());
142+
break;
143+
case InputFeaturesMulticharm::pi2cDCAxy:
144+
inputFeatures.emplace_back(multicharm.pi2cDCAxy());
145+
break;
146+
case InputFeaturesMulticharm::piccDCAxy:
147+
inputFeatures.emplace_back(multicharm.piccDCAxy());
148+
break;
149+
case InputFeaturesMulticharm::pi1cDCAz:
150+
inputFeatures.emplace_back(multicharm.pi1cDCAz());
151+
break;
152+
case InputFeaturesMulticharm::pi2cDCAz:
153+
inputFeatures.emplace_back(multicharm.pi2cDCAz());
154+
break;
155+
case InputFeaturesMulticharm::piccDCAz:
156+
inputFeatures.emplace_back(multicharm.piccDCAz());
157+
break;
158+
case InputFeaturesMulticharm::xicDecayRadius2D:
159+
inputFeatures.emplace_back(multicharm.xicDecayRadius2D());
160+
break;
161+
case InputFeaturesMulticharm::xiccDecayRadius2D:
162+
inputFeatures.emplace_back(multicharm.xiccDecayRadius2D());
163+
break;
164+
case InputFeaturesMulticharm::xicProperLength:
165+
inputFeatures.emplace_back(multicharm.xicProperLength());
166+
break;
167+
case InputFeaturesMulticharm::xicDistanceFromPV:
168+
inputFeatures.emplace_back(multicharm.xicDistanceFromPV());
169+
break;
170+
case InputFeaturesMulticharm::xiccProperLength:
171+
inputFeatures.emplace_back(multicharm.xiccProperLength());
172+
break;
173+
}
174+
}
175+
return inputFeatures;
176+
}
177+
178+
protected:
179+
void setAvailableInputFeatures() override
180+
{
181+
MlResponse<TypeOutputScore>::mAvailableInputFeatures = {
182+
{"xicDauDCA", static_cast<uint8_t>(InputFeaturesMulticharm::xicDauDCA)},
183+
{"xiccDauDCA", static_cast<uint8_t>(InputFeaturesMulticharm::xiccDauDCA)},
184+
{"xiDCAxy", static_cast<uint8_t>(InputFeaturesMulticharm::xiDCAxy)},
185+
{"xicDCAxy", static_cast<uint8_t>(InputFeaturesMulticharm::xicDCAxy)},
186+
{"xiccDCAxy", static_cast<uint8_t>(InputFeaturesMulticharm::xiccDCAxy)},
187+
{"xiDCAz", static_cast<uint8_t>(InputFeaturesMulticharm::xiDCAz)},
188+
{"xicDCAz", static_cast<uint8_t>(InputFeaturesMulticharm::xicDCAz)},
189+
{"xiccDCAz", static_cast<uint8_t>(InputFeaturesMulticharm::xiccDCAz)},
190+
{"pi1cDCAxy", static_cast<uint8_t>(InputFeaturesMulticharm::pi1cDCAxy)},
191+
{"pi2cDCAxy", static_cast<uint8_t>(InputFeaturesMulticharm::pi2cDCAxy)},
192+
{"piccDCAxy", static_cast<uint8_t>(InputFeaturesMulticharm::piccDCAxy)},
193+
{"pi1cDCAz", static_cast<uint8_t>(InputFeaturesMulticharm::pi1cDCAz)},
194+
{"pi2cDCAz", static_cast<uint8_t>(InputFeaturesMulticharm::pi2cDCAz)},
195+
{"piccDCAz", static_cast<uint8_t>(InputFeaturesMulticharm::piccDCAz)},
196+
{"xicDecayRadius2D", static_cast<uint8_t>(InputFeaturesMulticharm::xicDecayRadius2D)},
197+
{"xiccDecayRadius2D", static_cast<uint8_t>(InputFeaturesMulticharm::xiccDecayRadius2D)},
198+
{"xicProperLength", static_cast<uint8_t>(InputFeaturesMulticharm::xicProperLength)},
199+
{"xicDistanceFromPV", static_cast<uint8_t>(InputFeaturesMulticharm::xicDistanceFromPV)},
200+
{"xiccProperLength", static_cast<uint8_t>(InputFeaturesMulticharm::xiccProperLength)}};
201+
}
202+
};
203+
204+
} // namespace o2::analysis
205+
206+
#endif // ALICE3_ML_MULTICHARMRESPONSE_H_

0 commit comments

Comments
 (0)