Logo ROOT  
Reference Guide
 
Loading...
Searching...
No Matches
TMVA_SOFIE_RDataFrame.C File Reference

Detailed Description

View in nbviewer Open in SWAN
This macro provides an example of using a trained model with PyTorch and make inference using SOFIE and RDataFrame This macro uses as input an ONNX model generated with the Python tutorial TMVA_SOFIE_PyTorch_HiggsModel.py You need to run that macro before to generate the trained PyTorch model and also the corresponding header file with SOFIE which can then be used for inference

Execute in this order:

// need to add the current directory (from where we are running this macro)
// to the include path for Cling
#include "HiggsModel.hxx"
#include <array>
#include <vector>
void TMVA_SOFIE_RDataFrame(int nthreads = 2){
std::string inputFileName = "Higgs_data.root";
std::string inputFile = std::string{gROOT->GetTutorialDir()} + "/machine_learning/data/" + inputFileName;
int nslots = df1.GetNSlots();
std::cout << "Running using " << nslots << " threads" << std::endl;
// A SOFIE Session holds the model weights and the intermediate buffers and is
// not thread-safe: create one Session per RDataFrame processing slot and use
// the slot number in the DefineSlot functor to dispatch to the right one.
// The Session default constructor reads the weights from the default weight
// file (HiggsModel.dat in this case).
std::vector<TMVA_SOFIE_HiggsModel::Session> sessions(nslots);
// The functor assembles the model input tensor from the RDataFrame columns
// and evaluates the model. The column order must match the ordering of the
// model input tensor.
auto evalModel = [&sessions](unsigned int slot, float m_jj, float m_jjj, float m_lv, float m_jlv, float m_bb,
float m_wbb, float m_wwbb) {
std::array<float, 7> input{m_jj, m_jjj, m_lv, m_jlv, m_bb, m_wbb, m_wwbb};
auto result = sessions[slot].infer(input.data());
return result[0];
};
auto h1 = df1.DefineSlot("DNN_Value", evalModel, {"m_jj", "m_jjj", "m_lv", "m_jlv", "m_bb", "m_wbb", "m_wwbb"})
.Histo1D({"h_sig", "", 100, 0, 1}, "DNN_Value");
auto h2 = df2.DefineSlot("DNN_Value", evalModel, {"m_jj", "m_jjj", "m_lv", "m_jlv", "m_bb", "m_wbb", "m_wwbb"})
.Histo1D({"h_bkg", "", 100, 0, 1}, "DNN_Value");
h2->SetLineColor(kBlue);
auto c1 = new TCanvas();
h2->DrawClone();
h1->DrawClone("SAME");
c1->BuildLegend();
}
#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
Running using 2 threads
Author
Lorenzo Moneta

Definition in file TMVA_SOFIE_RDataFrame.C.