16namespace Experimental {
53 ss <<
"Failed to open TMVA XML file "
55 throw std::runtime_error(
ss.str());
58 for (
auto node =
xml.GetChild(
mainNode); node; node =
xml.GetNext(node)) {
59 const auto nodeName = std::string(
xml.GetNodeName(node));
61 if (
nodeName.compare(
"Variables") == 0) {
62 c.numVariables = std::atoi(
xml.GetAttr(node,
"NVar"));
63 c.variables = std::vector<std::string>(
c.numVariables);
64 c.variable_expressions = std::vector<std::string>(
c.numVariables);
72 else if (
nodeName.compare(
"Spectators") == 0) {
73 c.numSpectators = std::atoi(
xml.GetAttr(node,
"NSpec"));
74 c.spectators = std::vector<std::string>(
c.numSpectators);
75 c.spectator_expressions = std::vector<std::string>(
c.numSpectators);
83 else if (
nodeName.compare(
"Classes") == 0) {
84 c.numClasses = std::atoi(
xml.GetAttr(node,
"NClass"));
90 else if (
nodeName.compare(
"GeneralInfo") == 0) {
91 std::string analysisType =
"";
93 if (std::string(
"AnalysisType").compare(
xml.GetAttr(
thisNode,
"name")) == 0) {
97 if (analysisType.compare(
"Classification") == 0) {
99 }
else if (analysisType.compare(
"Regression") == 0) {
101 }
else if (analysisType.compare(
"Multiclass") == 0) {
109 if (
c.numVariables !=
c.variables.size() ||
c.numVariables == 0) {
110 std::stringstream
ss;
111 ss <<
"Failed to parse input variables from TMVA config " <<
filename <<
".";
112 throw std::runtime_error(
ss.str());
114 if (
c.numSpectators !=
c.spectators.size()) {
115 std::stringstream
ss;
116 ss <<
"Failed to parse input spectators from TMVA config " <<
filename <<
".";
117 throw std::runtime_error(
ss.str());
119 if (
c.numClasses !=
c.classes.size() ||
c.numClasses == 0) {
120 std::stringstream
ss;
121 ss <<
"Failed to parse output classes from TMVA config " <<
filename <<
".";
122 throw std::runtime_error(
ss.str());
125 std::stringstream
ss;
126 ss <<
"Failed to parse analysis type from TMVA config " <<
filename <<
".";
127 throw std::runtime_error(
ss.str());
165 fReader = std::make_unique<Reader>(
"Silent");
168 for (std::size_t i = 0; i <
numVars; i++) {
173 for (std::size_t i = 0; i <
numSpecs; i++) {
180 std::vector<float>
Compute(
const std::vector<float> &
x)
186 throw std::runtime_error(
"Size of input vector is not equal to number of variables.");
190 for (std::size_t i = 0; i !=
nVars ; ++i) {
193 for (std::size_t i = 0; i !=
fSpectators.size(); ++i) {
200 return std::vector<float>({
static_cast<float>(
fReader->EvaluateMVA(
name))});
212 throw std::runtime_error(
"RReader has undefined analysis type.");
213 return std::vector<float>();
222 std::vector<float>
Compute(std::span<const float>
x)
226 throw std::runtime_error(
227 "Size of input vector is not a multiple of the number of variables, which must be nonzero, or the input "
230 const std::size_t numEntries =
x.size() /
numCols;
233 unsigned int numClasses = 1;
236 std::vector<float>
y(numEntries * numClasses);
241 for (std::size_t i = 0; i < numEntries; i++) {
242 for (std::size_t
j = 0;
j <
nVars;
j++) {
259 for (std::size_t k = 0; k < numClasses; k++)
260 y[i * numClasses + k] =
p[k];
ROOT::Detail::TRangeCast< T, true > TRangeDynCast
TRangeDynCast is an adapter class that allows the typed iteration through a TCollection.
winID h TVirtualViewer3D TVirtualGLPainter p
Option_t Option_t TPoint TPoint const char GetTextMagnitude GetFillStyle GetLineColor GetLineWidth GetMarkerStyle GetTextAlign GetTextColor GetTextSize void char Point_t Rectangle_t WindowAttributes_t Float_t Float_t Float_t Int_t Int_t UInt_t UInt_t Rectangle_t Int_t Int_t Window_t TString Int_t GCValues_t GetPrimarySelectionOwner GetDisplay GetScreen GetColormap GetNativeEvent const char const char dpyName wid window const char font_name cursor keysym reg const char only_if_exist regb h Point_t winding char text const char depth char const char Int_t count const char ColorStruct_t color const char filename
#define R__WRITE_LOCKGUARD(mutex)
A replacement for the TMVA::Reader legacy interface.
std::vector< float > Compute(std::span< const float > x)
Compute model prediction on a flat batch of events The input is the concatenation of the events' inpu...
std::vector< float > Compute(const std::vector< float > &x)
Compute model prediction on vector.
std::vector< std::string > GetSpectatorNames()
Internal::AnalysisType fAnalysisType
std::vector< float > fVariableValues
std::vector< float > fSpectatorValues
std::vector< std::string > fVariableExpressions
std::vector< std::string > GetVariableNames()
std::vector< std::string > fSpectatorExpressions
std::vector< std::string > fSpectators
std::vector< std::string > fVariables
RReader(const std::string &path)
Create TMVA model from XML file.
std::unique_ptr< Reader > fReader
R__EXTERN TVirtualRWMutex * gCoreMutex
XMLConfig ParseXMLConfig(const std::string &filename)
Parse TMVA XML config.
AnalysisType
Internal definition of analysis types.
create variable transformations
Container for information extracted from TMVA XML config.
std::vector< std::string > classes
unsigned int numVariables
AnalysisType analysisType
std::vector< std::string > spectators
std::vector< std::string > spectator_expressions
unsigned int numSpectators
std::vector< std::string > variable_expressions
std::vector< std::string > variables