Logo ROOT  
Reference Guide
 
Loading...
Searching...
No Matches
RBDT.cxx
Go to the documentation of this file.
1/**********************************************************************************
2 * Project: ROOT - a Root-integrated toolkit for multivariate data analysis *
3 * Package: TMVA *
4 * *
5 * *
6 * Description: *
7 * *
8 * Authors: *
9 * Jonas Rembser (jonas.rembser@cern.ch) *
10 * *
11 * Copyright (c) 2024: *
12 * CERN, Switzerland *
13 * *
14 * Redistribution and use in source and binary forms, with or without *
15 * modification, are permitted according to the terms listed in LICENSE *
16 * (see tmva/doc/LICENSE) *
17 **********************************************************************************/
18
19#include <TMVA/RBDT.hxx>
20
21#include <ROOT/StringUtils.hxx>
22
23#include <TSystem.h>
24
25#include <nlohmann/json.hpp>
26
27#include <cmath>
28#include <fstream>
29#include <iostream>
30#include <sstream>
31#include <stdexcept>
32#include <cstdlib>
33
34namespace {
35
36template <class Value_t>
37void softmaxTransformInplace(Value_t *out, int nOut)
38{
39 // Do softmax transformation inplace, mimicing exactly the Softmax function
40 // in the src/common/math.h source file of xgboost.
41 double norm = 0.;
42 Value_t wmax = *out;
43 for (int i = 1; i < nOut; ++i) {
44 wmax = std::max(out[i], wmax);
45 }
46 for (int i = 0; i < nOut; ++i) {
47 Value_t &x = out[i];
48 x = std::exp(x - wmax);
49 norm += x;
50 }
51 for (int i = 0; i < nOut; ++i) {
52 out[i] /= static_cast<float>(norm);
53 }
54}
55
56namespace util {
57
58template <class NumericType>
59struct NumericAfterSubstrOutput {
60 explicit NumericAfterSubstrOutput()
61 {
62 value = 0;
63 found = false;
64 failed = true;
65 }
67 bool found;
68 bool failed;
69 std::string rest;
70};
71
72template <class NumericType>
73inline NumericAfterSubstrOutput<NumericType> numericAfterSubstr(std::string const &str, std::string const &substr)
74{
75 std::string rest;
77 output.rest = str;
78
79 std::size_t found = str.find(substr);
80 if (found != std::string::npos) {
81 output.found = true;
82 std::stringstream ss(str.substr(found + substr.size(), str.size() - found + substr.size()));
83 ss >> output.value;
84 if (!ss.fail()) {
85 output.failed = false;
86 output.rest = ss.str();
87 }
88 }
89 return output;
90}
91
92} // namespace util
93
94} // namespace
95
96/// Compute model prediction on a flat batch of events
97std::vector<TMVA::Experimental::RBDT::Value_t>
98TMVA::Experimental::RBDT::Compute(std::span<const Value_t> x, unsigned int cols) const
99{
100 if (cols == 0 || x.empty() || x.size() % cols != 0) {
101 throw std::runtime_error(
102 "TMVA::Experimental::RBDT: the number of columns is zero or the batch input is empty or its size is not a "
103 "multiple of the number of columns.");
104 }
105 const std::size_t nOut = fBaseResponses.size() > 2 ? fBaseResponses.size() : 1;
106 const std::size_t rows = x.size() / cols;
107 std::vector<Value_t> y(rows * nOut);
108 for (std::size_t iRow = 0; iRow < rows; ++iRow) {
109 ComputeImpl(x.data() + iRow * cols, y.data() + iRow * nOut);
110 }
111 return y;
112}
113
115{
116 std::size_t nOut = fBaseResponses.size() > 2 ? fBaseResponses.size() : 1;
117 if (nOut == 1) {
118 throw std::runtime_error(
119 "Error in RBDT::softmax : binary classification models don't support softmax evaluation. Plase set "
120 "the number of classes in the RBDT-creating function if this is a multiclassification model.");
121 }
122
123 for (std::size_t i = 0; i < nOut; ++i) {
124 out[i] = fBaseScore + fBaseResponses[i];
125 }
126
127 int iRootIndex = 0;
128 for (int index : fRootIndices) {
129 do {
130 int r = fRightIndices[index];
131 int l = fLeftIndices[index];
132 index = array[fCutIndices[index]] < fCutValues[index] ? l : r;
133 } while (index > 0);
134 out[fTreeNumbers[iRootIndex] % nOut] += fResponses[-index];
135 ++iRootIndex;
136 }
137
139}
140
142{
143 std::size_t nOut = fBaseResponses.size() > 2 ? fBaseResponses.size() : 1;
144 if (nOut > 1) {
145 Softmax(array, out);
146 } else {
147 out[0] = EvaluateBinary(array);
148 if (fLogistic) {
149 out[0] = 1.0 / (1.0 + std::exp(-out[0]));
150 }
151 }
152}
153
155{
156 Value_t out = fBaseScore + fBaseResponses[0];
157
158 for (std::vector<int>::const_iterator indexIter = fRootIndices.begin(); indexIter != fRootIndices.end();
159 ++indexIter) {
160 int index = *indexIter;
161 do {
162 int r = fRightIndices[index];
163 int l = fLeftIndices[index];
164 index = array[fCutIndices[index]] < fCutValues[index] ? l : r;
165 } while (index > 0);
166 out += fResponses[-index];
167 }
168
169 return out;
170}
171
172/// RBDT uses a more efficient representation of the BDT in flat arrays. This
173/// function translates the indices to the RBDT indices. In RBDT, leaf nodes
174/// are stored in separate arrays. To encode this, the sign of the index is
175/// flipped.
177 IndexMap const &leafIndices)
178{
179 for (int &idx : indices) {
180 auto foundNode = nodeIndices.find(idx);
181 if (foundNode != nodeIndices.end()) {
182 idx = foundNode->second;
183 continue;
184 }
185 auto foundLeaf = leafIndices.find(idx);
186 if (foundLeaf != leafIndices.end()) {
187 idx = -foundLeaf->second;
188 continue;
189 } else {
190 std::stringstream errMsg;
191 errMsg << "RBDT: something is wrong in the node structure - node with index " << idx << " doesn't exist";
192 throw std::runtime_error(errMsg.str());
193 }
194 }
195}
196
199{
200 correctIndices({ff.fRightIndices.begin() + nPreviousNodes, ff.fRightIndices.end()}, nodeIndices, leafIndices);
201 correctIndices({ff.fLeftIndices.begin() + nPreviousNodes, ff.fLeftIndices.end()}, nodeIndices, leafIndices);
202
203 if (nPreviousNodes != static_cast<int>(ff.fCutValues.size())) {
204 ff.fTreeNumbers.push_back(ff.fRootIndices.size() + treesSkipped);
205 ff.fRootIndices.push_back(nPreviousNodes);
206 } else {
207 int treeNumbers = ff.fRootIndices.size() + treesSkipped;
208 ++treesSkipped;
209 ff.fBaseResponses[treeNumbers % ff.fBaseResponses.size()] += ff.fResponses.back();
210 ff.fResponses.pop_back();
211 }
212
213 nodeIndices.clear();
214 leafIndices.clear();
215 nPreviousNodes = ff.fCutValues.size();
216 nPreviousLeaves = ff.fResponses.size();
217}
218
219/// Construct an RBDT from an XGBoost model in its native JSON serialization.
220///
221/// This reads the structured model that XGBoost writes with Booster.save_model().
222/// That format stores each tree as a set of parallel arrays and references
223/// features by index, so no feature-name resolution is needed. Everything else
224/// (objective, base score, number of classes) is taken from the file, which
225/// makes this a self-contained, Python-free entry point.
227{
228 const std::string info = "constructing RBDT from '" + jsonPath + "': ";
229
230 if (gSystem->AccessPathName(jsonPath.c_str())) {
231 throw std::runtime_error(info + "file does not exist");
232 }
233
234 nlohmann::json j;
235 {
236 std::ifstream jsonFile(jsonPath.c_str());
237 jsonFile >> j;
238 }
239
240 auto const &learner = j.at("learner");
241 auto const &modelParam = learner.at("learner_model_param");
242
243 // Map the XGBoost objective to the RBDT one.
244 std::string const xgbObjective = learner.at("objective").at("name").get<std::string>();
245 static const std::unordered_map<std::string, std::string> objectiveMap{
246 {"multi:softprob", "softmax"}, // Naming the objective softmax is more common today
247 {"binary:logistic", "logistic"},
248 {"reg:linear", "identity"},
249 {"reg:squarederror", "identity"},
250 };
253 std::string supported;
254 for (auto const &item : objectiveMap) {
255 supported += (supported.empty() ? "" : ", ") + item.first;
256 }
257 throw std::runtime_error(info + "XGBoost model has unsupported objective \"" + xgbObjective +
258 "\". Supported objectives are " + supported + ".");
259 }
260 bool const logistic = foundObjective->second == "logistic";
261
262 // The base score is stored as a string, e.g. "5.14E-1". Since XGBoost 3.1.0 it
263 // is always serialized as a JSON array embedded in that string (e.g.
264 // "[5.14E-1]"), even for single-output models. Only a genuine multi-element
265 // array (multi-target base score) is unsupported.
266 std::string const baseScoreStr = modelParam.at("base_score").get<std::string>();
267 double baseScoreProb;
268 if (baseScoreStr.find('[') != std::string::npos) {
269 nlohmann::json const baseScoreArr = nlohmann::json::parse(baseScoreStr);
270 if (baseScoreArr.size() > 1) {
271 throw std::runtime_error(info + "model contains multiple base scores, which is not supported. This "
272 "typically occurs with XGBoost >= 3.1.0, which supports multi-target base "
273 "scores.");
274 }
275 baseScoreProb = baseScoreArr.at(0).get<double>();
276 } else {
277 baseScoreProb = std::stod(baseScoreStr);
278 }
279 // For a logistic objective the base score is a probability, but RBDT works on
280 // the raw margin, so we apply the logit transform (as the Python code does).
281 Value_t const baseScore = logistic ? std::log(baseScoreProb / (1.0 - baseScoreProb)) : baseScoreProb;
282
283 // Only multiclass models produce more than one output.
284 int nClasses = 1;
285 if (xgbObjective.rfind("multi:", 0) == 0) {
286 nClasses = std::stoi(modelParam.at("num_class").get<std::string>());
287 }
288
289 RBDT ff;
290 ff.fLogistic = logistic;
291 ff.fBaseScore = baseScore;
292 ff.fBaseResponses.resize(nClasses <= 2 ? 1 : nClasses);
293
294 auto const &trees = learner.at("gradient_booster").at("model").at("trees");
295
296 int treesSkipped = 0;
297 int nPreviousNodes = 0;
298 int nPreviousLeaves = 0;
301
302 // Fill the flat RBDT arrays tree by tree, keying the index maps by the node's
303 // position in the XGBoost arrays. terminateTree() then remaps the child
304 // references to the RBDT indexing (negated for leaves), exactly as for the
305 // text dump. Node 0 is always the tree root, so iterating in array order
306 // makes it the first internal node of the tree, which is what fRootIndices
307 // expects.
308 for (auto const &tree : trees) {
309 auto const &leftChildren = tree.at("left_children");
310 auto const &rightChildren = tree.at("right_children");
311 auto const &splitIndices = tree.at("split_indices");
312 auto const &splitConditions = tree.at("split_conditions");
313
314 std::size_t const nNodes = leftChildren.size();
315 for (std::size_t i = 0; i < nNodes; ++i) {
316 int const left = leftChildren[i].get<int>();
317 if (left == -1) {
318 // Leaf node: the split condition holds the leaf response.
319 ff.fResponses.push_back(splitConditions[i].get<Value_t>());
320 std::size_t const nLeafIndices = leafIndices.size();
322 } else {
323 // Internal node: x < cut goes left (yes), otherwise right (no).
324 ff.fCutValues.push_back(splitConditions[i].get<Value_t>());
325 ff.fCutIndices.push_back(splitIndices[i].get<unsigned int>());
326 ff.fLeftIndices.push_back(left);
327 ff.fRightIndices.push_back(rightChildren[i].get<int>());
328 std::size_t const nNodeIndices = nodeIndices.size();
330 }
331 }
332
334 }
335
336 if (nClasses > 2 && (ff.fRootIndices.size() + treesSkipped) % nClasses != 0) {
337 std::stringstream ss;
338 ss << info << "Forest has " << ff.fRootIndices.size() << " trees, which is not compatible with " << nClasses
339 << " classes!";
340 throw std::runtime_error(ss.str());
341 }
342
343 return ff;
344}
ROOT::Detail::TRangeCast< T, true > TRangeDynCast
TRangeDynCast is an adapter class that allows the typed iteration through a TCollection.
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 r
Option_t Option_t TPoint TPoint const char GetTextMagnitude GetFillStyle GetLineColor GetLineWidth GetMarkerStyle GetTextAlign GetTextColor GetTextSize void char Point_t Rectangle_t WindowAttributes_t index
Option_t Option_t TPoint TPoint const char GetTextMagnitude GetFillStyle GetLineColor GetLineWidth GetMarkerStyle GetTextAlign GetTextColor GetTextSize void value
Option_t Option_t TPoint TPoint const char GetTextMagnitude GetFillStyle GetLineColor GetLineWidth GetMarkerStyle GetTextAlign GetTextColor GetTextSize void char Point_t Rectangle_t wmax
R__EXTERN TSystem * gSystem
Definition TSystem.h:582
const_iterator begin() const
const_iterator end() const
static void terminateTree(TMVA::Experimental::RBDT &ff, int &nPreviousNodes, int &nPreviousLeaves, IndexMap &nodeIndices, IndexMap &leafIndices, int &treesSkipped)
Definition RBDT.cxx:197
static RBDT LoadXGBoost(std::string const &jsonPath)
Construct an RBDT from an XGBoost model in its native JSON serialization.
Definition RBDT.cxx:226
static void correctIndices(std::span< int > indices, IndexMap const &nodeIndices, IndexMap const &leafIndices)
RBDT uses a more efficient representation of the BDT in flat arrays.
Definition RBDT.cxx:176
std::unordered_map< int, int > IndexMap
Map from XGBoost to RBDT indices.
Definition RBDT.hxx:71
void Softmax(const Value_t *array, Value_t *out) const
Definition RBDT.cxx:114
Value_t EvaluateBinary(const Value_t *array) const
Definition RBDT.cxx:154
std::vector< Value_t > fBaseResponses
Definition RBDT.hxx:87
Vector Compute(const Vector &x) const
Compute model prediction on a single event.
Definition RBDT.hxx:44
void ComputeImpl(const Value_t *array, Value_t *out) const
Definition RBDT.cxx:141
virtual Bool_t AccessPathName(const char *path, EAccessMode mode=kFileExists)
Returns FALSE if one can access a file using the specified access mode.
Definition TSystem.cxx:1312
Double_t y[n]
Definition legend1.C:17
Double_t x[n]
Definition legend1.C:17
Definition RBDT.cxx:56
TLine l
Definition textangle.C:4