Belle II Software light-2609-luna
basf2_mva_util.py
1
8
9import tempfile
10import numpy as np
11
12from basf2 import B2WARNING
13import basf2_mva
14
15
16def chain2dict(chain, tree_columns, dict_columns=None, max_entries=None):
17 """
18 Convert a ROOT.TChain into a dictionary of np.arrays
19 @param chain the ROOT.TChain
20 @param tree_columns the column (or branch) names in the tree
21 @param dict_columns the corresponding column names in the dictionary
22 """
23 if len(tree_columns) == 0:
24 return dict()
25 if dict_columns is None:
26 dict_columns = tree_columns
27 try:
28 from ROOT import RDataFrame
29 rdf = RDataFrame(chain)
30 if max_entries is not None:
31 nEntries = rdf.Count().GetValue()
32 if nEntries > max_entries:
33 B2WARNING(
34 "basf2_mva_util (chain2dict): Number of entries in the chain is larger than the maximum allowed entries: " +
35 str(nEntries) +
36 " > " +
37 str(max_entries))
38 skip = nEntries // max_entries
39 rdf_subset = rdf.Filter("rdfentry_ % " + str(skip) + " == 0")
40 rdf = rdf_subset
41
42 d = np.column_stack(list(rdf.AsNumpy(tree_columns).values()))
43 d = np.core.records.fromarrays(d.transpose(), names=dict_columns)
44 except ImportError:
45 d = {column: np.zeros((chain.GetEntries(),)) for column in dict_columns}
46 for iEvent, event in enumerate(chain):
47 for dict_column, tree_column in zip(dict_columns, tree_columns):
48 d[dict_column][iEvent] = getattr(event, tree_column)
49 return d
50
51
52def calculate_auc_efficiency_vs_purity(p, t, w=None):
53 """
54 Calculates the area under the efficiency-purity curve
55 @param p np.array filled with the probability output of a classifier
56 @param t np.array filled with the target (0 or 1)
57 @param w None or np.array filled with weights
58 """
59 if w is None:
60 w = np.ones(t.shape)
61
62 wt = w * t
63
64 N = np.sum(w)
65 T = np.sum(wt)
66
67 index = np.argsort(p)
68 efficiency = (T - np.cumsum(wt[index])) / float(T)
69 purity = (T - np.cumsum(wt[index])) / (N - np.cumsum(w[index]))
70 purity = np.where(np.isnan(purity), 0, purity)
71 return np.abs(np.trapz(purity, efficiency))
72
73
74def calculate_auc_efficiency_vs_background_retention(p, t, w=None):
75 """
76 Calculates the area under the efficiency-background_retention curve (AUC ROC)
77 @param p np.array filled with the probability output of a classifier
78 @param t np.array filled with the target (0 or 1)
79 @param w None or np.array filled with weights
80 """
81 if w is None:
82 w = np.ones(t.shape)
83
84 wt = w * t
85
86 N = np.sum(w)
87 T = np.sum(wt)
88
89 index = np.argsort(p)
90 efficiency = (T - np.cumsum(wt[index])) / float(T)
91 background_retention = (N - T - np.cumsum((np.abs(1 - t) * w)[index])) / float(N - T)
92 return np.abs(np.trapz(efficiency, background_retention))
93
94
95def calculate_flatness(f, p, w=None):
96 """
97 Calculates the flatness of a feature under cuts on a signal probability
98 @param f the feature values
99 @param p the probability values
100 @param w optional weights
101 @return the mean standard deviation between the local and global cut selection efficiency
102 """
103 quantiles = list(range(101))
104 binning_feature = np.unique(np.percentile(f, q=quantiles))
105 binning_probability = np.unique(np.percentile(p, q=quantiles))
106 if len(binning_feature) < 2:
107 binning_feature = np.array([np.min(f) - 1, np.max(f) + 1])
108 if len(binning_probability) < 2:
109 binning_probability = np.array([np.min(p) - 1, np.max(p) + 1])
110 hist_n, _ = np.histogramdd(np.c_[p, f],
111 bins=[binning_probability, binning_feature],
112 weights=w)
113 hist_inc = hist_n.sum(axis=1)
114 hist_inc /= hist_inc.sum(axis=0)
115 hist_n /= hist_n.sum(axis=0)
116 hist_n = hist_n.cumsum(axis=0)
117 hist_inc = hist_inc.cumsum(axis=0)
118 diff = (hist_n.T - hist_inc)**2
119 return np.sqrt(diff.sum() / (100 * 99))
120
121
122class Method:
123 """
124 Wrapper class providing an interface to the method stored under the given identifier.
125 It loads the Options, can apply the expert and train new ones using the current as a prototype.
126 This class is used by the basf_mva_evaluation tools
127 """
128
129 def __init__(self, identifier):
130 """
131 Load a method stored under the given identifier
132 @param identifier identifying the method
133 """
134 # Always avoid the top-level 'import ROOT'.
135 import ROOT # noqa
136 # Initialize all the available interfaces
137 ROOT.Belle2.MVA.AbstractInterface.initSupportedInterfaces()
138
139 self.identifier = identifier
140
141 self.weightfile = ROOT.Belle2.MVA.Weightfile.load(self.identifier)
142
143 self.general_options = basf2_mva.GeneralOptions()
144 self.general_options.load(self.weightfile.getXMLTree())
145
146 # This piece of code should be correct but leads to random segmentation faults
147 # inside python, llvm or pyroot, therefore we use the more dirty code below
148 # Ideas why this is happening:
149 # 1. Ownership of the unique_ptr returned by getOptions()
150 # 2. Some kind of object slicing, although pyroot identifies the correct type
151 # 3. Bug in pyroot
152 # interfaces = ROOT.Belle2.MVA.AbstractInterface.getSupportedInterfaces()
153 # self.interface = interfaces[self.general_options.m_method]
154 # self.specific_options = self.interface.getOptions()
155
156
158 if self.general_options.m_method == "FastBDT":
159 self.specific_options = basf2_mva.FastBDTOptions()
160 elif self.general_options.m_method == "TMVAClassification":
161 self.specific_options = basf2_mva.TMVAOptionsClassification()
162 elif self.general_options.m_method == "TMVARegression":
163 self.specific_options = basf2_mva.TMVAOptionsRegression()
164 elif self.general_options.m_method == "Python":
165 self.specific_options = basf2_mva.PythonOptions()
166 elif self.general_options.m_method == "PDF":
167 self.specific_options = basf2_mva.PDFOptions()
168 elif self.general_options.m_method == "Combination":
169 self.specific_options = basf2_mva.CombinationOptions()
170 elif self.general_options.m_method == "Reweighter":
171 self.specific_options = basf2_mva.ReweighterOptions()
172 elif self.general_options.m_method == "Trivial":
173 self.specific_options = basf2_mva.TrivialOptions()
174 elif self.general_options.m_method == "ONNX":
175 self.specific_options = basf2_mva.ONNXOptions()
176 else:
177 raise RuntimeError("Unknown method " + self.general_options.m_method)
178
179 self.specific_options.load(self.weightfile.getXMLTree())
180
181 variables = [str(v) for v in self.general_options.m_variables]
182 importances = self.weightfile.getFeatureImportance()
183
184
185 self.importances = {k: importances[k] for k in variables}
186
187 self.variables = list(sorted(variables, key=lambda v: self.importances.get(v, 0.0)))
188
189 self.root_variables = [ROOT.Belle2.MakeROOTCompatible.makeROOTCompatible(v) for v in self.variables]
190
191 self.root_importances = {k: importances[k] for k in self.root_variables}
192
193 self.description = str(basf2_mva.info(self.identifier))
194
195 self.spectators = [str(v) for v in self.general_options.m_spectators]
196
197 self.root_spectators = [ROOT.Belle2.MakeROOTCompatible.makeROOTCompatible(v) for v in self.spectators]
198
199 def train_teacher(self, datafiles, treename, general_options=None, specific_options=None):
200 """
201 Train a new method using this method as a prototype
202 @param datafiles the training datafiles
203 @param treename the name of the tree containing the training data
204 @param general_options general options given to basf2_mva.teacher
205 (if None the options of this method are used)
206 @param specific_options specific options given to basf2_mva.teacher
207 (if None the options of this method are used)
208 """
209 # Always avoid the top-level 'import ROOT'.
210 import ROOT # noqa
211 if isinstance(datafiles, str):
212 datafiles = [datafiles]
213 if general_options is None:
214 general_options = self.general_options
215 if specific_options is None:
216 specific_options = self.specific_options
217
218 with tempfile.TemporaryDirectory() as tempdir:
219 identifier = tempdir + "/weightfile.xml"
220
221 general_options.m_datafiles = basf2_mva.vector(*datafiles)
222 general_options.m_identifier = identifier
223
224 basf2_mva.teacher(general_options, specific_options)
225
226 method = Method(identifier)
227 return method
228
229 def apply_expert(self, datafiles, treename):
230 """
231 Apply the expert of the method to data and return the calculated probability and the target
232 @param datafiles the datafiles
233 @param treename the name of the tree containing the data
234 """
235 import ROOT # noqa
236 if isinstance(datafiles, str):
237 datafiles = [datafiles]
238 with tempfile.TemporaryDirectory() as tempdir:
239 identifier = tempdir + "/weightfile.xml"
240 ROOT.Belle2.MVA.Weightfile.save(self.weightfile, identifier)
241
242 rootfilename = tempdir + '/expert.root'
243 basf2_mva.expert(basf2_mva.vector(identifier),
244 basf2_mva.vector(*datafiles),
245 treename,
246 rootfilename)
247 chain = ROOT.TChain("variables")
248 chain.Add(rootfilename)
249
250 expert_target = identifier + '_' + self.general_options.m_target_variable
251 stripped_expert_target = self.identifier + '_' + self.general_options.m_target_variable
252
253 output_names = [self.identifier]
254 branch_names = [
255 ROOT.Belle2.MakeROOTCompatible.makeROOTCompatible(identifier),
256 ]
257 if self.general_options.m_nClasses > 2:
258 output_names = [self.identifier+f'_{i}' for i in range(self.general_options.m_nClasses)]
259 branch_names = [
260 ROOT.Belle2.MakeROOTCompatible.makeROOTCompatible(
261 identifier +
262 f'_{i}') for i in range(
263 self.general_options.m_nClasses)]
264
265 d = chain2dict(
266 chain,
267 [*branch_names, ROOT.Belle2.MakeROOTCompatible.makeROOTCompatible(expert_target)],
268 [*output_names, stripped_expert_target])
269
270 return (d[str(self.identifier)] if self.general_options.m_nClasses <= 2 else np.array([d[x]
271 for x in output_names]).T), d[stripped_expert_target]
272
273
274def create_onnx_mva_weightfile(onnx_model_path, **kwargs):
275 """
276 Create an MVA Weightfile for ONNX
277
278 Parameters:
279 kwargs: keyword arguments to set the options in the weightfile. They are
280 directly mapped to member variable names of the option classes with ``m_``
281 added automatically. First, GeneralOptions are tried and the remaining
282 arguments are passed to ONNXOptions.
283
284 Returns:
285 Weightfile object containing the ONNX model and options
286
287 Example:
288 .. code-block:: python
289
290 >>> weightfile = create_onnx_mva_weightfile(
291 ... "model.onnx",
292 ... outputName="probabilities",
293 ... variables=["variable1", "variable2"],
294 ... target_variable="isSignal"
295 ...)
296 >>> weightfile.save("model.root")
297
298 """
299 general_options = basf2_mva.GeneralOptions()
300 onnx_options = basf2_mva.ONNXOptions()
301 general_options.m_method = onnx_options.getMethod()
302
303 # fill everything that exists in general options from kwargs
304 for k, v in list(kwargs.items()):
305 m_k = f"m_{k}"
306 if hasattr(general_options, m_k):
307 setattr(general_options, m_k, v)
308 kwargs.pop(k)
309
310 # for the rest try to set members of specific options
311 for k, v in list(kwargs.items()):
312 m_k = f"m_{k}"
313 if not hasattr(onnx_options, m_k):
314 raise AttributeError(f"No member named {m_k} in ONNXOptions.")
315 setattr(onnx_options, m_k, v)
316
317 w = basf2_mva.Weightfile()
318 w.addOptions(general_options)
319 w.addOptions(onnx_options)
320 w.addFile("ONNX_Modelfile", str(onnx_model_path))
321 return w
dict importances
Dictionary of the variable importances calculated by the method.
specific_options
Specific options of the method.
description
Description of the method as a xml string returned by basf2_mva.info.
__init__(self, identifier)
dict root_importances
Dictionary of the variables sorted by their importance but with root compatoble variable names.
variables
List of variables sorted by their importance.
weightfile
Weightfile of the method.
list root_spectators
List of spectators with root compatible names.
list root_variables
List of the variable importances calculated by the method, but with the root compatible variable name...
list spectators
List of spectators.
train_teacher(self, datafiles, treename, general_options=None, specific_options=None)
general_options
General options of the method.
apply_expert(self, datafiles, treename)
identifier
Identifier of the method.