Logo ROOT  
Reference Guide
 
Loading...
Searching...
No Matches
TMVA_SOFIE_RDataFrame.C
Go to the documentation of this file.
1/// \file
2/// \ingroup tutorial_ml
3/// \notebook -nodraw
4/// This macro provides an example of using a trained model with PyTorch
5/// and make inference using SOFIE and RDataFrame
6/// This macro uses as input an ONNX model generated with the
7/// Python tutorial TMVA_SOFIE_PyTorch_HiggsModel.py
8/// You need to run that macro before to generate the trained PyTorch model
9/// and also the corresponding header file with SOFIE which can then be used for inference
10///
11/// Execute in this order:
12/// ```
13/// python3 TMVA_SOFIE_PyTorch_HiggsModel.py
14/// root TMVA_SOFIE_RDataFrame.C
15/// ```
16///
17/// \macro_code
18/// \macro_output
19/// \author Lorenzo Moneta
20
21// need to add the current directory (from where we are running this macro)
22// to the include path for Cling
24#include "HiggsModel.hxx"
25
26#include <array>
27#include <vector>
28
29void TMVA_SOFIE_RDataFrame(int nthreads = 2){
30
31 std::string inputFileName = "Higgs_data.root";
32 std::string inputFile = std::string{gROOT->GetTutorialDir()} + "/machine_learning/data/" + inputFileName;
33
35
36 ROOT::RDataFrame df1("sig_tree", inputFile);
37 int nslots = df1.GetNSlots();
38 std::cout << "Running using " << nslots << " threads" << std::endl;
39
40 // A SOFIE Session holds the model weights and the intermediate buffers and is
41 // not thread-safe: create one Session per RDataFrame processing slot and use
42 // the slot number in the DefineSlot functor to dispatch to the right one.
43 // The Session default constructor reads the weights from the default weight
44 // file (HiggsModel.dat in this case).
45 std::vector<TMVA_SOFIE_HiggsModel::Session> sessions(nslots);
46
47 // The functor assembles the model input tensor from the RDataFrame columns
48 // and evaluates the model. The column order must match the ordering of the
49 // model input tensor.
50 auto evalModel = [&sessions](unsigned int slot, float m_jj, float m_jjj, float m_lv, float m_jlv, float m_bb,
51 float m_wbb, float m_wwbb) {
52 std::array<float, 7> input{m_jj, m_jjj, m_lv, m_jlv, m_bb, m_wbb, m_wwbb};
53 auto result = sessions[slot].infer(input.data());
54 return result[0];
55 };
56
57 auto h1 = df1.DefineSlot("DNN_Value", evalModel, {"m_jj", "m_jjj", "m_lv", "m_jlv", "m_bb", "m_wbb", "m_wwbb"})
58 .Histo1D({"h_sig", "", 100, 0, 1}, "DNN_Value");
59
60 ROOT::RDataFrame df2("bkg_tree", inputFile);
61 auto h2 = df2.DefineSlot("DNN_Value", evalModel, {"m_jj", "m_jjj", "m_lv", "m_jlv", "m_bb", "m_wbb", "m_wwbb"})
62 .Histo1D({"h_bkg", "", 100, 0, 1}, "DNN_Value");
63
65 h2->SetLineColor(kBlue);
66
67 auto c1 = new TCanvas();
69
70 h2->DrawClone();
71 h1->DrawClone("SAME");
72 c1->BuildLegend();
73
74}
#define R__ADD_INCLUDE_PATH(PATH)
Definition Rtypes.h:474
@ kRed
Definition Rtypes.h:66
@ kBlue
Definition Rtypes.h:66
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 input
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 result
#define gROOT
Definition TROOT.h:417
R__EXTERN TStyle * gStyle
Definition TStyle.h:442
ROOT's RDataFrame offers a modern, high-level interface for analysis of data stored in TTree ,...
virtual void SetLineColor(Color_t lcolor)
Set the line color.
Definition TAttLine.h:44
The Canvas class.
Definition TCanvas.h:23
virtual TObject * DrawClone(Option_t *option="") const
Draw a clone of this object in the current selected pad with: gROOT->SetSelectedPad(c1).
Definition TObject.cxx:318
void SetOptStat(Int_t stat=1)
The type of information printed in the histogram statistics box can be selected via the parameter mod...
Definition TStyle.cxx:1641
return c1
Definition legend1.C:41
TH1F * h1
Definition legend1.C:5
void EnableImplicitMT(UInt_t numthreads=0)
Enable ROOT's implicit multi-threading for all objects and methods that provide an internal paralleli...
Definition TROOT.cxx:617