15void train(
const std::string &
filename)
18 auto output =
TFile::Open(
"TMVARR.root",
"RECREATE");
20 output,
"!V:!DrawProgressBar:AnalysisType=Classification");
29 const std::vector<std::string>
variables = {
"var1",
"var2",
"var3",
"var4"};
35 dataloader->PrepareTrainingAndTestTree(
"",
"");
39 factory->TrainAllMethods();
45 const std::string
filename = std::string(
gROOT->GetTutorialDir()) +
"/machine_learning/data/tmva_class_example.root";
49 RReader model(
"tmva003_BDT/weights/tmva003_BDT.weights.xml");
53 auto variables = model.GetVariableNames();
67 auto prediction = model.Compute(std::vector<float>{0.5f, 1.0f, -0.2f, 1.5f});
68 std::cout <<
"Single-event inference: " <<
prediction[0] <<
"\n\n";
77 auto df2 = df.Range(3);
78 const std::size_t nEvents = 3;
81 for (std::size_t
v = 0;
v <
nVars;
v++)
85 std::vector<float>
x(nEvents *
nVars);
86 for (std::size_t i = 0; i < nEvents; i++)
87 for (std::size_t
v = 0;
v <
nVars;
v++)
92 auto y = model.Compute(std::span<const float>(
x.data(),
x.size()));
94 std::cout <<
"Flat input for inference on " << nEvents <<
" events with " <<
nVars <<
" variables each:\n";
95 for (std::size_t i = 0; i < nEvents; i++) {
96 std::cout <<
" Event " << i <<
":";
97 for (std::size_t
v = 0;
v <
nVars;
v++)
98 std::cout <<
" " <<
x[i *
nVars +
v];
102 std::cout <<
"Prediction performed on multiple events:\n";
103 for (std::size_t i = 0; i < nEvents; i++)
104 std::cout <<
" Event " << i <<
": " <<
y[i] <<
"\n";
113 return df2.Histo1D({
treename.c_str(),
";BDT score;N_{Events}", 30, -0.5, 0.5},
"y");
121 auto c =
new TCanvas(
"",
"", 800, 800);
123 sig->SetLineColor(
kRed);
125 sig->SetLineWidth(2);
126 bkg->SetLineWidth(2);
128 sig->Draw(
"HIST SAME");
132 legend.AddEntry(
"TreeS",
"Signal",
"l");
133 legend.AddEntry(
"TreeB",
"Background",
"l");
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 data
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
R__EXTERN TStyle * gStyle
ROOT's RDataFrame offers a modern, high-level interface for analysis of data stored in TTree ,...
static TFile * Open(const char *name, Option_t *option="", const char *ftitle="", Int_t compress=ROOT::RCompressionSetting::EDefaults::kUseCompiledDefault, Int_t netopt=0)
Create / open a file.
This class displays a legend box (TPaveText) containing several legend entries.
A replacement for the TMVA::Reader legacy interface.
This is the main MVA steering class.
void SetOptStat(Int_t stat=1)
The type of information printed in the histogram statistics box can be selected via the parameter mod...
A TTree represents a columnar dataset.
void variables(TString dataset, TString fin="TMVA.root", TString dirName="InputVariables_Id", TString title="TMVA Input Variables", Bool_t isRegression=kFALSE, Bool_t useTMVAStyle=kTRUE)