Logo ROOT  
Reference Guide
 
Loading...
Searching...
No Matches
TMVA_SOFIE_RDataFrame.py
Go to the documentation of this file.
1### \file
2### \ingroup tutorial_ml
3### \notebook -nodraw
4### Example of inference with SOFIE and RDataFrame, of a model trained with PyTorch.
5### First, generate the input ONNX model by running `TMVA_SOFIE_PyTorch_HiggsModel.py`.
6###
7### This tutorial parses the input model and runs the inference using ROOT's JITing capability.
8###
9### \macro_code
10### \macro_output
11### \author Lorenzo Moneta
12
13from os.path import exists
14
15import ROOT
16
17# check if the input file exists
18modelFile = "HiggsModel.onnx"
19modelName = "HiggsModel"
20
21if not exists(modelFile):
22 raise FileNotFoundError("You need to run TMVA_SOFIE_PyTorch_HiggsModel.py to generate the ONNX trained model")
23
24# parse the input ONNX model into RModel object
26model = parser.Parse(modelFile)
27
28# generating inference code
30model.OutputGenerated("Higgs_trained_model_generated.hxx")
32
33# compile using ROOT JIT trained model
34print("compiling SOFIE model and inference helper....")
35ROOT.gInterpreter.Declare('#include "Higgs_trained_model_generated.hxx"\n#include <array>\n#include <vector>')
36
37# A SOFIE Session holds the model weights and the intermediate buffers and is not
38# thread-safe: create one Session per RDataFrame processing slot and dispatch on
39# the slot number. This tutorial runs single-threaded, so a single Session is enough.
40# The weights file name is passed explicitly because the generated header was
41# written under a custom name.
43 'std::vector<TMVA_SOFIE_' + modelName + '::Session> sofie_sessions{TMVA_SOFIE_' + modelName +
44 '::Session("Higgs_trained_model_generated.dat")};')
45
46# Declare the inference function for RDataFrame: it assembles the model input
47# tensor from the columns and evaluates the model. The column order must match
48# the ordering of the model input tensor.
50double sofie_eval(unsigned int slot, float m_jj, float m_jjj, float m_lv, float m_jlv, float m_bb, float m_wbb,
51 float m_wwbb)
52{
53 std::array<float, 7> input{m_jj, m_jjj, m_lv, m_jlv, m_bb, m_wbb, m_wwbb};
54 return sofie_sessions[slot].infer(input.data())[0];
55}
56""")
57
58# run inference over input data
59inputFile = str(ROOT.gROOT.GetTutorialDir()) + "/machine_learning/data/Higgs_data.root"
60df1 = ROOT.RDataFrame("sig_tree", inputFile)
61h1 = df1.Define("DNN_Value", "sofie_eval(rdfslot_,m_jj, m_jjj, m_lv, m_jlv, m_bb, m_wbb, m_wwbb)").Histo1D(("h_sig", "", 100, 0, 1),"DNN_Value")
62
63df2 = ROOT.RDataFrame("bkg_tree", inputFile)
64h2 = df2.Define("DNN_Value", "sofie_eval(rdfslot_,m_jj, m_jjj, m_lv, m_jlv, m_bb, m_wbb, m_wwbb)").Histo1D(("h_bkg", "", 100, 0, 1),"DNN_Value")
65
66# run over the input data once, combining both RDataFrame graphs.
67ROOT.RDF.RunGraphs([h1, h2])
68
69print("Number of signal entries",h1.GetEntries())
70print("Number of background entries",h2.GetEntries())
71
72h1.SetLineColor("kRed")
73h2.SetLineColor("kBlue")
74
75c1 = ROOT.TCanvas()
77
79h1.DrawClone("SAME")
ROOT::Detail::TRangeCast< T, true > TRangeDynCast
TRangeDynCast is an adapter class that allows the typed iteration through a TCollection.
ROOT's RDataFrame offers a modern, high-level interface for analysis of data stored in TTree ,...