Logo ROOT  
Reference Guide
 
Loading...
Searching...
No Matches
TMVA_SOFIE_RDataFrame_JIT.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 the SOFIE header generated from the ONNX model
7/// with the TMVA_SOFIE_PyTorch_HiggsModel.py tutorial
8/// You need to run that macro before this one.
9/// In this case we are parsing the input file and then run the inference in the same
10/// macro making use of the ROOT JITing capability
11///
12///
13/// \macro_code
14/// \macro_output
15/// \author Lorenzo Moneta
16
17/// Function to compile the generated model with the ROOT JIT and to declare the
18/// Session objects and the inference function used by RDataFrame.
19/// A SOFIE Session holds the model weights and the intermediate buffers and is
20/// not thread-safe: one Session per RDataFrame processing slot is created and
21/// the slot number is used to dispatch to the right one.
22/// Assume that the model name is the same as the header file name.
23void CompileModelForRDF(const std::string &headerModelFile, unsigned int ninputs, unsigned int nslots = 0)
24{
25
26 std::string modelName = headerModelFile.substr(0,headerModelFile.find(".hxx"));
27 std::string cmd =
28 std::string("#include \"") + headerModelFile + std::string("\"\n#include <array>\n#include <vector>");
29 auto ret = gInterpreter->Declare(cmd.c_str());
30 if (!ret)
31 throw std::runtime_error("Error compiling : " + cmd);
32 std::cout << "compiled : " << cmd << std::endl;
33
34 // Declare one Session per processing slot. The Session default constructor
35 // reads the weights from the default weight file (<modelName>.dat here).
36 if (nslots < 1)
37 nslots = 1;
38 cmd = "std::vector<TMVA_SOFIE_" + modelName + "::Session> sofie_sessions(" + std::to_string(nslots) + ");";
39 ret = gInterpreter->Declare(cmd.c_str());
40 if (!ret)
41 throw std::runtime_error("Error compiling : " + cmd);
42
43 // Declare the inference function for RDataFrame: it assembles the model
44 // input tensor from the columns and evaluates the model of the given slot.
45 std::string params;
46 std::string inputValues;
47 for (unsigned int i = 0; i < ninputs; i++) {
48 if (i > 0) {
49 params += ", ";
50 inputValues += ", ";
51 }
52 params += "float x" + std::to_string(i);
53 inputValues += "x" + std::to_string(i);
54 }
55 cmd = "double sofie_eval(unsigned int slot, " + params +
56 ") {\n"
57 " std::array<float, " +
58 std::to_string(ninputs) + "> input{" + inputValues +
59 "};\n"
60 " return sofie_sessions[slot].infer(input.data())[0];\n"
61 "}";
62 ret = gInterpreter->Declare(cmd.c_str());
63 if (!ret)
64 throw std::runtime_error("Error compiling : " + cmd);
65 std::cout << "compiled : " << cmd << std::endl;
66 std::cout << "Model is ready to be evaluated" << std::endl;
67 return;
68}
69
70void TMVA_SOFIE_RDataFrame_JIT(std::string modelName = "HiggsModel"){
71
72 // check if the input file exists
73 std::string modelHeaderFile = modelName + ".hxx";
74 if (gSystem->AccessPathName(modelHeaderFile.c_str())) {
75 Info("TMVA_SOFIE_RDataFrame", "You need to run TMVA_SOFIE_PyTorch_HiggsModel.py to generate the SOFIE header "
76 "for the PyTorch trained model");
77 return;
78 }
79
80 // check that also weigh file exists
81 std::string modelWeightFile = modelName + std::string(".dat");
83 Error("TMVA_SOFIE_RDataFrame","Generated weight file is missing");
84 return;
85 }
86
87 // now compile using ROOT JIT trained model (see function above)
88 CompileModelForRDF(modelHeaderFile,7);
89
90 std::string inputFileName = "Higgs_data.root";
91 std::string inputFile = std::string{gROOT->GetTutorialDir()} + "/machine_learning/data/" + inputFileName;
92
93 // The column order in the Define expressions must match the ordering of the
94 // model input tensor.
95 ROOT::RDataFrame df1("sig_tree", inputFile);
96 auto h1 = df1.Define("DNN_Value", "sofie_eval(rdfslot_,m_jj, m_jjj, m_lv, m_jlv, m_bb, m_wbb, m_wwbb)")
97 .Histo1D({"h_sig", "", 100, 0, 1}, "DNN_Value");
98
99 ROOT::RDataFrame df2("bkg_tree", inputFile);
100 auto h2 = df2.Define("DNN_Value", "sofie_eval(rdfslot_,m_jj, m_jjj, m_lv, m_jlv, m_bb, m_wbb, m_wwbb)")
101 .Histo1D({"h_bkg", "", 100, 0, 1}, "DNN_Value");
102
104 h2->SetLineColor(kBlue);
105
106 auto c1 = new TCanvas();
107 gStyle->SetOptStat(0);
108
109 h2->DrawClone();
110 h1->DrawClone("SAME");
111 c1->BuildLegend();
112
113
114}
@ 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.
void Info(const char *location, const char *msgfmt,...)
Use this function for informational messages.
Definition TError.cxx:241
void Error(const char *location, const char *msgfmt,...)
Use this function in case an error occurred.
Definition TError.cxx:208
#define gInterpreter
#define gROOT
Definition TROOT.h:417
R__EXTERN TStyle * gStyle
Definition TStyle.h:442
R__EXTERN TSystem * gSystem
Definition TSystem.h:582
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
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
return c1
Definition legend1.C:41
TH1F * h1
Definition legend1.C:5
modelName
Step 2 : Parse model and generate inference code with SOFIE.