58template <
class NumericType>
59struct NumericAfterSubstrOutput {
60 explicit NumericAfterSubstrOutput()
72template <
class NumericType>
73inline NumericAfterSubstrOutput<NumericType> numericAfterSubstr(std::string
const &str, std::string
const &substr)
76 NumericAfterSubstrOutput<NumericType> output;
79 std::size_t found = str.find(substr);
80 if (found != std::string::npos) {
82 std::stringstream ss(str.substr(found + substr.size(), str.size() - found + substr.size()));
85 output.failed =
false;
86 output.rest = ss.str();
116 std::size_t nOut = fBaseResponses.size() > 2 ? fBaseResponses.size() : 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.");
123 for (std::size_t i = 0; i < nOut; ++i) {
124 out[i] = fBaseScore + fBaseResponses[i];
128 for (
int index : fRootIndices) {
130 int r = fRightIndices[
index];
131 int l = fLeftIndices[
index];
134 out[fTreeNumbers[iRootIndex] % nOut] += fResponses[-
index];
138 softmaxTransformInplace(out, nOut);
179 for (
int &idx : indices) {
180 auto foundNode = nodeIndices.find(idx);
181 if (foundNode != nodeIndices.end()) {
182 idx = foundNode->second;
185 auto foundLeaf = leafIndices.find(idx);
186 if (foundLeaf != leafIndices.end()) {
187 idx = -foundLeaf->second;
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());
203 if (nPreviousNodes !=
static_cast<int>(ff.
fCutValues.size())) {
207 int treeNumbers = ff.
fRootIndices.size() + treesSkipped;
228 const std::string info =
"constructing RBDT from '" + jsonPath +
"': ";
231 throw std::runtime_error(info +
"file does not exist");
236 std::ifstream jsonFile(jsonPath.c_str());
240 auto const &learner = j.at(
"learner");
241 auto const &modelParam = learner.at(
"learner_model_param");
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"},
247 {
"binary:logistic",
"logistic"},
248 {
"reg:linear",
"identity"},
249 {
"reg:squarederror",
"identity"},
251 auto foundObjective = objectiveMap.find(xgbObjective);
252 if (foundObjective == objectiveMap.end()) {
253 std::string supported;
254 for (
auto const &item : objectiveMap) {
255 supported += (supported.empty() ?
"" :
", ") + item.first;
257 throw std::runtime_error(info +
"XGBoost model has unsupported objective \"" + xgbObjective +
258 "\". Supported objectives are " + supported +
".");
260 bool const logistic = foundObjective->second ==
"logistic";
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 "
275 baseScoreProb = baseScoreArr.at(0).get<
double>();
277 baseScoreProb = std::stod(baseScoreStr);
281 Value_t const baseScore = logistic ? std::log(baseScoreProb / (1.0 - baseScoreProb)) : baseScoreProb;
285 if (xgbObjective.rfind(
"multi:", 0) == 0) {
286 nClasses = std::stoi(modelParam.at(
"num_class").get<std::string>());
294 auto const &trees = learner.at(
"gradient_booster").at(
"model").at(
"trees");
296 int treesSkipped = 0;
297 int nPreviousNodes = 0;
298 int nPreviousLeaves = 0;
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");
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>();
319 ff.
fResponses.push_back(splitConditions[i].get<Value_t>());
320 std::size_t
const nLeafIndices = leafIndices.size();
321 leafIndices[i] = nLeafIndices + nPreviousLeaves;
324 ff.
fCutValues.push_back(splitConditions[i].get<Value_t>());
325 ff.
fCutIndices.push_back(splitIndices[i].get<unsigned int>());
328 std::size_t
const nNodeIndices = nodeIndices.size();
329 nodeIndices[i] = nNodeIndices + nPreviousNodes;
333 terminateTree(ff, nPreviousNodes, nPreviousLeaves, nodeIndices, leafIndices, treesSkipped);
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
340 throw std::runtime_error(ss.str());
static void terminateTree(TMVA::Experimental::RBDT &ff, int &nPreviousNodes, int &nPreviousLeaves, IndexMap &nodeIndices, IndexMap &leafIndices, int &treesSkipped)
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.