Logo ROOT  
Reference Guide
 
Loading...
Searching...
No Matches
RReader.hxx
Go to the documentation of this file.
1#ifndef TMVA_RREADER
2#define TMVA_RREADER
3
4#include "TString.h"
5#include "TXMLEngine.h"
6
7#include "TMVA/Reader.h"
8
9#include <ROOT/RSpan.hxx>
10
11#include <memory> // std::unique_ptr
12#include <sstream> // std::stringstream
13#include <vector>
14
15namespace TMVA {
16namespace Experimental {
17
18namespace Internal {
19
20/// Internal definition of analysis types
22
23/// Container for information extracted from TMVA XML config
24struct XMLConfig {
25 unsigned int numVariables;
26 std::vector<std::string> variables;
27 std::vector<std::string> variable_expressions;
28 unsigned int numSpectators;
29 std::vector<std::string> spectators;
30 std::vector<std::string> spectator_expressions;
31 unsigned int numClasses;
32 std::vector<std::string> classes;
35 : numVariables(0), variables(std::vector<std::string>(0)),
36 numSpectators(0), spectators(std::vector<std::string>(0)),
37 numClasses(0), classes(std::vector<std::string>(0)),
39 {
40 }
41};
42
43/// Parse TMVA XML config
44inline XMLConfig ParseXMLConfig(const std::string &filename)
45{
47
48 // Parse XML file and find root node
50 auto xmldoc = xml.ParseFile(filename.c_str());
51 if (!xmldoc) {
52 std::stringstream ss;
53 ss << "Failed to open TMVA XML file "
54 << filename << ".";
55 throw std::runtime_error(ss.str());
56 }
57 auto mainNode = xml.DocGetRootElement(xmldoc);
58 for (auto node = xml.GetChild(mainNode); node; node = xml.GetNext(node)) {
59 const auto nodeName = std::string(xml.GetNodeName(node));
60 // Read out input variables
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);
65 for (auto thisNode = xml.GetChild(node); thisNode; thisNode = xml.GetNext(thisNode)) {
66 const auto iVariable = std::atoi(xml.GetAttr(thisNode, "VarIndex"));
67 c.variables[iVariable] = xml.GetAttr(thisNode, "Title");
68 c.variable_expressions[iVariable] = xml.GetAttr(thisNode, "Expression");
69 }
70 }
71 // Read out input spectators
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);
76 for (auto thisNode = xml.GetChild(node); thisNode; thisNode = xml.GetNext(thisNode)) {
77 const auto iVariable = std::atoi(xml.GetAttr(thisNode, "SpecIndex"));
78 c.spectators[iVariable] = xml.GetAttr(thisNode, "Title");
79 c.spectator_expressions[iVariable] = xml.GetAttr(thisNode, "Expression");
80 }
81 }
82 // Read out output classes
83 else if (nodeName.compare("Classes") == 0) {
84 c.numClasses = std::atoi(xml.GetAttr(node, "NClass"));
85 for (auto thisNode = xml.GetChild(node); thisNode; thisNode = xml.GetNext(thisNode)) {
86 c.classes.push_back(xml.GetAttr(thisNode, "Name"));
87 }
88 }
89 // Read out analysis type
90 else if (nodeName.compare("GeneralInfo") == 0) {
91 std::string analysisType = "";
92 for (auto thisNode = xml.GetChild(node); thisNode; thisNode = xml.GetNext(thisNode)) {
93 if (std::string("AnalysisType").compare(xml.GetAttr(thisNode, "name")) == 0) {
94 analysisType = xml.GetAttr(thisNode, "value");
95 }
96 }
97 if (analysisType.compare("Classification") == 0) {
99 } else if (analysisType.compare("Regression") == 0) {
101 } else if (analysisType.compare("Multiclass") == 0) {
103 }
104 }
105 }
106 xml.FreeDoc(xmldoc);
107
108 // Error-handling
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());
113 }
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());
118 }
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());
123 }
124 if (c.analysisType == Internal::AnalysisType::Undefined) {
125 std::stringstream ss;
126 ss << "Failed to parse analysis type from TMVA config " << filename << ".";
127 throw std::runtime_error(ss.str());
128 }
129
130 return c;
131}
132
133} // namespace Internal
134
135/// A replacement for the TMVA::Reader legacy interface.
136/// Performs inference for TMVA models stored as XML files.
137/// For neural network inference consider using [SOFIE](https://github.com/root-project/root/blob/master/tmva/sofie/README.md) instead.
138class RReader {
139private:
140 std::unique_ptr<Reader> fReader;
141 std::vector<float> fVariableValues;
142 std::vector<std::string> fVariables;
143 std::vector<std::string> fVariableExpressions;
144 std::vector<float> fSpectatorValues;
145 std::vector<std::string> fSpectators;
146 std::vector<std::string> fSpectatorExpressions;
147 unsigned int fNumClasses;
148 const char *name = "RReader";
150
151public:
152 /// Create TMVA model from XML file
153 RReader(const std::string &path)
154 {
155 // Load config
156 auto c = Internal::ParseXMLConfig(path);
157 fVariables = c.variables;
158 fVariableExpressions = c.variable_expressions;
159 fSpectators = c.spectators;
160 fSpectatorExpressions = c.spectator_expressions;
161 fAnalysisType = c.analysisType;
162 fNumClasses = c.numClasses;
163
164 // Setup reader
165 fReader = std::make_unique<Reader>("Silent");
166 const auto numVars = fVariables.size();
167 fVariableValues = std::vector<float>(numVars);
168 for (std::size_t i = 0; i < numVars; i++) {
170 }
171 const auto numSpecs = fSpectators.size();
172 fSpectatorValues = std::vector<float>(numSpecs);
173 for (std::size_t i = 0; i < numSpecs; i++) {
175 }
176 fReader->BookMVA(name, path.c_str());
177 }
178
179 /// Compute model prediction on vector
180 std::vector<float> Compute(const std::vector<float> &x)
181 {
182 // Take lock to protect model evaluation
184
185 if (x.size() != (fVariables.size()+fSpectators.size()))
186 throw std::runtime_error("Size of input vector is not equal to number of variables.");
187
188 // Copy over inputs to memory used by TMVA reader
189 const auto nVars = fVariables.size();
190 for (std::size_t i = 0; i != nVars ; ++i) {
191 fVariableValues[i] = x[i];
192 }
193 for (std::size_t i = 0; i != fSpectators.size(); ++i) {
194 fSpectatorValues[i] = x[nVars+i];
195 }
196
197 // Evaluate TMVA model
198 // Classification
200 return std::vector<float>({static_cast<float>(fReader->EvaluateMVA(name))});
201 }
202 // Regression
204 return fReader->EvaluateRegression(name);
205 }
206 // Multiclass
208 return fReader->EvaluateMulticlass(name);
209 }
210 // Throw error
211 else {
212 throw std::runtime_error("RReader has undefined analysis type.");
213 return std::vector<float>();
214 }
215 }
216
217 /// Compute model prediction on a flat batch of events
218 /// The input is the concatenation of the events' input variables and spectators
219 /// in row-major layout and the returned vector is flat row-major as well, with
220 /// size nEvents * numClasses (numClasses = 1 for classification and regression)
221 /// and the outputs of one event contiguous at y[event * numClasses ...].
222 std::vector<float> Compute(std::span<const float> x)
223 {
224 const std::size_t numCols = fVariables.size() + fSpectators.size();
225 if (numCols == 0 || x.empty() || x.size() % numCols != 0)
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 "
228 "is empty.");
229
230 const std::size_t numEntries = x.size() / numCols;
231
232 // Define size of output vector based on analysis type
233 unsigned int numClasses = 1;
235 numClasses = fNumClasses;
236 std::vector<float> y(numEntries * numClasses);
237
238 // Fill output vector
239 const auto nVars = fVariables.size(); // number of non-spectator variables
241 for (std::size_t i = 0; i < numEntries; i++) {
242 for (std::size_t j = 0; j < nVars; j++) {
243 fVariableValues[j] = x[i * numCols + j];
244 }
245 for (std::size_t j = 0; j < fSpectators.size(); ++j) {
246 fSpectatorValues[j] = x[i * numCols + nVars + j];
247 }
248 // Classification
250 y[i] = fReader->EvaluateMVA(name);
251 }
252 // Regression
254 y[i] = fReader->EvaluateRegression(name)[0];
255 }
256 // Multiclass
258 const auto p = fReader->EvaluateMulticlass(name);
259 for (std::size_t k = 0; k < numClasses; k++)
260 y[i * numClasses + k] = p[k];
261 }
262 }
263
264 return y;
265 }
266
267 std::vector<std::string> GetVariableNames() { return fVariables; }
268 std::vector<std::string> GetSpectatorNames() { return fSpectators; }
269};
270
271} // namespace Experimental
272} // namespace TMVA
273
274#endif // TMVA_RREADER
#define c(i)
Definition RSha256.hxx:101
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.
Definition RReader.hxx:138
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...
Definition RReader.hxx:222
std::vector< float > Compute(const std::vector< float > &x)
Compute model prediction on vector.
Definition RReader.hxx:180
std::vector< std::string > GetSpectatorNames()
Definition RReader.hxx:268
Internal::AnalysisType fAnalysisType
Definition RReader.hxx:149
std::vector< float > fVariableValues
Definition RReader.hxx:141
std::vector< float > fSpectatorValues
Definition RReader.hxx:144
std::vector< std::string > fVariableExpressions
Definition RReader.hxx:143
std::vector< std::string > GetVariableNames()
Definition RReader.hxx:267
std::vector< std::string > fSpectatorExpressions
Definition RReader.hxx:146
std::vector< std::string > fSpectators
Definition RReader.hxx:145
std::vector< std::string > fVariables
Definition RReader.hxx:142
RReader(const std::string &path)
Create TMVA model from XML file.
Definition RReader.hxx:153
std::unique_ptr< Reader > fReader
Definition RReader.hxx:140
Basic string class.
Definition TString.h:137
Double_t y[n]
Definition legend1.C:17
Double_t x[n]
Definition legend1.C:17
R__EXTERN TVirtualRWMutex * gCoreMutex
XMLConfig ParseXMLConfig(const std::string &filename)
Parse TMVA XML config.
Definition RReader.hxx:44
AnalysisType
Internal definition of analysis types.
Definition RReader.hxx:21
create variable transformations
Container for information extracted from TMVA XML config.
Definition RReader.hxx:24
std::vector< std::string > classes
Definition RReader.hxx:32
std::vector< std::string > spectators
Definition RReader.hxx:29
std::vector< std::string > spectator_expressions
Definition RReader.hxx:30
std::vector< std::string > variable_expressions
Definition RReader.hxx:27
std::vector< std::string > variables
Definition RReader.hxx:26