ROOT
master
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
13
from
os.path
import
exists
14
15
import
ROOT
16
17
# check if the input file exists
18
modelFile =
"HiggsModel.onnx"
19
modelName =
"HiggsModel"
20
21
if
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
25
parser =
ROOT.TMVA.Experimental.SOFIE.RModelParser_ONNX
()
26
model =
parser.Parse
(modelFile)
27
28
# generating inference code
29
model.Generate
()
30
model.OutputGenerated
(
"Higgs_trained_model_generated.hxx"
)
31
model.PrintGenerated
()
32
33
# compile using ROOT JIT trained model
34
print(
"compiling SOFIE model and inference helper...."
)
35
ROOT.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.
42
ROOT.gInterpreter.Declare
(
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.
49
ROOT.gInterpreter.Declare
(
"""
50
double 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
59
inputFile = str(
ROOT.gROOT.GetTutorialDir
()) +
"/machine_learning/data/Higgs_data.root"
60
df1 =
ROOT.RDataFrame
(
"sig_tree"
, inputFile)
61
h1 =
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
63
df2 =
ROOT.RDataFrame
(
"bkg_tree"
, inputFile)
64
h2 =
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.
67
ROOT.RDF.RunGraphs
([h1, h2])
68
69
print(
"Number of signal entries"
,
h1.GetEntries
())
70
print(
"Number of background entries"
,
h2.GetEntries
())
71
72
h1.SetLineColor
(
"kRed"
)
73
h2.SetLineColor
(
"kBlue"
)
74
75
c1 =
ROOT.TCanvas
()
76
ROOT.gStyle.SetOptStat
(0)
77
78
h2.DrawClone
()
79
h1.DrawClone
(
"SAME"
)
TRangeDynCast
ROOT::Detail::TRangeCast< T, true > TRangeDynCast
TRangeDynCast is an adapter class that allows the typed iteration through a TCollection.
Definition
TCollection.h:359
ROOT::Detail::TRangeCast
Definition
TCollection.h:312
ROOT::RDataFrame
ROOT's RDataFrame offers a modern, high-level interface for analysis of data stored in TTree ,...
Definition
RDataFrame.hxx:50
tutorials
machine_learning
TMVA_SOFIE_RDataFrame.py
ROOTmaster - Reference Guide Generated on Fri Oct 2 2026 04:42:53 (GVA Time) using Doxygen 1.10.0