Logo ROOT  
Reference Guide
 
Loading...
Searching...
No Matches
RBDT.hxx
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 * Stefan Wunsch (stefan.wunsch@cern.ch) *
10 * Jonas Rembser (jonas.rembser@cern.ch) *
11 * *
12 * Copyright (c) 2024: *
13 * CERN, Switzerland *
14 * *
15 * Redistribution and use in source and binary forms, with or without *
16 * modification, are permitted according to the terms listed in LICENSE *
17 * (see tmva/doc/LICENSE) *
18 **********************************************************************************/
19
20#ifndef TMVA_RBDT
21#define TMVA_RBDT
22
23#include <ROOT/RSpan.hxx>
24
25#include <array>
26#include <istream>
27#include <string>
28#include <unordered_map>
29#include <vector>
30
31namespace TMVA {
32
33namespace Experimental {
34
35class RBDT final {
36public:
37 typedef float Value_t;
38
39 /// Compute model prediction on a single event.
40 ///
41 /// The method is intended to be used with std::vectors-like containers,
42 /// for example RVecs.
43 template <typename Vector>
44 Vector Compute(const Vector &x) const
45 {
46 std::size_t nOut = fBaseResponses.size() > 2 ? fBaseResponses.size() : 1;
47 Vector y(nOut);
48 ComputeImpl(x.data(), y.data());
49 return y;
50 }
51
52 /// Compute model prediction on a single event.
53 inline std::vector<Value_t> Compute(std::vector<Value_t> const &x) const { return Compute<std::vector<Value_t>>(x); }
54
55 /// Compute model prediction on a flat batch of events.
56 ///
57 /// The input must be flat row-major with `cols` features per event and the
58 /// output is a flat row-major vector with one row of outputs per event,
59 /// contiguous at `y[row * nOut + k]`. `nOut` is the number of model output
60 /// values per event: one for regression and binary classification models,
61 /// the number of classes for multiclass ones.
62 std::vector<Value_t> Compute(std::span<const Value_t> x, unsigned int cols) const;
63
64 static RBDT LoadXGBoost(std::string const &jsonPath);
65
66private:
67 /// Private default constructor, used by the public LoadXGBoost() factory.
68 RBDT() = default;
69
70 /// Map from XGBoost to RBDT indices.
71 using IndexMap = std::unordered_map<int, int>;
72
73 void Softmax(const Value_t *array, Value_t *out) const;
74 void ComputeImpl(const Value_t *array, Value_t *out) const;
75 Value_t EvaluateBinary(const Value_t *array) const;
76 static void correctIndices(std::span<int> indices, IndexMap const &nodeIndices, IndexMap const &leafIndices);
79
80 std::vector<int> fRootIndices;
81 std::vector<unsigned int> fCutIndices;
82 std::vector<Value_t> fCutValues;
83 std::vector<int> fLeftIndices;
84 std::vector<int> fRightIndices;
85 std::vector<Value_t> fResponses;
86 std::vector<int> fTreeNumbers;
87 std::vector<Value_t> fBaseResponses;
89 bool fLogistic = false;
90};
91
92} // namespace Experimental
93
94} // namespace TMVA
95
96#endif // TMVA_RBDT
std::vector< Value_t > fCutValues
Definition RBDT.hxx:82
static void terminateTree(TMVA::Experimental::RBDT &ff, int &nPreviousNodes, int &nPreviousLeaves, IndexMap &nodeIndices, IndexMap &leafIndices, int &treesSkipped)
Definition RBDT.cxx:197
RBDT()=default
Private default constructor, used by the public LoadXGBoost() factory.
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::vector< int > fRightIndices
Definition RBDT.hxx:84
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
std::vector< int > fTreeNumbers
Definition RBDT.hxx:86
Value_t EvaluateBinary(const Value_t *array) const
Definition RBDT.cxx:154
std::vector< Value_t > fResponses
Definition RBDT.hxx:85
std::vector< Value_t > fBaseResponses
Definition RBDT.hxx:87
std::vector< Value_t > Compute(std::vector< Value_t > const &x) const
Compute model prediction on a single event.
Definition RBDT.hxx:53
Vector Compute(const Vector &x) const
Compute model prediction on a single event.
Definition RBDT.hxx:44
std::vector< unsigned int > fCutIndices
Definition RBDT.hxx:81
void ComputeImpl(const Value_t *array, Value_t *out) const
Definition RBDT.cxx:141
std::vector< int > fRootIndices
Definition RBDT.hxx:80
std::vector< int > fLeftIndices
Definition RBDT.hxx:83
Double_t y[n]
Definition legend1.C:17
Double_t x[n]
Definition legend1.C:17
create variable transformations