Logo ROOT  
Reference Guide
 
Loading...
Searching...
No Matches
MethodBase.cxx
Go to the documentation of this file.
1// @(#)root/tmva $Id$
2// Author: Andreas Hoecker, Peter Speckmayer, Joerg Stelzer, Helge Voss, Kai Voss, Eckhard von Toerne, Jan Therhaag
3
4/**********************************************************************************
5 * Project: TMVA - a Root-integrated toolkit for multivariate data analysis *
6 * Package: TMVA *
7 * Class : MethodBase *
8 * *
9 * *
10 * Description: *
11 * Implementation (see header for description) *
12 * *
13 * Authors (alphabetical): *
14 * Andreas Hoecker <Andreas.Hocker@cern.ch> - CERN, Switzerland *
15 * Joerg Stelzer <Joerg.Stelzer@cern.ch> - CERN, Switzerland *
16 * Peter Speckmayer <Peter.Speckmayer@cern.ch> - CERN, Switzerland *
17 * Helge Voss <Helge.Voss@cern.ch> - MPI-K Heidelberg, Germany *
18 * Kai Voss <Kai.Voss@cern.ch> - U. of Victoria, Canada *
19 * Jan Therhaag <Jan.Therhaag@cern.ch> - U of Bonn, Germany *
20 * Eckhard v. Toerne <evt@uni-bonn.de> - U of Bonn, Germany *
21 * *
22 * Copyright (c) 2005-2011: *
23 * CERN, Switzerland *
24 * U. of Victoria, Canada *
25 * MPI-K Heidelberg, Germany *
26 * U. of Bonn, Germany *
27 * *
28 * Redistribution and use in source and binary forms, with or without *
29 * modification, are permitted according to the terms listed in LICENSE *
30 * (see tmva/doc/LICENSE) *
31 * *
32 **********************************************************************************/
33
34/*! \class TMVA::MethodBase
35\ingroup TMVA
36
37 Virtual base Class for all MVA method
38
39 MethodBase hosts several specific evaluation methods.
40
41 The kind of MVA that provides optimal performance in an analysis strongly
42 depends on the particular application. The evaluation factory provides a
43 number of numerical benchmark results to directly assess the performance
44 of the MVA training on the independent test sample. These are:
45
46 - The _signal efficiency_ at three representative background efficiencies
47 (which is 1 &minus; rejection).
48 - The _significance_ of an MVA estimator, defined by the difference
49 between the MVA mean values for signal and background, divided by the
50 quadratic sum of their root mean squares.
51 - The _separation_ of an MVA _x_, defined by the integral
52 \f[
53 \frac{1}{2} \int \frac{(S(x) - B(x))^2}{(S(x) + B(x))} dx
54 \f]
55 where
56 \f$ S(x) \f$ and \f$ B(x) \f$ are the signal and background distributions,
57 respectively. The separation is zero for identical signal and background MVA
58 shapes, and it is one for disjunctive shapes.
59 - The average, \f$ \int x \mu (S(x)) dx \f$, of the signal \f$ \mu_{transform} \f$.
60 The \f$ \mu_{transform} \f$ of an MVA denotes the transformation that yields
61 a uniform background distribution. In this way, the signal distributions
62 \f$ S(x) \f$ can be directly compared among the various MVAs. The stronger
63 \f$ S(x) \f$ peaks towards one, the better is the discrimination of the MVA.
64 The \f$ \mu_{transform} \f$ is
65 [documented here](http://tel.ccsd.cnrs.fr/documents/archives0/00/00/29/91/index_fr.html).
66
67 The MVA standard output also prints the linear correlation coefficients between
68 signal and background, which can be useful to eliminate variables that exhibit too
69 strong correlations.
70*/
71
72#include "TMVA/MethodBase.h"
73
74#include "TMVA/Config.h"
75#include "TMVA/Configurable.h"
76#include "TMVA/DataSetInfo.h"
77#include "TMVA/DataSet.h"
78#include "TMVA/Factory.h"
79#include "TMVA/IMethod.h"
80#include "TMVA/MsgLogger.h"
81#include "TMVA/PDF.h"
82#include "TMVA/Ranking.h"
83#include "TMVA/DataLoader.h"
84#include "TMVA/Tools.h"
85#include "TMVA/Results.h"
89#include "TMVA/RootFinder.h"
90#include "TMVA/Timer.h"
91#include "TMVA/TSpline1.h"
92#include "TMVA/Types.h"
96#include "TMVA/VariableInfo.h"
100#include "TMVA/Version.h"
101
102#include "TROOT.h"
103#include "TSystem.h"
104#include "TObjString.h"
105#include "TQObject.h"
106#include "TSpline.h"
107#include "TMatrix.h"
108#include "TMath.h"
109#include "TH1F.h"
110#include "TH2F.h"
111#include "TFile.h"
112#include "TGraph.h"
113#include "TXMLEngine.h"
114
115#include <iomanip>
116#include <iostream>
117#include <fstream>
118#include <sstream>
119#include <cstdlib>
120#include <algorithm>
121#include <limits>
122
123
124
125using std::endl;
126using std::atof;
127
128//const Int_t MethodBase_MaxIterations_ = 200;
130
131//const Int_t NBIN_HIST_PLOT = 100;
132const Int_t NBIN_HIST_HIGH = 10000;
133
134#ifdef _WIN32
135/* Disable warning C4355: 'this' : used in base member initializer list */
136#pragma warning ( disable : 4355 )
137#endif
138
139////////////////////////////////////////////////////////////////////////////////
140/// standard constructor
141
143 Types::EMVA methodType,
144 const TString& methodTitle,
145 DataSetInfo& dsi,
146 const TString& theOption) :
147 IMethod(),
148 Configurable ( theOption ),
149 fTmpEvent ( 0 ),
150 fRanking ( 0 ),
151 fInputVars ( 0 ),
152 fAnalysisType ( Types::kNoAnalysisType ),
153 fRegressionReturnVal ( 0 ),
154 fMulticlassReturnVal ( 0 ),
155 fDataSetInfo ( dsi ),
156 fSignalReferenceCut ( 0.5 ),
157 fSignalReferenceCutOrientation( 1. ),
158 fVariableTransformType ( Types::kSignal ),
159 fJobName ( jobName ),
160 fMethodName ( methodTitle ),
161 fMethodType ( methodType ),
162 fTestvar ( "" ),
163 fTMVATrainingVersion ( TMVA_VERSION_CODE ),
164 fROOTTrainingVersion ( ROOT_VERSION_CODE ),
165 fConstructedFromWeightFile ( kFALSE ),
166 fBaseDir ( 0 ),
167 fMethodBaseDir ( 0 ),
168 fFile ( 0 ),
169 fSilentFile (kFALSE),
170 fModelPersistence (kTRUE),
171 fWeightFile ( "" ),
172 fEffS ( 0 ),
173 fDefaultPDF ( 0 ),
174 fMVAPdfS ( 0 ),
175 fMVAPdfB ( 0 ),
176 fSplS ( 0 ),
177 fSplB ( 0 ),
178 fSpleffBvsS ( 0 ),
179 fSplTrainS ( 0 ),
180 fSplTrainB ( 0 ),
181 fSplTrainEffBvsS ( 0 ),
182 fVarTransformString ( "None" ),
183 fTransformationPointer ( 0 ),
184 fTransformation ( dsi, methodTitle ),
185 fVerbose ( kFALSE ),
186 fVerbosityLevelString ( "Default" ),
187 fHelp ( kFALSE ),
188 fHasMVAPdfs ( kFALSE ),
189 fIgnoreNegWeightsInTraining( kFALSE ),
190 fSignalClass ( 0 ),
191 fBackgroundClass ( 0 ),
192 fSplRefS ( 0 ),
193 fSplRefB ( 0 ),
194 fSplTrainRefS ( 0 ),
195 fSplTrainRefB ( 0 ),
196 fSetupCompleted (kFALSE)
197{
198 SetTestvarName();
199 fLogger->SetSource(GetName());
200
201// // default extension for weight files
202}
203
204////////////////////////////////////////////////////////////////////////////////
205/// constructor used for Testing + Application of the MVA,
206/// only (no training), using given WeightFiles
207
209 DataSetInfo& dsi,
210 const TString& weightFile ) :
211 IMethod(),
212 Configurable(""),
213 fTmpEvent ( 0 ),
214 fRanking ( 0 ),
215 fInputVars ( 0 ),
216 fAnalysisType ( Types::kNoAnalysisType ),
217 fRegressionReturnVal ( 0 ),
218 fMulticlassReturnVal ( 0 ),
219 fDataSetInfo ( dsi ),
220 fSignalReferenceCut ( 0.5 ),
221 fVariableTransformType ( Types::kSignal ),
222 fJobName ( "" ),
223 fMethodName ( "MethodBase" ),
224 fMethodType ( methodType ),
225 fTestvar ( "" ),
226 fTMVATrainingVersion ( 0 ),
227 fROOTTrainingVersion ( 0 ),
228 fConstructedFromWeightFile ( kTRUE ),
229 fBaseDir ( 0 ),
230 fMethodBaseDir ( 0 ),
231 fFile ( 0 ),
232 fSilentFile (kFALSE),
233 fModelPersistence (kTRUE),
234 fWeightFile ( weightFile ),
235 fEffS ( 0 ),
236 fDefaultPDF ( 0 ),
237 fMVAPdfS ( 0 ),
238 fMVAPdfB ( 0 ),
239 fSplS ( 0 ),
240 fSplB ( 0 ),
241 fSpleffBvsS ( 0 ),
242 fSplTrainS ( 0 ),
243 fSplTrainB ( 0 ),
244 fSplTrainEffBvsS ( 0 ),
245 fVarTransformString ( "None" ),
246 fTransformationPointer ( 0 ),
247 fTransformation ( dsi, "" ),
248 fVerbose ( kFALSE ),
249 fVerbosityLevelString ( "Default" ),
250 fHelp ( kFALSE ),
251 fHasMVAPdfs ( kFALSE ),
252 fIgnoreNegWeightsInTraining( kFALSE ),
253 fSignalClass ( 0 ),
254 fBackgroundClass ( 0 ),
255 fSplRefS ( 0 ),
256 fSplRefB ( 0 ),
257 fSplTrainRefS ( 0 ),
258 fSplTrainRefB ( 0 ),
259 fSetupCompleted (kFALSE)
260{
262// // constructor used for Testing + Application of the MVA,
263// // only (no training), using given WeightFiles
264}
265
266////////////////////////////////////////////////////////////////////////////////
267/// destructor
268
270{
271 // destructor
272 if (!fSetupCompleted) Log() << kWARNING <<Form("Dataset[%s] : ",DataInfo().GetName())<< "Calling destructor of method which got never setup" << Endl;
273
274 // destructor
275 if (fInputVars != 0) { fInputVars->clear(); delete fInputVars; }
276 if (fRanking != 0) delete fRanking;
277
278 // PDFs
279 if (fDefaultPDF!= 0) { delete fDefaultPDF; fDefaultPDF = 0; }
280 if (fMVAPdfS != 0) { delete fMVAPdfS; fMVAPdfS = 0; }
281 if (fMVAPdfB != 0) { delete fMVAPdfB; fMVAPdfB = 0; }
282
283 // Splines
284 if (fSplS) { delete fSplS; fSplS = 0; }
285 if (fSplB) { delete fSplB; fSplB = 0; }
286 if (fSpleffBvsS) { delete fSpleffBvsS; fSpleffBvsS = 0; }
287 if (fSplRefS) { delete fSplRefS; fSplRefS = 0; }
288 if (fSplRefB) { delete fSplRefB; fSplRefB = 0; }
289 if (fSplTrainRefS) { delete fSplTrainRefS; fSplTrainRefS = 0; }
290 if (fSplTrainRefB) { delete fSplTrainRefB; fSplTrainRefB = 0; }
291 if (fSplTrainEffBvsS) { delete fSplTrainEffBvsS; fSplTrainEffBvsS = 0; }
292
293 for (size_t i = 0; i < fEventCollections.size(); i++ ) {
294 if (fEventCollections.at(i)) {
295 for (std::vector<Event*>::const_iterator it = fEventCollections.at(i)->begin();
296 it != fEventCollections.at(i)->end(); ++it) {
297 delete (*it);
298 }
299 delete fEventCollections.at(i);
300 fEventCollections.at(i) = nullptr;
301 }
302 }
303
304 if (fRegressionReturnVal) delete fRegressionReturnVal;
305 if (fMulticlassReturnVal) delete fMulticlassReturnVal;
306}
307
308////////////////////////////////////////////////////////////////////////////////
309/// setup of methods
310
312{
313 // setup of methods
314
315 if (fSetupCompleted) Log() << kFATAL <<Form("Dataset[%s] : ",DataInfo().GetName())<< "Calling SetupMethod for the second time" << Endl;
316 InitBase();
317 DeclareBaseOptions();
318 Init();
319 DeclareOptions();
320 fSetupCompleted = kTRUE;
321}
322
323////////////////////////////////////////////////////////////////////////////////
324/// process all options
325/// the "CheckForUnusedOptions" is done in an independent call, since it may be overridden by derived class
326/// (sometimes, eg, fitters are used which can only be implemented during training phase)
327
329{
330 ProcessBaseOptions();
331 ProcessOptions();
332}
333
334////////////////////////////////////////////////////////////////////////////////
335/// check may be overridden by derived class
336/// (sometimes, eg, fitters are used which can only be implemented during training phase)
337
339{
340 CheckForUnusedOptions();
341}
342
343////////////////////////////////////////////////////////////////////////////////
344/// default initialization called by all constructors
345
347{
348 SetConfigDescription( "Configuration options for classifier architecture and tuning" );
349
351 fNbinsMVAoutput = gConfig().fVariablePlotting.fNbinsMVAoutput;
352 fNbinsH = NBIN_HIST_HIGH;
353
354 fSplTrainS = 0;
355 fSplTrainB = 0;
356 fSplTrainEffBvsS = 0;
357 fMeanS = -1;
358 fMeanB = -1;
359 fRmsS = -1;
360 fRmsB = -1;
361 fXmin = DBL_MAX;
362 fXmax = -DBL_MAX;
363 fTxtWeightsOnly = kTRUE;
364 fSplRefS = 0;
365 fSplRefB = 0;
366
367 fTrainTime = -1.;
368 fTestTime = -1.;
369
370 fRanking = 0;
371
372 // temporary until the move to DataSet is complete
373 fInputVars = new std::vector<TString>;
374 for (UInt_t ivar=0; ivar<GetNvar(); ivar++) {
375 fInputVars->push_back(DataInfo().GetVariableInfo(ivar).GetLabel());
376 }
377 fRegressionReturnVal = 0;
378 fMulticlassReturnVal = 0;
379
380 fEventCollections.resize( 2 );
381 fEventCollections.at(0) = 0;
382 fEventCollections.at(1) = 0;
383
384 // retrieve signal and background class index
385 if (DataInfo().GetClassInfo("Signal") != 0) {
386 fSignalClass = DataInfo().GetClassInfo("Signal")->GetNumber();
387 }
388 if (DataInfo().GetClassInfo("Background") != 0) {
389 fBackgroundClass = DataInfo().GetClassInfo("Background")->GetNumber();
390 }
391
392 SetConfigDescription( "Configuration options for MVA method" );
393 SetConfigName( TString("Method") + GetMethodTypeName() );
394}
395
396////////////////////////////////////////////////////////////////////////////////
397/// define the options (their key words) that can be set in the option string
398/// here the options valid for ALL MVA methods are declared.
399///
400/// know options:
401///
402/// - VariableTransform=None,Decorrelated,PCA to use transformed variables
403/// instead of the original ones
404/// - VariableTransformType=Signal,Background which decorrelation matrix to use
405/// in the method. Only the Likelihood
406/// Method can make proper use of independent
407/// transformations of signal and background
408/// - fNbinsMVAPdf = 50 Number of bins used to create a PDF of MVA
409/// - fNsmoothMVAPdf = 2 Number of times a histogram is smoothed before creating the PDF
410/// - fHasMVAPdfs create PDFs for the MVA outputs
411/// - V for Verbose output (!V) for non verbos
412/// - H for Help message
413
415{
416 DeclareOptionRef( fVerbose, "V", "Verbose output (short form of \"VerbosityLevel\" below - overrides the latter one)" );
417
418 DeclareOptionRef( fVerbosityLevelString="Default", "VerbosityLevel", "Verbosity level" );
419 AddPreDefVal( TString("Default") ); // uses default defined in MsgLogger header
420 AddPreDefVal( TString("Debug") );
421 AddPreDefVal( TString("Verbose") );
422 AddPreDefVal( TString("Info") );
423 AddPreDefVal( TString("Warning") );
424 AddPreDefVal( TString("Error") );
425 AddPreDefVal( TString("Fatal") );
426
427 // If True (default): write all training results (weights) as text files only;
428 // if False: write also in ROOT format (not available for all methods - will abort if not
429 fTxtWeightsOnly = kTRUE; // OBSOLETE !!!
430 fNormalise = kFALSE; // OBSOLETE !!!
431
432 DeclareOptionRef( fVarTransformString, "VarTransform", "List of variable transformations performed before training, e.g., \"D_Background,P_Signal,G,N_AllClasses\" for: \"Decorrelation, PCA-transformation, Gaussianisation, Normalisation, each for the given class of events ('AllClasses' denotes all events of all classes, if no class indication is given, 'All' is assumed)\"" );
433
434 DeclareOptionRef( fHelp, "H", "Print method-specific help message" );
435
436 DeclareOptionRef( fHasMVAPdfs, "CreateMVAPdfs", "Create PDFs for classifier outputs (signal and background)" );
437
438 DeclareOptionRef( fIgnoreNegWeightsInTraining, "IgnoreNegWeightsInTraining",
439 "Events with negative weights are ignored in the training (but are included for testing and performance evaluation)" );
440}
441
442////////////////////////////////////////////////////////////////////////////////
443/// the option string is decoded, for available options see "DeclareOptions"
444
446{
447 if (HasMVAPdfs()) {
448 // setting the default bin num... maybe should be static ? ==> Please no static (JS)
449 // You can't use the logger in the constructor!!! Log() << kINFO << "Create PDFs" << Endl;
450 // reading every PDF's definition and passing the option string to the next one to be read and marked
451 fDefaultPDF = new PDF( TString(GetName())+"_PDF", GetOptions(), "MVAPdf" );
452 fDefaultPDF->DeclareOptions();
453 fDefaultPDF->ParseOptions();
454 fDefaultPDF->ProcessOptions();
455 fMVAPdfB = new PDF( TString(GetName())+"_PDFBkg", fDefaultPDF->GetOptions(), "MVAPdfBkg", fDefaultPDF );
456 fMVAPdfB->DeclareOptions();
457 fMVAPdfB->ParseOptions();
458 fMVAPdfB->ProcessOptions();
459 fMVAPdfS = new PDF( TString(GetName())+"_PDFSig", fMVAPdfB->GetOptions(), "MVAPdfSig", fDefaultPDF );
460 fMVAPdfS->DeclareOptions();
461 fMVAPdfS->ParseOptions();
462 fMVAPdfS->ProcessOptions();
463
464 // the final marked option string is written back to the original methodbase
465 SetOptions( fMVAPdfS->GetOptions() );
466 }
467
468 TMVA::CreateVariableTransforms( fVarTransformString,
469 DataInfo(),
470 GetTransformationHandler(),
471 Log() );
472
473 if (!HasMVAPdfs()) {
474 if (fDefaultPDF!= 0) { delete fDefaultPDF; fDefaultPDF = 0; }
475 if (fMVAPdfS != 0) { delete fMVAPdfS; fMVAPdfS = 0; }
476 if (fMVAPdfB != 0) { delete fMVAPdfB; fMVAPdfB = 0; }
477 }
478
479 if (fVerbose) { // overwrites other settings
480 fVerbosityLevelString = TString("Verbose");
481 Log().SetMinType( kVERBOSE );
482 }
483 else if (fVerbosityLevelString == "Debug" ) Log().SetMinType( kDEBUG );
484 else if (fVerbosityLevelString == "Verbose" ) Log().SetMinType( kVERBOSE );
485 else if (fVerbosityLevelString == "Info" ) Log().SetMinType( kINFO );
486 else if (fVerbosityLevelString == "Warning" ) Log().SetMinType( kWARNING );
487 else if (fVerbosityLevelString == "Error" ) Log().SetMinType( kERROR );
488 else if (fVerbosityLevelString == "Fatal" ) Log().SetMinType( kFATAL );
489 else if (fVerbosityLevelString != "Default" ) {
490 Log() << kFATAL << "<ProcessOptions> Verbosity level type '"
491 << fVerbosityLevelString << "' unknown." << Endl;
492 }
493 Event::SetIgnoreNegWeightsInTraining(fIgnoreNegWeightsInTraining);
494}
495
496////////////////////////////////////////////////////////////////////////////////
497/// options that are used ONLY for the READER to ensure backward compatibility
498/// they are hence without any effect (the reader is only reading the training
499/// options that HAD been used at the training of the .xml weight file at hand
500
502{
503 DeclareOptionRef( fNormalise=kFALSE, "Normalise", "Normalise input variables" ); // don't change the default !!!
504 DeclareOptionRef( fUseDecorr=kFALSE, "D", "Use-decorrelated-variables flag" );
505 DeclareOptionRef( fVariableTransformTypeString="Signal", "VarTransformType",
506 "Use signal or background events to derive for variable transformation (the transformation is applied on both types of, course)" );
507 AddPreDefVal( TString("Signal") );
508 AddPreDefVal( TString("Background") );
509 DeclareOptionRef( fTxtWeightsOnly=kTRUE, "TxtWeightFilesOnly", "If True: write all training results (weights) as text files (False: some are written in ROOT format)" );
510 // Why on earth ?? was this here? Was the verbosity level option meant to 'disappear? Not a good idea i think..
511 // DeclareOptionRef( fVerbosityLevelString="Default", "VerboseLevel", "Verbosity level" );
512 // AddPreDefVal( TString("Default") ); // uses default defined in MsgLogger header
513 // AddPreDefVal( TString("Debug") );
514 // AddPreDefVal( TString("Verbose") );
515 // AddPreDefVal( TString("Info") );
516 // AddPreDefVal( TString("Warning") );
517 // AddPreDefVal( TString("Error") );
518 // AddPreDefVal( TString("Fatal") );
519 DeclareOptionRef( fNbinsMVAPdf = 60, "NbinsMVAPdf", "Number of bins used for the PDFs of classifier outputs" );
520 DeclareOptionRef( fNsmoothMVAPdf = 2, "NsmoothMVAPdf", "Number of smoothing iterations for classifier PDFs" );
521}
522
523
524////////////////////////////////////////////////////////////////////////////////
525/// call the Optimizer with the set of parameters and ranges that
526/// are meant to be tuned.
527
528std::map<TString,Double_t> TMVA::MethodBase::OptimizeTuningParameters(TString /* fomType */ , TString /* fitType */)
529{
530 // this is just a dummy... needs to be implemented for each method
531 // individually (as long as we don't have it automatized via the
532 // configuration string
533
534 Log() << kWARNING <<Form("Dataset[%s] : ",DataInfo().GetName())<< "Parameter optimization is not yet implemented for method "
535 << GetName() << Endl;
536 Log() << kWARNING <<Form("Dataset[%s] : ",DataInfo().GetName())<< "Currently we need to set hardcoded which parameter is tuned in which ranges"<<Endl;
537
538 return std::map<TString,Double_t>();
539}
540
541////////////////////////////////////////////////////////////////////////////////
542/// set the tuning parameters according to the argument
543/// This is just a dummy .. have a look at the MethodBDT how you could
544/// perhaps implement the same thing for the other Classifiers..
545
546void TMVA::MethodBase::SetTuneParameters(std::map<TString,Double_t> /* tuneParameters */)
547{
548}
549
550////////////////////////////////////////////////////////////////////////////////
551
553{
554 Data()->SetCurrentType(Types::kTraining);
555 Event::SetIsTraining(kTRUE); // used to set negative event weights to zero if chosen to do so
556
557 // train the MVA method
558 if (Help()) PrintHelpMessage();
559
560 // all histograms should be created in the method's subdirectory
561 if(!IsSilentFile()) BaseDir()->cd();
562
563 // once calculate all the transformation (e.g. the sequence of Decorr:Gauss:Decorr)
564 // needed for this classifier
565 GetTransformationHandler().CalcTransformations(Data()->GetEventCollection());
566
567 // call training of derived MVA
568 Log() << kDEBUG //<<Form("\tDataset[%s] : ",DataInfo().GetName())
569 << "Begin training" << Endl;
570 Long64_t nEvents = Data()->GetNEvents();
571 Timer traintimer( nEvents, GetName(), kTRUE );
572 Train();
573 Log() << kDEBUG //<<Form("Dataset[%s] : ",DataInfo().GetName()
574 << "\tEnd of training " << Endl;
575 SetTrainTime(traintimer.ElapsedSeconds());
576 Log() << kINFO //<<Form("Dataset[%s] : ",DataInfo().GetName())
577 << "Elapsed time for training with " << nEvents << " events: "
578 << traintimer.GetElapsedTime() << " " << Endl;
579
580 Log() << kDEBUG //<<Form("Dataset[%s] : ",DataInfo().GetName())
581 << "\tCreate MVA output for ";
582
583 // create PDFs for the signal and background MVA distributions (if required)
584 if (DoMulticlass()) {
585 Log() <<Form("[%s] : ",DataInfo().GetName())<< "Multiclass classification on training sample" << Endl;
586 AddMulticlassOutput(Types::kTraining);
587 }
588 else if (!DoRegression()) {
589
590 Log() <<Form("[%s] : ",DataInfo().GetName())<< "classification on training sample" << Endl;
591 AddClassifierOutput(Types::kTraining);
592 if (HasMVAPdfs()) {
593 CreateMVAPdfs();
594 AddClassifierOutputProb(Types::kTraining);
595 }
596
597 } else {
598
599 Log() <<Form("Dataset[%s] : ",DataInfo().GetName())<< "regression on training sample" << Endl;
600 AddRegressionOutput( Types::kTraining );
601
602 if (HasMVAPdfs() ) {
603 Log() <<Form("Dataset[%s] : ",DataInfo().GetName())<< "Create PDFs" << Endl;
604 CreateMVAPdfs();
605 }
606 }
607
608 // write the current MVA state into stream
609 // produced are one text file and one ROOT file
610 if (fModelPersistence ) WriteStateToFile();
611
612 // produce standalone make class (presently only supported for classification)
613 if ((!DoRegression()) && (fModelPersistence)) MakeClass();
614
615 // write additional monitoring histograms to main target file (not the weight file)
616 // again, make sure the histograms go into the method's subdirectory
617 if(!IsSilentFile())
618 {
619 BaseDir()->cd();
620 WriteMonitoringHistosToFile();
621 }
622}
623
624////////////////////////////////////////////////////////////////////////////////
625
627{
628 if (!DoRegression()) Log() << kFATAL <<Form("Dataset[%s] : ",DataInfo().GetName())<< "Trying to use GetRegressionDeviation() with a classification job" << Endl;
629 Log() << kINFO <<Form("Dataset[%s] : ",DataInfo().GetName())<< "Create results for " << (type==Types::kTraining?"training":"testing") << Endl;
630 ResultsRegression* regRes = (ResultsRegression*)Data()->GetResults(GetMethodName(), Types::kTesting, Types::kRegression);
631 bool truncate = false;
632 TH1F* h1 = regRes->QuadraticDeviation( tgtNum , truncate, 1.);
633 stddev = sqrt(h1->GetMean());
634 truncate = true;
635 Double_t yq[1], xq[]={0.9};
636 h1->GetQuantiles(1,yq,xq);
637 TH1F* h2 = regRes->QuadraticDeviation( tgtNum , truncate, yq[0]);
638 stddev90Percent = sqrt(h2->GetMean());
639 delete h1;
640 delete h2;
641}
642
643////////////////////////////////////////////////////////////////////////////////
644/// Get al regression values in one call
646{
647 Long64_t nEvents = Data()->GetNEvents();
648 // use timer
649 Timer timer( nEvents, GetName(), kTRUE );
650
651 // Drawing the progress bar every event was causing a huge slowdown in the evaluation time
652 // So we set some parameters to draw the progress bar a total of totalProgressDraws, i.e. only draw every 1 in 100
653
654 Int_t totalProgressDraws = 100; // total number of times to update the progress bar
655 Int_t drawProgressEvery = 1; // draw every nth event such that we have a total of totalProgressDraws
656 if(nEvents >= totalProgressDraws) drawProgressEvery = nEvents/totalProgressDraws;
657
658 size_t ntargets = Data()->GetEvent(0)->GetNTargets();
659 std::vector<float> output(nEvents*ntargets);
660 auto itr = output.begin();
661 for (Int_t ievt=0; ievt<nEvents; ievt++) {
662
663 Data()->SetCurrentEvent(ievt);
664 std::vector< Float_t > vals = GetRegressionValues();
665 if (vals.size() != ntargets)
666 Log() << kFATAL << "Output regression vector with size " << vals.size() << " is not consistent with target size of "
667 << ntargets << std::endl;
668
669 std::copy(vals.begin(), vals.end(), itr);
670 itr += vals.size();
671
672 // Only draw the progress bar once in a while, doing this every event causes the evaluation to be ridiculously slow
673 if(ievt % drawProgressEvery == 0 || ievt==nEvents-1) timer.DrawProgressBar( ievt );
674 }
675
676 return output;
677}
678
679////////////////////////////////////////////////////////////////////////////////
680/// prepare tree branch with the method's discriminating variable
681
683{
684 Data()->SetCurrentType(type);
685
686 Log() << kINFO <<Form("Dataset[%s] : ",DataInfo().GetName())<< "Create results for " << (type==Types::kTraining?"training":"testing") << Endl;
687
688 ResultsRegression* regRes = (ResultsRegression*)Data()->GetResults(GetMethodName(), type, Types::kRegression);
689
690 Long64_t nEvents = Data()->GetNEvents();
691
692 // use timer
693 Timer timer( nEvents, GetName(), kTRUE );
694
695 Log() << kINFO <<Form("Dataset[%s] : ",DataInfo().GetName()) << "Evaluation of " << GetMethodName() << " on "
696 << (type==Types::kTraining?"training":"testing") << " sample" << Endl;
697
698 regRes->Resize( nEvents );
699
700 std::vector<float> output = GetAllRegressionValues();
701 // assume we have all number of targets for all events
702 Data()->SetCurrentEvent(0);
703 size_t nTargets = GetEvent()->GetNTargets();
704
705 // Log() << kFATAL throws, so after this check we can rely on
706 // output.size() == nTargets * nEvents when forming iterators below
707 if (output.size() != nTargets * size_t(nEvents))
708 Log() << kFATAL << "Output regression vector with size " << output.size() << " is not consistent with target size of "
709 << nTargets << " and number of events " << nEvents << std::endl;
710
711 for (Int_t ievt=0; ievt<nEvents; ievt++) {
712 // Form the iterators per event so that they are never advanced past end()
713 auto valsBegin = output.begin() + size_t(ievt) * nTargets;
714 std::vector<Float_t> vals(valsBegin, valsBegin + nTargets);
715 regRes->SetValue(vals, ievt);
716 }
717
718 Log() << kINFO <<Form("Dataset[%s] : ",DataInfo().GetName())
719 << "Elapsed time for evaluation of " << nEvents << " events: "
720 << timer.GetElapsedTime() << " " << Endl;
721
722 // store time used for testing
724 SetTestTime(timer.ElapsedSeconds());
725
726 TString histNamePrefix(GetTestvarName());
727 histNamePrefix += (type==Types::kTraining?"train":"test");
728 regRes->CreateDeviationHistograms( histNamePrefix );
729}
730////////////////////////////////////////////////////////////////////////////////
731/// Get all multi-class values
733{
734 // use timer for progress bar
735
736 Long64_t nEvents = Data()->GetNEvents();
737 Timer timer( nEvents, GetName(), kTRUE );
738
739 Int_t modulo = Int_t(nEvents/100) + 1;
740 // call first time to get number of classes
741 Data()->SetCurrentEvent(0);
742 std::vector< Float_t > vals = GetMulticlassValues();
743 std::vector<float> output(nEvents * vals.size());
744 auto itr = output.begin();
745 std::copy(vals.begin(), vals.end(), itr);
746 for (Int_t ievt=1; ievt<nEvents; ievt++) {
747 itr += vals.size();
748 Data()->SetCurrentEvent(ievt);
749 vals = GetMulticlassValues();
750
751 std::copy(vals.begin(), vals.end(), itr);
752
753 if (ievt%modulo == 0) timer.DrawProgressBar( ievt );
754 }
755 return output;
756}
757////////////////////////////////////////////////////////////////////////////////
758/// prepare tree branch with the method's discriminating variable
759
761{
762 Data()->SetCurrentType(type);
763
764 Log() << kINFO <<Form("Dataset[%s] : ",DataInfo().GetName())<< "Create results for " << (type==Types::kTraining?"training":"testing") << Endl;
765
766 ResultsMulticlass* resMulticlass = dynamic_cast<ResultsMulticlass*>(Data()->GetResults(GetMethodName(), type, Types::kMulticlass));
767 if (!resMulticlass) Log() << kFATAL<<Form("Dataset[%s] : ",DataInfo().GetName())<< "unable to create pointer in AddMulticlassOutput, exiting."<<Endl;
768
769 Long64_t nEvents = Data()->GetNEvents();
770
771 // use timer
772 Timer timer( nEvents, GetName(), kTRUE );
773
774 Log() << kINFO <<Form("Dataset[%s] : ",DataInfo().GetName())<< "Multiclass evaluation of " << GetMethodName() << " on "
775 << (type==Types::kTraining?"training":"testing") << " sample" << Endl;
776
777 resMulticlass->Resize( nEvents );
778 std::vector<Float_t> output = GetAllMulticlassValues();
779 size_t nClasses = output.size()/nEvents;
780 for (Int_t ievt=0; ievt<nEvents; ievt++) {
781 std::vector< Float_t > vals(output.begin()+ievt*nClasses, output.begin()+(ievt+1)*nClasses);
782 resMulticlass->SetValue( vals, ievt );
783 }
784
785
786 Log() << kINFO <<Form("Dataset[%s] : ",DataInfo().GetName())
787 << "Elapsed time for evaluation of " << nEvents << " events: "
788 << timer.GetElapsedTime() << " " << Endl;
789
790 // store time used for testing
792 SetTestTime(timer.ElapsedSeconds());
793
794 TString histNamePrefix(GetTestvarName());
795 histNamePrefix += (type==Types::kTraining?"_Train":"_Test");
796
797 resMulticlass->CreateMulticlassHistos( histNamePrefix, fNbinsMVAoutput, fNbinsH );
798 resMulticlass->CreateMulticlassPerformanceHistos(histNamePrefix);
799}
800
801////////////////////////////////////////////////////////////////////////////////
802
803void TMVA::MethodBase::NoErrorCalc(Double_t* const err, Double_t* const errUpper) {
804 if (err) *err=-1;
805 if (errUpper) *errUpper=-1;
806}
807
808////////////////////////////////////////////////////////////////////////////////
809
810Double_t TMVA::MethodBase::GetMvaValue( const Event* const ev, Double_t* err, Double_t* errUpper ) {
811 fTmpEvent = ev;
812 Double_t val = GetMvaValue(err, errUpper);
813 fTmpEvent = 0;
814 return val;
815}
816
817////////////////////////////////////////////////////////////////////////////////
818/// uses a pre-set cut on the MVA output (SetSignalReferenceCut and SetSignalReferenceCutOrientation)
819/// for a quick determination if an event would be selected as signal or background
820
822 return GetMvaValue()*GetSignalReferenceCutOrientation() > GetSignalReferenceCut()*GetSignalReferenceCutOrientation() ? kTRUE : kFALSE;
823}
824////////////////////////////////////////////////////////////////////////////////
825/// uses a pre-set cut on the MVA output (SetSignalReferenceCut and SetSignalReferenceCutOrientation)
826/// for a quick determination if an event with this mva output value would be selected as signal or background
827
829 return mvaVal*GetSignalReferenceCutOrientation() > GetSignalReferenceCut()*GetSignalReferenceCutOrientation() ? kTRUE : kFALSE;
830}
831
832////////////////////////////////////////////////////////////////////////////////
833/// prepare tree branch with the method's discriminating variable
834
836{
837 Data()->SetCurrentType(type);
838
839 ResultsClassification* clRes =
840 (ResultsClassification*)Data()->GetResults(GetMethodName(), type, Types::kClassification );
841
842 Long64_t nEvents = Data()->GetNEvents();
843 clRes->Resize( nEvents );
844
845 // use timer
846 Timer timer( nEvents, GetName(), kTRUE );
847
848 Log() << kHEADER << Form("[%s] : ",DataInfo().GetName())
849 << "Evaluation of " << GetMethodName() << " on "
850 << (Data()->GetCurrentType() == Types::kTraining ? "training" : "testing")
851 << " sample (" << nEvents << " events)" << Endl;
852
853 std::vector<Double_t> mvaValues = GetMvaValues(0, nEvents, true);
854
855 Log() << kINFO
856 << "Elapsed time for evaluation of " << nEvents << " events: "
857 << timer.GetElapsedTime() << " " << Endl;
858
859 // store time used for testing
861 SetTestTime(timer.ElapsedSeconds());
862
863 // load mva values and type to results object
864 for (Int_t ievt = 0; ievt < nEvents; ievt++) {
865 // note we do not need the trasformed event to get the signal/background information
866 // by calling Data()->GetEvent instead of this->GetEvent we access the untransformed one
867 auto ev = Data()->GetEvent(ievt);
868 clRes->SetValue(mvaValues[ievt], ievt, DataInfo().IsSignal(ev));
869 }
870}
871
872////////////////////////////////////////////////////////////////////////////////
873/// get all the MVA values for the events of the current Data type
874std::vector<Double_t> TMVA::MethodBase::GetMvaValues(Long64_t firstEvt, Long64_t lastEvt, Bool_t logProgress)
875{
876
877 Long64_t nEvents = Data()->GetNEvents();
878 if (firstEvt > lastEvt || lastEvt > nEvents) lastEvt = nEvents;
879 if (firstEvt < 0) firstEvt = 0;
880 std::vector<Double_t> values(lastEvt-firstEvt);
881 // log in case of looping on all the events
882 nEvents = values.size();
883
884 // use timer
885 Timer timer( nEvents, GetName(), kTRUE );
886
887 if (logProgress)
888 Log() << kHEADER << Form("[%s] : ",DataInfo().GetName())
889 << "Evaluation of " << GetMethodName() << " on "
890 << (Data()->GetCurrentType() == Types::kTraining ? "training" : "testing")
891 << " sample (" << nEvents << " events)" << Endl;
892
893 for (Int_t ievt=firstEvt; ievt<lastEvt; ievt++) {
894 Data()->SetCurrentEvent(ievt);
895 values[ievt] = GetMvaValue();
896
897 // print progress
898 if (logProgress) {
899 Int_t modulo = Int_t(nEvents/100);
900 if (modulo <= 0 ) modulo = 1;
901 if (ievt%modulo == 0) timer.DrawProgressBar( ievt );
902 }
903 }
904 if (logProgress) {
905 Log() << kINFO //<<Form("Dataset[%s] : ",DataInfo().GetName())
906 << "Elapsed time for evaluation of " << nEvents << " events: "
907 << timer.GetElapsedTime() << " " << Endl;
908 }
909
910 return values;
911}
912
913////////////////////////////////////////////////////////////////////////////////
914/// get all the MVA values for the events of the given Data type
915// (this is used by Method Category and it does not need to be re-implemented by derived classes )
916std::vector<Double_t> TMVA::MethodBase::GetDataMvaValues(DataSet * data, Long64_t firstEvt, Long64_t lastEvt, Bool_t logProgress)
917{
918 fTmpData = data;
919 auto result = GetMvaValues(firstEvt, lastEvt, logProgress);
920 fTmpData = nullptr;
921 return result;
922}
923
924////////////////////////////////////////////////////////////////////////////////
925/// prepare tree branch with the method's discriminating variable
926
928{
929 Data()->SetCurrentType(type);
930
931 ResultsClassification* mvaProb =
932 (ResultsClassification*)Data()->GetResults(TString("prob_")+GetMethodName(), type, Types::kClassification );
933
934 Long64_t nEvents = Data()->GetNEvents();
935
936 // use timer
937 Timer timer( nEvents, GetName(), kTRUE );
938
939 Log() << kINFO <<Form("Dataset[%s] : ",DataInfo().GetName()) << "Evaluation of " << GetMethodName() << " on "
940 << (type==Types::kTraining?"training":"testing") << " sample" << Endl;
941
942 mvaProb->Resize( nEvents );
943 Int_t modulo = Int_t(nEvents/100);
944 if (modulo <= 0 ) modulo = 1;
945 for (Int_t ievt=0; ievt<nEvents; ievt++) {
946
947 Data()->SetCurrentEvent(ievt);
948 Float_t proba = ((Float_t)GetProba( GetMvaValue(), 0.5 ));
949 if (proba < 0) break;
950 mvaProb->SetValue( proba, ievt, DataInfo().IsSignal( Data()->GetEvent()) );
951
952 // print progress
953 if (ievt%modulo == 0) timer.DrawProgressBar( ievt );
954 }
955
956 Log() << kDEBUG <<Form("Dataset[%s] : ",DataInfo().GetName())
957 << "Elapsed time for evaluation of " << nEvents << " events: "
958 << timer.GetElapsedTime() << " " << Endl;
959}
960
961////////////////////////////////////////////////////////////////////////////////
962/// calculate <sum-of-deviation-squared> of regression output versus "true" value from test sample
963///
964/// - bias = average deviation
965/// - dev = average absolute deviation
966/// - rms = rms of deviation
967
969 Double_t& dev, Double_t& devT,
970 Double_t& rms, Double_t& rmsT,
971 Double_t& mInf, Double_t& mInfT,
972 Double_t& corr,
974{
975 Types::ETreeType savedType = Data()->GetCurrentType();
976 Data()->SetCurrentType(type);
977
978 bias = 0; biasT = 0; dev = 0; devT = 0; rms = 0; rmsT = 0;
979 Double_t sumw = 0;
980 Double_t m1 = 0, m2 = 0, s1 = 0, s2 = 0, s12 = 0; // for correlation
981 const Int_t nevt = GetNEvents();
982 Float_t* rV = new Float_t[nevt];
983 Float_t* tV = new Float_t[nevt];
984 Float_t* wV = new Float_t[nevt];
985 Float_t xmin = 1e30, xmax = -1e30;
986 Log() << kINFO << "Calculate regression for all events" << Endl;
987 Timer timer( nevt, GetName(), kTRUE );
988 Long64_t modulo = Long64_t(nevt / 100) + 1;
989 auto output = GetAllRegressionValues();
990 int ntargets = Data()->GetEvent(0)->GetNTargets();
991 for (Long64_t ievt=0; ievt<nevt; ievt++) {
992 const Event* ev = Data()->GetEvent(ievt); // NOTE: need untransformed event here !
993 Float_t t = ev->GetTarget(0);
994 Float_t w = ev->GetWeight();
995 Float_t r = output[ievt*ntargets];
996 Float_t d = (r-t);
997
998 // find min/max
1001
1002 // store for truncated RMS computation
1003 rV[ievt] = r;
1004 tV[ievt] = t;
1005 wV[ievt] = w;
1006
1007 // compute deviation-squared
1008 sumw += w;
1009 bias += w * d;
1010 dev += w * TMath::Abs(d);
1011 rms += w * d * d;
1012
1013 // compute correlation between target and regression estimate
1014 m1 += t*w; s1 += t*t*w;
1015 m2 += r*w; s2 += r*r*w;
1016 s12 += t*r;
1017 // print progress
1018 if (ievt % modulo == 0)
1019 timer.DrawProgressBar(ievt);
1020 }
1021 timer.DrawProgressBar(nevt - 1);
1022 Log() << kINFO << "Elapsed time for evaluation of " << nevt << " events: "
1023 << timer.GetElapsedTime() << " " << Endl;
1024
1025 // standard quantities
1026 bias /= sumw;
1027 dev /= sumw;
1028 rms /= sumw;
1029 rms = TMath::Sqrt(rms - bias*bias);
1030
1031 // correlation
1032 m1 /= sumw;
1033 m2 /= sumw;
1034 corr = s12/sumw - m1*m2;
1035 corr /= TMath::Sqrt( (s1/sumw - m1*m1) * (s2/sumw - m2*m2) );
1036
1037 // create histogram required for computation of mutual information
1038 TH2F* hist = new TH2F( "hist", "hist", 150, xmin, xmax, 100, xmin, xmax );
1039 TH2F* histT = new TH2F( "histT", "histT", 150, xmin, xmax, 100, xmin, xmax );
1040
1041 // compute truncated RMS and fill histogram
1042 Double_t devMax = bias + 2*rms;
1043 Double_t devMin = bias - 2*rms;
1044 sumw = 0;
1045 for (Long64_t ievt=0; ievt<nevt; ievt++) {
1046 Float_t d = (rV[ievt] - tV[ievt]);
1047 hist->Fill( rV[ievt], tV[ievt], wV[ievt] );
1048 if (d >= devMin && d <= devMax) {
1049 sumw += wV[ievt];
1050 biasT += wV[ievt] * d;
1051 devT += wV[ievt] * TMath::Abs(d);
1052 rmsT += wV[ievt] * d * d;
1053 histT->Fill( rV[ievt], tV[ievt], wV[ievt] );
1054 }
1055 }
1056 biasT /= sumw;
1057 devT /= sumw;
1058 rmsT /= sumw;
1059 rmsT = TMath::Sqrt(rmsT - biasT*biasT);
1060 mInf = gTools().GetMutualInformation( *hist );
1061 mInfT = gTools().GetMutualInformation( *histT );
1062
1063 delete hist;
1064 delete histT;
1065
1066 delete [] rV;
1067 delete [] tV;
1068 delete [] wV;
1069
1070 Data()->SetCurrentType(savedType);
1071}
1072
1073
1074////////////////////////////////////////////////////////////////////////////////
1075/// test multiclass classification
1076
1078{
1079 ResultsMulticlass* resMulticlass = dynamic_cast<ResultsMulticlass*>(Data()->GetResults(GetMethodName(), Types::kTesting, Types::kMulticlass));
1080 if (!resMulticlass) Log() << kFATAL<<Form("Dataset[%s] : ",DataInfo().GetName())<< "unable to create pointer in TestMulticlass, exiting."<<Endl;
1081
1082 // GA evaluation of best cut for sig eff * sig pur. Slow, disabled for now.
1083 // Log() << kINFO <<Form("Dataset[%s] : ",DataInfo().GetName())<< "Determine optimal multiclass cuts for test
1084 // data..." << Endl; for (UInt_t icls = 0; icls<DataInfo().GetNClasses(); ++icls) {
1085 // resMulticlass->GetBestMultiClassCuts(icls);
1086 // }
1087
1088 // Create histograms for use in TMVA GUI
1089 TString histNamePrefix(GetTestvarName());
1090 TString histNamePrefixTest{histNamePrefix + "_Test"};
1091 TString histNamePrefixTrain{histNamePrefix + "_Train"};
1092
1093 resMulticlass->CreateMulticlassHistos(histNamePrefixTest, fNbinsMVAoutput, fNbinsH);
1094 resMulticlass->CreateMulticlassPerformanceHistos(histNamePrefixTest);
1095
1096 resMulticlass->CreateMulticlassHistos(histNamePrefixTrain, fNbinsMVAoutput, fNbinsH);
1097 resMulticlass->CreateMulticlassPerformanceHistos(histNamePrefixTrain);
1098}
1099
1100
1101////////////////////////////////////////////////////////////////////////////////
1102/// initialization
1103
1105{
1106 Data()->SetCurrentType(Types::kTesting);
1107
1108 ResultsClassification* mvaRes = dynamic_cast<ResultsClassification*>
1109 ( Data()->GetResults(GetMethodName(),Types::kTesting, Types::kClassification) );
1110
1111 // sanity checks: tree must exist, and theVar must be in tree
1112 if (0==mvaRes && !(GetMethodTypeName().Contains("Cuts"))) {
1113 Log()<<Form("Dataset[%s] : ",DataInfo().GetName()) << "mvaRes " << mvaRes << " GetMethodTypeName " << GetMethodTypeName()
1114 << " contains " << !(GetMethodTypeName().Contains("Cuts")) << Endl;
1115 Log() << kFATAL<<Form("Dataset[%s] : ",DataInfo().GetName()) << "<TestInit> Test variable " << GetTestvarName()
1116 << " not found in tree" << Endl;
1117 }
1118
1119 // basic statistics operations are made in base class
1120 gTools().ComputeStat( GetEventCollection(Types::kTesting), mvaRes->GetValueVector(),
1121 fMeanS, fMeanB, fRmsS, fRmsB, fXmin, fXmax, fSignalClass );
1122
1123 // choose reasonable histogram ranges, by removing outliers
1124 Double_t nrms = 10;
1125 fXmin = TMath::Max( TMath::Min( fMeanS - nrms*fRmsS, fMeanB - nrms*fRmsB ), fXmin );
1126 fXmax = TMath::Min( TMath::Max( fMeanS + nrms*fRmsS, fMeanB + nrms*fRmsB ), fXmax );
1127
1128 // determine cut orientation
1129 fCutOrientation = (fMeanS > fMeanB) ? kPositive : kNegative;
1130
1131 // fill 2 types of histograms for the various analyses
1132 // this one is for actual plotting
1133
1134 Double_t sxmax = fXmax+0.00001;
1135
1136 // classifier response distributions for training sample
1137 // MVA plots used for graphics representation (signal)
1138 TString TestvarName;
1139 if(IsSilentFile()) {
1140 TestvarName = TString::Format("[%s]%s",DataInfo().GetName(),GetTestvarName().Data());
1141 } else {
1142 TestvarName=GetTestvarName();
1143 }
1144 TH1* mva_s = new TH1D( TestvarName + "_S",TestvarName + "_S", fNbinsMVAoutput, fXmin, sxmax );
1145 TH1* mva_b = new TH1D( TestvarName + "_B",TestvarName + "_B", fNbinsMVAoutput, fXmin, sxmax );
1146 mvaRes->Store(mva_s, "MVA_S");
1147 mvaRes->Store(mva_b, "MVA_B");
1148 mva_s->Sumw2();
1149 mva_b->Sumw2();
1150
1151 TH1* proba_s = 0;
1152 TH1* proba_b = 0;
1153 TH1* rarity_s = 0;
1154 TH1* rarity_b = 0;
1155 if (HasMVAPdfs()) {
1156 // P(MVA) plots used for graphics representation
1157 proba_s = new TH1D( TestvarName + "_Proba_S", TestvarName + "_Proba_S", fNbinsMVAoutput, 0.0, 1.0 );
1158 proba_b = new TH1D( TestvarName + "_Proba_B", TestvarName + "_Proba_B", fNbinsMVAoutput, 0.0, 1.0 );
1159 mvaRes->Store(proba_s, "Prob_S");
1160 mvaRes->Store(proba_b, "Prob_B");
1161 proba_s->Sumw2();
1162 proba_b->Sumw2();
1163
1164 // R(MVA) plots used for graphics representation
1165 rarity_s = new TH1D( TestvarName + "_Rarity_S", TestvarName + "_Rarity_S", fNbinsMVAoutput, 0.0, 1.0 );
1166 rarity_b = new TH1D( TestvarName + "_Rarity_B", TestvarName + "_Rarity_B", fNbinsMVAoutput, 0.0, 1.0 );
1167 mvaRes->Store(rarity_s, "Rar_S");
1168 mvaRes->Store(rarity_b, "Rar_B");
1169 rarity_s->Sumw2();
1170 rarity_b->Sumw2();
1171 }
1172
1173 // MVA plots used for efficiency calculations (large number of bins)
1174 TH1* mva_eff_s = new TH1D( TestvarName + "_S_high", TestvarName + "_S_high", fNbinsH, fXmin, sxmax );
1175 TH1* mva_eff_b = new TH1D( TestvarName + "_B_high", TestvarName + "_B_high", fNbinsH, fXmin, sxmax );
1176 mvaRes->Store(mva_eff_s, "MVA_HIGHBIN_S");
1177 mvaRes->Store(mva_eff_b, "MVA_HIGHBIN_B");
1178 mva_eff_s->Sumw2();
1179 mva_eff_b->Sumw2();
1180
1181 // fill the histograms
1182
1183 ResultsClassification* mvaProb = dynamic_cast<ResultsClassification*>
1184 (Data()->GetResults( TString("prob_")+GetMethodName(), Types::kTesting, Types::kMaxAnalysisType ) );
1185
1186 Log() << kHEADER <<Form("[%s] : ",DataInfo().GetName())<< "Loop over test events and fill histograms with classifier response..." << Endl << Endl;
1187 if (mvaProb) Log() << kINFO << "Also filling probability and rarity histograms (on request)..." << Endl;
1188 //std::vector<Bool_t>* mvaResTypes = mvaRes->GetValueVectorTypes();
1189
1190 //LM: this is needed to avoid crashes in ROOCCURVE
1191 if ( mvaRes->GetSize() != GetNEvents() ) {
1192 Log() << kFATAL << TString::Format("Inconsistent result size %lld with number of events %u ", mvaRes->GetSize() , GetNEvents() ) << Endl;
1193 assert(mvaRes->GetSize() == GetNEvents());
1194 }
1195
1196 for (Long64_t ievt=0; ievt<GetNEvents(); ievt++) {
1197
1198 const Event* ev = GetEvent(ievt);
1199 Float_t v = (*mvaRes)[ievt][0];
1200 Float_t w = ev->GetWeight();
1201
1202 if (DataInfo().IsSignal(ev)) {
1203 //mvaResTypes->push_back(kTRUE);
1204 mva_s ->Fill( v, w );
1205 if (mvaProb) {
1206 proba_s->Fill( (*mvaProb)[ievt][0], w );
1207 rarity_s->Fill( GetRarity( v ), w );
1208 }
1209
1210 mva_eff_s ->Fill( v, w );
1211 }
1212 else {
1213 //mvaResTypes->push_back(kFALSE);
1214 mva_b ->Fill( v, w );
1215 if (mvaProb) {
1216 proba_b->Fill( (*mvaProb)[ievt][0], w );
1217 rarity_b->Fill( GetRarity( v ), w );
1218 }
1219 mva_eff_b ->Fill( v, w );
1220 }
1221 }
1222
1223 // uncomment those (and several others if you want unnormalized output
1224 gTools().NormHist( mva_s );
1225 gTools().NormHist( mva_b );
1226 gTools().NormHist( proba_s );
1227 gTools().NormHist( proba_b );
1228 gTools().NormHist( rarity_s );
1229 gTools().NormHist( rarity_b );
1230 gTools().NormHist( mva_eff_s );
1231 gTools().NormHist( mva_eff_b );
1232
1233 // create PDFs from histograms, using default splines, and no additional smoothing
1234 if (fSplS) { delete fSplS; fSplS = 0; }
1235 if (fSplB) { delete fSplB; fSplB = 0; }
1236 fSplS = new PDF( TString(GetName()) + " PDF Sig", mva_s, PDF::kSpline2 );
1237 fSplB = new PDF( TString(GetName()) + " PDF Bkg", mva_b, PDF::kSpline2 );
1238}
1239
1240////////////////////////////////////////////////////////////////////////////////
1241/// general method used in writing the header of the weight files where
1242/// the used variables, variable transformation type etc. is specified
1243
1244void TMVA::MethodBase::WriteStateToStream( std::ostream& tf ) const
1245{
1246 TString prefix = "";
1247 UserGroup_t * userInfo = gSystem->GetUserInfo();
1248
1249 tf << prefix << "#GEN -*-*-*-*-*-*-*-*-*-*-*- general info -*-*-*-*-*-*-*-*-*-*-*-" << std::endl << prefix << std::endl;
1250 tf << prefix << "Method : " << GetMethodTypeName() << "::" << GetMethodName() << std::endl;
1251 tf.setf(std::ios::left);
1252 tf << prefix << "TMVA Release : " << std::setw(10) << GetTrainingTMVAVersionString() << " ["
1253 << GetTrainingTMVAVersionCode() << "]" << std::endl;
1254 tf << prefix << "ROOT Release : " << std::setw(10) << GetTrainingROOTVersionString() << " ["
1255 << GetTrainingROOTVersionCode() << "]" << std::endl;
1256 tf << prefix << "Creator : " << userInfo->fUser << std::endl;
1257 tf << prefix << "Date : "; TDatime *d = new TDatime; tf << d->AsString() << std::endl; delete d;
1258 tf << prefix << "Host : " << gSystem->GetBuildNode() << std::endl;
1259 tf << prefix << "Dir : " << gSystem->WorkingDirectory() << std::endl;
1260 tf << prefix << "Training events: " << Data()->GetNTrainingEvents() << std::endl;
1261
1262 TString analysisType(((const_cast<TMVA::MethodBase*>(this)->GetAnalysisType()==Types::kRegression) ? "Regression" : "Classification"));
1263
1264 tf << prefix << "Analysis type : " << "[" << ((GetAnalysisType()==Types::kRegression) ? "Regression" : "Classification") << "]" << std::endl;
1265 tf << prefix << std::endl;
1266
1267 delete userInfo;
1268
1269 // First write all options
1270 tf << prefix << std::endl << prefix << "#OPT -*-*-*-*-*-*-*-*-*-*-*-*- options -*-*-*-*-*-*-*-*-*-*-*-*-" << std::endl << prefix << std::endl;
1271 WriteOptionsToStream( tf, prefix );
1272 tf << prefix << std::endl;
1273
1274 // Second write variable info
1275 tf << prefix << std::endl << prefix << "#VAR -*-*-*-*-*-*-*-*-*-*-*-* variables *-*-*-*-*-*-*-*-*-*-*-*-" << std::endl << prefix << std::endl;
1276 WriteVarsToStream( tf, prefix );
1277 tf << prefix << std::endl;
1278}
1279
1280////////////////////////////////////////////////////////////////////////////////
1281/// xml writing
1282
1283void TMVA::MethodBase::AddInfoItem( void* gi, const TString& name, const TString& value) const
1284{
1285 void* it = gTools().AddChild(gi,"Info");
1286 gTools().AddAttr(it,"name", name);
1287 gTools().AddAttr(it,"value", value);
1288}
1289
1290////////////////////////////////////////////////////////////////////////////////
1291
1293 if (analysisType == Types::kRegression) {
1294 AddRegressionOutput( type );
1295 } else if (analysisType == Types::kMulticlass) {
1296 AddMulticlassOutput( type );
1297 } else {
1298 AddClassifierOutput( type );
1299 if (HasMVAPdfs())
1300 AddClassifierOutputProb( type );
1301 }
1302}
1303
1304////////////////////////////////////////////////////////////////////////////////
1305/// general method used in writing the header of the weight files where
1306/// the used variables, variable transformation type etc. is specified
1307
1308void TMVA::MethodBase::WriteStateToXML( void* parent ) const
1309{
1310 if (!parent) return;
1311
1312 UserGroup_t* userInfo = gSystem->GetUserInfo();
1313
1314 void* gi = gTools().AddChild(parent, "GeneralInfo");
1315 AddInfoItem( gi, "TMVA Release", GetTrainingTMVAVersionString() + " [" + gTools().StringFromInt(GetTrainingTMVAVersionCode()) + "]" );
1316 AddInfoItem( gi, "ROOT Release", GetTrainingROOTVersionString() + " [" + gTools().StringFromInt(GetTrainingROOTVersionCode()) + "]");
1317 AddInfoItem( gi, "Creator", userInfo->fUser);
1318 TDatime dt; AddInfoItem( gi, "Date", dt.AsString());
1319 AddInfoItem( gi, "Host", gSystem->GetBuildNode() );
1320 AddInfoItem( gi, "Dir", gSystem->WorkingDirectory());
1321 AddInfoItem( gi, "Training events", gTools().StringFromInt(Data()->GetNTrainingEvents()));
1322 AddInfoItem( gi, "TrainingTime", gTools().StringFromDouble(const_cast<TMVA::MethodBase*>(this)->GetTrainTime()));
1323
1324 Types::EAnalysisType aType = const_cast<TMVA::MethodBase*>(this)->GetAnalysisType();
1325 TString analysisType((aType==Types::kRegression) ? "Regression" :
1326 (aType==Types::kMulticlass ? "Multiclass" : "Classification"));
1327 AddInfoItem( gi, "AnalysisType", analysisType );
1328 delete userInfo;
1329
1330 // write options
1331 AddOptionsXMLTo( parent );
1332
1333 // write variable info
1334 AddVarsXMLTo( parent );
1335
1336 // write spectator info
1337 if (fModelPersistence)
1338 AddSpectatorsXMLTo( parent );
1339
1340 // write class info if in multiclass mode
1341 AddClassesXMLTo(parent);
1342
1343 // write target info if in regression mode
1344 if (DoRegression()) AddTargetsXMLTo(parent);
1345
1346 // write transformations
1347 GetTransformationHandler(false).AddXMLTo( parent );
1348
1349 // write MVA variable distributions
1350 void* pdfs = gTools().AddChild(parent, "MVAPdfs");
1351 if (fMVAPdfS) fMVAPdfS->AddXMLTo(pdfs);
1352 if (fMVAPdfB) fMVAPdfB->AddXMLTo(pdfs);
1353
1354 // write weights
1355 AddWeightsXMLTo( parent );
1356}
1357
1358////////////////////////////////////////////////////////////////////////////////
1359/// write reference MVA distributions (and other information)
1360/// to a ROOT type weight file
1361
1363{
1364 TDirectory::TContext dirCtx{nullptr}; // Don't register histograms to current directory
1365 fMVAPdfS = (TMVA::PDF*)rf.Get( "MVA_PDF_Signal" );
1366 fMVAPdfB = (TMVA::PDF*)rf.Get( "MVA_PDF_Background" );
1367
1368 ReadWeightsFromStream( rf );
1369
1370 SetTestvarName();
1371}
1372
1373////////////////////////////////////////////////////////////////////////////////
1374/// write options and weights to file
1375/// note that each one text file for the main configuration information
1376/// and one ROOT file for ROOT objects are created
1377
1379{
1380 // ---- create the text file
1381 TString tfname( GetWeightFileName() );
1382
1383 // writing xml file
1384 TString xmlfname( tfname ); xmlfname.ReplaceAll( ".txt", ".xml" );
1385 Log() << kINFO //<<Form("Dataset[%s] : ",DataInfo().GetName())
1386 << "Creating xml weight file: "
1387 << gTools().Color("lightblue") << xmlfname << gTools().Color("reset") << Endl;
1388 void* doc = gTools().xmlengine().NewDoc();
1389 void* rootnode = gTools().AddChild(0,"MethodSetup", "", true);
1390 gTools().xmlengine().DocSetRootElement(doc,rootnode);
1391 gTools().AddAttr(rootnode,"Method", GetMethodTypeName() + "::" + GetMethodName());
1392 WriteStateToXML(rootnode);
1393 gTools().xmlengine().SaveDoc(doc,xmlfname);
1394 gTools().xmlengine().FreeDoc(doc);
1395}
1396
1397////////////////////////////////////////////////////////////////////////////////
1398/// Function to write options and weights to file
1399
1401{
1402 // get the filename
1403
1404 TString tfname(GetWeightFileName());
1405
1406 Log() << kINFO //<<Form("Dataset[%s] : ",DataInfo().GetName())
1407 << "Reading weight file: "
1408 << gTools().Color("lightblue") << tfname << gTools().Color("reset") << Endl;
1409
1410 if (tfname.EndsWith(".xml") ) {
1411 void* doc = gTools().xmlengine().ParseFile(tfname,gTools().xmlenginebuffersize()); // the default buffer size in TXMLEngine::ParseFile is 100k. Starting with ROOT 5.29 one can set the buffer size, see: http://savannah.cern.ch/bugs/?78864. This might be necessary for large XML files
1412 if (!doc) {
1413 Log() << kFATAL << "Error parsing XML file " << tfname << Endl;
1414 }
1415 void* rootnode = gTools().xmlengine().DocGetRootElement(doc); // node "MethodSetup"
1416 ReadStateFromXML(rootnode);
1417 gTools().xmlengine().FreeDoc(doc);
1418 }
1419 else {
1420 std::filebuf fb;
1421 fb.open(tfname.Data(),std::ios::in);
1422 if (!fb.is_open()) { // file not found --> Error
1423 Log() << kFATAL <<Form("Dataset[%s] : ",DataInfo().GetName())<< "<ReadStateFromFile> "
1424 << "Unable to open input weight file: " << tfname << Endl;
1425 }
1426 std::istream fin(&fb);
1427 ReadStateFromStream(fin);
1428 fb.close();
1429 }
1430 if (!fTxtWeightsOnly) {
1431 // ---- read the ROOT file
1432 TString rfname( tfname ); rfname.ReplaceAll( ".txt", ".root" );
1433 Log() << kINFO <<Form("Dataset[%s] : ",DataInfo().GetName())<< "Reading root weight file: "
1434 << gTools().Color("lightblue") << rfname << gTools().Color("reset") << Endl;
1435 TFile* rfile = TFile::Open( rfname, "READ" );
1436 ReadStateFromStream( *rfile );
1437 rfile->Close();
1438 }
1439}
1440////////////////////////////////////////////////////////////////////////////////
1441/// for reading from memory
1442
1444 void* doc = gTools().xmlengine().ParseString(xmlstr);
1445 void* rootnode = gTools().xmlengine().DocGetRootElement(doc); // node "MethodSetup"
1446 ReadStateFromXML(rootnode);
1447 gTools().xmlengine().FreeDoc(doc);
1448
1449 return;
1450}
1451
1452////////////////////////////////////////////////////////////////////////////////
1453
1455{
1456
1457 TString fullMethodName;
1458 gTools().ReadAttr( methodNode, "Method", fullMethodName );
1459
1460 fMethodName = fullMethodName(fullMethodName.Index("::")+2,fullMethodName.Length());
1461
1462 // update logger
1463 Log().SetSource( GetName() );
1464 Log() << kDEBUG//<<Form("Dataset[%s] : ",DataInfo().GetName())
1465 << "Read method \"" << GetMethodName() << "\" of type \"" << GetMethodTypeName() << "\"" << Endl;
1466
1467 // after the method name is read, the testvar can be set
1468 SetTestvarName();
1469
1470 TString nodeName("");
1471 void* ch = gTools().GetChild(methodNode);
1472 while (ch!=0) {
1473 nodeName = TString( gTools().GetName(ch) );
1474
1475 if (nodeName=="GeneralInfo") {
1476 // read analysis type
1477
1478 TString name(""),val("");
1479 void* antypeNode = gTools().GetChild(ch);
1480 while (antypeNode) {
1481 gTools().ReadAttr( antypeNode, "name", name );
1482
1483 if (name == "TrainingTime")
1484 gTools().ReadAttr( antypeNode, "value", fTrainTime );
1485
1486 if (name == "AnalysisType") {
1487 gTools().ReadAttr( antypeNode, "value", val );
1488 val.ToLower();
1489 if (val == "regression" ) SetAnalysisType( Types::kRegression );
1490 else if (val == "classification" ) SetAnalysisType( Types::kClassification );
1491 else if (val == "multiclass" ) SetAnalysisType( Types::kMulticlass );
1492 else Log() << kFATAL <<Form("Dataset[%s] : ",DataInfo().GetName())<< "Analysis type " << val << " is not known." << Endl;
1493 }
1494
1495 if (name == "TMVA Release" || name == "TMVA") {
1496 TString s;
1497 gTools().ReadAttr( antypeNode, "value", s);
1498 fTMVATrainingVersion = TString(s(s.Index("[")+1,s.Index("]")-s.Index("[")-1)).Atoi();
1499 Log() << kDEBUG <<Form("[%s] : ",DataInfo().GetName()) << "MVA method was trained with TMVA Version: " << GetTrainingTMVAVersionString() << Endl;
1500 }
1501
1502 if (name == "ROOT Release" || name == "ROOT") {
1503 TString s;
1504 gTools().ReadAttr( antypeNode, "value", s);
1505 fROOTTrainingVersion = TString(s(s.Index("[")+1,s.Index("]")-s.Index("[")-1)).Atoi();
1506 Log() << kDEBUG //<<Form("Dataset[%s] : ",DataInfo().GetName())
1507 << "MVA method was trained with ROOT Version: " << GetTrainingROOTVersionString() << Endl;
1508 }
1509 antypeNode = gTools().GetNextChild(antypeNode);
1510 }
1511 }
1512 else if (nodeName=="Options") {
1513 ReadOptionsFromXML(ch);
1514 ParseOptions();
1515
1516 }
1517 else if (nodeName=="Variables") {
1518 ReadVariablesFromXML(ch);
1519 }
1520 else if (nodeName=="Spectators") {
1521 ReadSpectatorsFromXML(ch);
1522 }
1523 else if (nodeName=="Classes") {
1524 if (DataInfo().GetNClasses()==0) ReadClassesFromXML(ch);
1525 }
1526 else if (nodeName=="Targets") {
1527 if (DataInfo().GetNTargets()==0 && DoRegression()) ReadTargetsFromXML(ch);
1528 }
1529 else if (nodeName=="Transformations") {
1530 GetTransformationHandler().ReadFromXML(ch);
1531 }
1532 else if (nodeName=="MVAPdfs") {
1533 TString pdfname;
1534 if (fMVAPdfS) { delete fMVAPdfS; fMVAPdfS=0; }
1535 if (fMVAPdfB) { delete fMVAPdfB; fMVAPdfB=0; }
1536 void* pdfnode = gTools().GetChild(ch);
1537 if (pdfnode) {
1538 gTools().ReadAttr(pdfnode, "Name", pdfname);
1539 fMVAPdfS = new PDF(pdfname);
1540 fMVAPdfS->ReadXML(pdfnode);
1541 pdfnode = gTools().GetNextChild(pdfnode);
1542 gTools().ReadAttr(pdfnode, "Name", pdfname);
1543 fMVAPdfB = new PDF(pdfname);
1544 fMVAPdfB->ReadXML(pdfnode);
1545 }
1546 }
1547 else if (nodeName=="Weights") {
1548 ReadWeightsFromXML(ch);
1549 }
1550 else {
1551 Log() << kWARNING <<Form("Dataset[%s] : ",DataInfo().GetName())<< "Unparsed XML node: '" << nodeName << "'" << Endl;
1552 }
1553 ch = gTools().GetNextChild(ch);
1554
1555 }
1556
1557 // update transformation handler
1558 if (GetTransformationHandler().GetCallerName() == "") GetTransformationHandler().SetCallerName( GetName() );
1559}
1560
1561////////////////////////////////////////////////////////////////////////////////
1562/// read the header from the weight files of the different MVA methods
1563
1565{
1566 char buf[512];
1567
1568 // when reading from stream, we assume the files are produced with TMVA<=397
1569 SetAnalysisType(Types::kClassification);
1570
1571
1572 // first read the method name
1573 GetLine(fin,buf);
1574 while (!TString(buf).BeginsWith("Method")) GetLine(fin,buf);
1575 TString namestr(buf);
1576
1577 TString methodType = namestr(0,namestr.Index("::"));
1578 methodType = methodType(methodType.Last(' '),methodType.Length());
1579 methodType = methodType.Strip(TString::kLeading);
1580
1581 TString methodName = namestr(namestr.Index("::")+2,namestr.Length());
1582 methodName = methodName.Strip(TString::kLeading);
1583 if (methodName == "") methodName = methodType;
1584 fMethodName = methodName;
1585
1586 Log() << kINFO <<Form("Dataset[%s] : ",DataInfo().GetName())<< "Read method \"" << GetMethodName() << "\" of type \"" << GetMethodTypeName() << "\"" << Endl;
1587
1588 // update logger
1589 Log().SetSource( GetName() );
1590
1591 // now the question is whether to read the variables first or the options (well, of course the order
1592 // of writing them needs to agree)
1593 //
1594 // the option "Decorrelation" is needed to decide if the variables we
1595 // read are decorrelated or not
1596 //
1597 // the variables are needed by some methods (TMLP) to build the NN
1598 // which is done in ProcessOptions so for the time being we first Read and Parse the options then
1599 // we read the variables, and then we process the options
1600
1601 // now read all options
1602 GetLine(fin,buf);
1603 while (!TString(buf).BeginsWith("#OPT")) GetLine(fin,buf);
1604 ReadOptionsFromStream(fin);
1605 ParseOptions();
1606
1607 // Now read variable info
1608 fin.getline(buf,512);
1609 while (!TString(buf).BeginsWith("#VAR")) fin.getline(buf,512);
1610 ReadVarsFromStream(fin);
1611
1612 // now we process the options (of the derived class)
1613 ProcessOptions();
1614
1615 if (IsNormalised()) {
1617 GetTransformationHandler().AddTransformation( new VariableNormalizeTransform(DataInfo()), -1 );
1618 norm->BuildTransformationFromVarInfo( DataInfo().GetVariableInfos() );
1619 }
1620 VariableTransformBase *varTrafo(0), *varTrafo2(0);
1621 if ( fVarTransformString == "None") {
1622 if (fUseDecorr)
1623 varTrafo = GetTransformationHandler().AddTransformation( new VariableDecorrTransform(DataInfo()), -1 );
1624 } else if ( fVarTransformString == "Decorrelate" ) {
1625 varTrafo = GetTransformationHandler().AddTransformation( new VariableDecorrTransform(DataInfo()), -1 );
1626 } else if ( fVarTransformString == "PCA" ) {
1627 varTrafo = GetTransformationHandler().AddTransformation( new VariablePCATransform(DataInfo()), -1 );
1628 } else if ( fVarTransformString == "Uniform" ) {
1629 varTrafo = GetTransformationHandler().AddTransformation( new VariableGaussTransform(DataInfo(),"Uniform"), -1 );
1630 } else if ( fVarTransformString == "Gauss" ) {
1631 varTrafo = GetTransformationHandler().AddTransformation( new VariableGaussTransform(DataInfo()), -1 );
1632 } else if ( fVarTransformString == "GaussDecorr" ) {
1633 varTrafo = GetTransformationHandler().AddTransformation( new VariableGaussTransform(DataInfo()), -1 );
1634 varTrafo2 = GetTransformationHandler().AddTransformation( new VariableDecorrTransform(DataInfo()), -1 );
1635 } else {
1636 Log() << kFATAL <<Form("Dataset[%s] : ",DataInfo().GetName())<< "<ProcessOptions> Variable transform '"
1637 << fVarTransformString << "' unknown." << Endl;
1638 }
1639 // Now read decorrelation matrix if available
1640 if (GetTransformationHandler().GetTransformationList().GetSize() > 0) {
1641 fin.getline(buf,512);
1642 while (!TString(buf).BeginsWith("#MAT")) fin.getline(buf,512);
1643 if (varTrafo) {
1644 TString trafo(fVariableTransformTypeString); trafo.ToLower();
1645 varTrafo->ReadTransformationFromStream(fin, trafo );
1646 }
1647 if (varTrafo2) {
1648 TString trafo(fVariableTransformTypeString); trafo.ToLower();
1649 varTrafo2->ReadTransformationFromStream(fin, trafo );
1650 }
1651 }
1652
1653
1654 if (HasMVAPdfs()) {
1655 // Now read the MVA PDFs
1656 fin.getline(buf,512);
1657 while (!TString(buf).BeginsWith("#MVAPDFS")) fin.getline(buf,512);
1658 if (fMVAPdfS != 0) { delete fMVAPdfS; fMVAPdfS = 0; }
1659 if (fMVAPdfB != 0) { delete fMVAPdfB; fMVAPdfB = 0; }
1660 fMVAPdfS = new PDF(TString(GetName()) + " MVA PDF Sig");
1661 fMVAPdfB = new PDF(TString(GetName()) + " MVA PDF Bkg");
1662 fMVAPdfS->SetReadingVersion( GetTrainingTMVAVersionCode() );
1663 fMVAPdfB->SetReadingVersion( GetTrainingTMVAVersionCode() );
1664
1665 fin >> *fMVAPdfS;
1666 fin >> *fMVAPdfB;
1667 }
1668
1669 // Now read weights
1670 fin.getline(buf,512);
1671 while (!TString(buf).BeginsWith("#WGT")) fin.getline(buf,512);
1672 fin.getline(buf,512);
1673 ReadWeightsFromStream( fin );
1674
1675 // update transformation handler
1676 if (GetTransformationHandler().GetCallerName() == "") GetTransformationHandler().SetCallerName( GetName() );
1677
1678}
1679
1680////////////////////////////////////////////////////////////////////////////////
1681/// write the list of variables (name, min, max) for a given data
1682/// transformation method to the stream
1683
1684void TMVA::MethodBase::WriteVarsToStream( std::ostream& o, const TString& prefix ) const
1685{
1686 o << prefix << "NVar " << DataInfo().GetNVariables() << std::endl;
1687 std::vector<VariableInfo>::const_iterator varIt = DataInfo().GetVariableInfos().begin();
1688 for (; varIt!=DataInfo().GetVariableInfos().end(); ++varIt) { o << prefix; varIt->WriteToStream(o); }
1689 o << prefix << "NSpec " << DataInfo().GetNSpectators() << std::endl;
1690 varIt = DataInfo().GetSpectatorInfos().begin();
1691 for (; varIt!=DataInfo().GetSpectatorInfos().end(); ++varIt) { o << prefix; varIt->WriteToStream(o); }
1692}
1693
1694////////////////////////////////////////////////////////////////////////////////
1695/// Read the variables (name, min, max) for a given data
1696/// transformation method from the stream. In the stream we only
1697/// expect the limits which will be set
1698
1700{
1701 TString dummy;
1702 UInt_t readNVar;
1703 istr >> dummy >> readNVar;
1704
1705 if (readNVar!=DataInfo().GetNVariables()) {
1706 Log() << kFATAL <<Form("Dataset[%s] : ",DataInfo().GetName())<< "You declared "<< DataInfo().GetNVariables() << " variables in the Reader"
1707 << " while there are " << readNVar << " variables declared in the file"
1708 << Endl;
1709 }
1710
1711 // we want to make sure all variables are read in the order they are defined
1712 VariableInfo varInfo;
1713 std::vector<VariableInfo>::iterator varIt = DataInfo().GetVariableInfos().begin();
1714 int varIdx = 0;
1715 for (; varIt!=DataInfo().GetVariableInfos().end(); ++varIt, ++varIdx) {
1716 varInfo.ReadFromStream(istr);
1717 if (varIt->GetExpression() == varInfo.GetExpression()) {
1718 varInfo.SetExternalLink((*varIt).GetExternalLink());
1719 (*varIt) = varInfo;
1720 }
1721 else {
1722 Log() << kINFO <<Form("Dataset[%s] : ",DataInfo().GetName())<< "ERROR in <ReadVarsFromStream>" << Endl;
1723 Log() << kINFO << "The definition (or the order) of the variables found in the input file is" << Endl;
1724 Log() << kINFO << "is not the same as the one declared in the Reader (which is necessary for" << Endl;
1725 Log() << kINFO << "the correct working of the method):" << Endl;
1726 Log() << kINFO << " var #" << varIdx <<" declared in Reader: " << varIt->GetExpression() << Endl;
1727 Log() << kINFO << " var #" << varIdx <<" declared in file : " << varInfo.GetExpression() << Endl;
1728 Log() << kFATAL << "The expression declared to the Reader needs to be checked (name or order are wrong)" << Endl;
1729 }
1730 }
1731}
1732
1733////////////////////////////////////////////////////////////////////////////////
1734/// write variable info to XML
1735
1736void TMVA::MethodBase::AddVarsXMLTo( void* parent ) const
1737{
1738 void* vars = gTools().AddChild(parent, "Variables");
1739 gTools().AddAttr( vars, "NVar", gTools().StringFromInt(DataInfo().GetNVariables()) );
1740
1741 for (UInt_t idx=0; idx<DataInfo().GetVariableInfos().size(); idx++) {
1742 VariableInfo& vi = DataInfo().GetVariableInfos()[idx];
1743 void* var = gTools().AddChild( vars, "Variable" );
1744 gTools().AddAttr( var, "VarIndex", idx );
1745 vi.AddToXML( var );
1746 }
1747}
1748
1749////////////////////////////////////////////////////////////////////////////////
1750/// write spectator info to XML
1751
1753{
1754 void* specs = gTools().AddChild(parent, "Spectators");
1755
1756 UInt_t writeIdx=0;
1757 for (UInt_t idx=0; idx<DataInfo().GetSpectatorInfos().size(); idx++) {
1758
1759 VariableInfo& vi = DataInfo().GetSpectatorInfos()[idx];
1760
1761 // we do not want to write spectators that are category-cuts,
1762 // except if the method is the category method and the spectators belong to it
1763 if (vi.GetVarType()=='C') continue;
1764
1765 void* spec = gTools().AddChild( specs, "Spectator" );
1766 gTools().AddAttr( spec, "SpecIndex", writeIdx++ );
1767 vi.AddToXML( spec );
1768 }
1769 gTools().AddAttr( specs, "NSpec", gTools().StringFromInt(writeIdx) );
1770}
1771
1772////////////////////////////////////////////////////////////////////////////////
1773/// write class info to XML
1774
1775void TMVA::MethodBase::AddClassesXMLTo( void* parent ) const
1776{
1777 UInt_t nClasses=DataInfo().GetNClasses();
1778
1779 void* classes = gTools().AddChild(parent, "Classes");
1780 gTools().AddAttr( classes, "NClass", nClasses );
1781
1782 for (UInt_t iCls=0; iCls<nClasses; ++iCls) {
1783 ClassInfo *classInfo=DataInfo().GetClassInfo (iCls);
1784 TString className =classInfo->GetName();
1785 UInt_t classNumber=classInfo->GetNumber();
1786
1787 void* classNode=gTools().AddChild(classes, "Class");
1788 gTools().AddAttr( classNode, "Name", className );
1789 gTools().AddAttr( classNode, "Index", classNumber );
1790 }
1791}
1792////////////////////////////////////////////////////////////////////////////////
1793/// write target info to XML
1794
1795void TMVA::MethodBase::AddTargetsXMLTo( void* parent ) const
1796{
1797 void* targets = gTools().AddChild(parent, "Targets");
1798 gTools().AddAttr( targets, "NTrgt", gTools().StringFromInt(DataInfo().GetNTargets()) );
1799
1800 for (UInt_t idx=0; idx<DataInfo().GetTargetInfos().size(); idx++) {
1801 VariableInfo& vi = DataInfo().GetTargetInfos()[idx];
1802 void* tar = gTools().AddChild( targets, "Target" );
1803 gTools().AddAttr( tar, "TargetIndex", idx );
1804 vi.AddToXML( tar );
1805 }
1806}
1807
1808////////////////////////////////////////////////////////////////////////////////
1809/// read variable info from XML
1810
1812{
1813 UInt_t readNVar;
1814 gTools().ReadAttr( varnode, "NVar", readNVar);
1815
1816 if (readNVar!=DataInfo().GetNVariables()) {
1817 Log() << kFATAL <<Form("Dataset[%s] : ",DataInfo().GetName())<< "You declared "<< DataInfo().GetNVariables() << " variables in the Reader"
1818 << " while there are " << readNVar << " variables declared in the file"
1819 << Endl;
1820 }
1821
1822 // we want to make sure all variables are read in the order they are defined
1823 VariableInfo readVarInfo, existingVarInfo;
1824 int varIdx = 0;
1825 void* ch = gTools().GetChild(varnode);
1826 while (ch) {
1827 gTools().ReadAttr( ch, "VarIndex", varIdx);
1828 existingVarInfo = DataInfo().GetVariableInfos()[varIdx];
1829 readVarInfo.ReadFromXML(ch);
1830
1831 if (existingVarInfo.GetExpression() == readVarInfo.GetExpression()) {
1832 readVarInfo.SetExternalLink(existingVarInfo.GetExternalLink());
1833 existingVarInfo = readVarInfo;
1834 }
1835 else {
1836 Log() << kINFO <<Form("Dataset[%s] : ",DataInfo().GetName())<< "ERROR in <ReadVariablesFromXML>" << Endl;
1837 Log() << kINFO << "The definition (or the order) of the variables found in the input file is" << Endl;
1838 Log() << kINFO << "not the same as the one declared in the Reader (which is necessary for the" << Endl;
1839 Log() << kINFO << "correct working of the method):" << Endl;
1840 Log() << kINFO << " var #" << varIdx <<" declared in Reader: " << existingVarInfo.GetExpression() << Endl;
1841 Log() << kINFO << " var #" << varIdx <<" declared in file : " << readVarInfo.GetExpression() << Endl;
1842 Log() << kFATAL << "The expression declared to the Reader needs to be checked (name or order are wrong)" << Endl;
1843 }
1844 ch = gTools().GetNextChild(ch);
1845 }
1846}
1847
1848////////////////////////////////////////////////////////////////////////////////
1849/// read spectator info from XML
1850
1852{
1853 UInt_t readNSpec;
1854 gTools().ReadAttr( specnode, "NSpec", readNSpec);
1855
1856 if (readNSpec!=DataInfo().GetNSpectators(kFALSE)) {
1857 Log() << kFATAL<<Form("Dataset[%s] : ",DataInfo().GetName()) << "You declared "<< DataInfo().GetNSpectators(kFALSE) << " spectators in the Reader"
1858 << " while there are " << readNSpec << " spectators declared in the file"
1859 << Endl;
1860 }
1861
1862 // we want to make sure all variables are read in the order they are defined
1863 VariableInfo readSpecInfo, existingSpecInfo;
1864 int specIdx = 0;
1865 void* ch = gTools().GetChild(specnode);
1866 while (ch) {
1867 gTools().ReadAttr( ch, "SpecIndex", specIdx);
1868 existingSpecInfo = DataInfo().GetSpectatorInfos()[specIdx];
1869 readSpecInfo.ReadFromXML(ch);
1870
1871 if (existingSpecInfo.GetExpression() == readSpecInfo.GetExpression()) {
1872 readSpecInfo.SetExternalLink(existingSpecInfo.GetExternalLink());
1873 existingSpecInfo = readSpecInfo;
1874 }
1875 else {
1876 Log() << kINFO <<Form("Dataset[%s] : ",DataInfo().GetName())<< "ERROR in <ReadSpectatorsFromXML>" << Endl;
1877 Log() << kINFO << "The definition (or the order) of the spectators found in the input file is" << Endl;
1878 Log() << kINFO << "not the same as the one declared in the Reader (which is necessary for the" << Endl;
1879 Log() << kINFO << "correct working of the method):" << Endl;
1880 Log() << kINFO << " spec #" << specIdx <<" declared in Reader: " << existingSpecInfo.GetExpression() << Endl;
1881 Log() << kINFO << " spec #" << specIdx <<" declared in file : " << readSpecInfo.GetExpression() << Endl;
1882 Log() << kFATAL << "The expression declared to the Reader needs to be checked (name or order are wrong)" << Endl;
1883 }
1884 ch = gTools().GetNextChild(ch);
1885 }
1886}
1887
1888////////////////////////////////////////////////////////////////////////////////
1889/// read number of classes from XML
1890
1892{
1893 UInt_t readNCls;
1894 // coverity[tainted_data_argument]
1895 gTools().ReadAttr( clsnode, "NClass", readNCls);
1896
1897 TString className="";
1898 UInt_t classIndex=0;
1899 void* ch = gTools().GetChild(clsnode);
1900 if (!ch) {
1901 for (UInt_t icls = 0; icls<readNCls;++icls) {
1902 TString classname = TString::Format("class%i",icls);
1903 DataInfo().AddClass(classname);
1904
1905 }
1906 }
1907 else{
1908 while (ch) {
1909 gTools().ReadAttr( ch, "Index", classIndex);
1910 gTools().ReadAttr( ch, "Name", className );
1911 DataInfo().AddClass(className);
1912
1913 ch = gTools().GetNextChild(ch);
1914 }
1915 }
1916
1917 // retrieve signal and background class index
1918 if (DataInfo().GetClassInfo("Signal") != 0) {
1919 fSignalClass = DataInfo().GetClassInfo("Signal")->GetNumber();
1920 }
1921 else
1922 fSignalClass=0;
1923 if (DataInfo().GetClassInfo("Background") != 0) {
1924 fBackgroundClass = DataInfo().GetClassInfo("Background")->GetNumber();
1925 }
1926 else
1927 fBackgroundClass=1;
1928}
1929
1930////////////////////////////////////////////////////////////////////////////////
1931/// read target info from XML
1932
1934{
1935 UInt_t readNTar;
1936 gTools().ReadAttr( tarnode, "NTrgt", readNTar);
1937
1938 int tarIdx = 0;
1939 TString expression;
1940 void* ch = gTools().GetChild(tarnode);
1941 while (ch) {
1942 gTools().ReadAttr( ch, "TargetIndex", tarIdx);
1943 gTools().ReadAttr( ch, "Expression", expression);
1944 DataInfo().AddTarget(expression,"","",0,0);
1945
1946 ch = gTools().GetNextChild(ch);
1947 }
1948}
1949
1950////////////////////////////////////////////////////////////////////////////////
1951/// returns the ROOT directory where info/histograms etc of the
1952/// corresponding MVA method instance are stored
1953
1955{
1956 if (fBaseDir != 0) return fBaseDir;
1957 Log()<<kDEBUG<<Form("Dataset[%s] : ",DataInfo().GetName())<<" Base Directory for " << GetMethodName() << " not set yet --> check if already there.." <<Endl;
1958
1959 if (IsSilentFile()) {
1960 Log() << kFATAL << Form("Dataset[%s] : ", DataInfo().GetName())
1961 << "MethodBase::BaseDir() - No directory exists when running a Method without output file. Enable the "
1962 "output when creating the factory"
1963 << Endl;
1964 }
1965
1966 TDirectory* methodDir = MethodBaseDir();
1967 if (methodDir==0)
1968 Log() << kFATAL <<Form("Dataset[%s] : ",DataInfo().GetName())<< "MethodBase::BaseDir() - MethodBaseDir() return a NULL pointer!" << Endl;
1969
1970 TString defaultDir = GetMethodName();
1971 TDirectory *sdir = methodDir->GetDirectory(defaultDir.Data());
1972 if(!sdir)
1973 {
1974 Log()<<kDEBUG<<Form("Dataset[%s] : ",DataInfo().GetName())<<" Base Directory for " << GetMethodTypeName() << " does not exist yet--> created it" <<Endl;
1975 sdir = methodDir->mkdir(defaultDir);
1976 sdir->cd();
1977 // write weight file name into target file
1978 if (fModelPersistence) {
1979 TObjString wfilePath( gSystem->WorkingDirectory() );
1980 TObjString wfileName( GetWeightFileName() );
1981 wfilePath.Write( "TrainingPath" );
1982 wfileName.Write( "WeightFileName" );
1983 }
1984 }
1985
1986 Log()<<kDEBUG<<Form("Dataset[%s] : ",DataInfo().GetName())<<" Base Directory for " << GetMethodTypeName() << " existed, return it.." <<Endl;
1987 return sdir;
1988}
1989
1990////////////////////////////////////////////////////////////////////////////////
1991/// returns the ROOT directory where all instances of the
1992/// corresponding MVA method are stored
1993
1995{
1996 if (fMethodBaseDir != 0) {
1997 return fMethodBaseDir;
1998 }
1999
2000 const char *datasetName = DataInfo().GetName();
2001
2002 Log() << kDEBUG << Form("Dataset[%s] : ", datasetName) << " Base Directory for " << GetMethodTypeName()
2003 << " not set yet --> check if already there.." << Endl;
2004
2005 TDirectory *factoryBaseDir = GetFile();
2006 if (!factoryBaseDir) return nullptr;
2007 fMethodBaseDir = factoryBaseDir->GetDirectory(datasetName);
2008 if (!fMethodBaseDir) {
2009 fMethodBaseDir = factoryBaseDir->mkdir(datasetName, TString::Format("Base directory for dataset %s", datasetName).Data());
2010 if (!fMethodBaseDir) {
2011 Log() << kFATAL << "Can not create dir " << datasetName;
2012 }
2013 }
2014 TString methodTypeDir = TString::Format("Method_%s", GetMethodTypeName().Data());
2015 fMethodBaseDir = fMethodBaseDir->GetDirectory(methodTypeDir.Data());
2016
2017 if (!fMethodBaseDir) {
2018 TDirectory *datasetDir = factoryBaseDir->GetDirectory(datasetName);
2019 TString methodTypeDirHelpStr = TString::Format("Directory for all %s methods", GetMethodTypeName().Data());
2020 fMethodBaseDir = datasetDir->mkdir(methodTypeDir.Data(), methodTypeDirHelpStr);
2021 Log() << kDEBUG << Form("Dataset[%s] : ", datasetName) << " Base Directory for " << GetMethodName()
2022 << " does not exist yet--> created it" << Endl;
2023 }
2024
2025 Log() << kDEBUG << Form("Dataset[%s] : ", datasetName)
2026 << "Return from MethodBaseDir() after creating base directory " << Endl;
2027 return fMethodBaseDir;
2028}
2029
2030////////////////////////////////////////////////////////////////////////////////
2031/// set directory of weight file
2032
2034{
2035 fFileDir = fileDir;
2036 gSystem->mkdir( fFileDir, kTRUE );
2037}
2038
2039////////////////////////////////////////////////////////////////////////////////
2040/// set the weight file name (depreciated)
2041
2043{
2044 fWeightFile = theWeightFile;
2045}
2046
2047////////////////////////////////////////////////////////////////////////////////
2048/// retrieve weight file name
2049
2051{
2052 if (fWeightFile!="") return fWeightFile;
2053
2054 // the default consists of
2055 // directory/jobname_methodname_suffix.extension.{root/txt}
2056 TString suffix = "";
2057 TString wFileDir(GetWeightFileDir());
2058 TString wFileName = GetJobName() + "_" + GetMethodName() +
2059 suffix + "." + gConfig().GetIONames().fWeightFileExtension + ".xml";
2060 if (wFileDir.IsNull() ) return wFileName;
2061 // add weight file directory of it is not null
2062 return ( wFileDir + (wFileDir[wFileDir.Length()-1]=='/' ? "" : "/")
2063 + wFileName );
2064}
2065////////////////////////////////////////////////////////////////////////////////
2066/// writes all MVA evaluation histograms to file
2067
2069{
2070 BaseDir()->cd();
2071
2072
2073 // write MVA PDFs to file - if exist
2074 if (0 != fMVAPdfS) {
2075 fMVAPdfS->GetOriginalHist()->Write();
2076 fMVAPdfS->GetSmoothedHist()->Write();
2077 fMVAPdfS->GetPDFHist()->Write();
2078 }
2079 if (0 != fMVAPdfB) {
2080 fMVAPdfB->GetOriginalHist()->Write();
2081 fMVAPdfB->GetSmoothedHist()->Write();
2082 fMVAPdfB->GetPDFHist()->Write();
2083 }
2084
2085 // write result-histograms
2086 Results* results = Data()->GetResults( GetMethodName(), treetype, Types::kMaxAnalysisType );
2087 if (!results)
2088 Log() << kFATAL <<Form("Dataset[%s] : ",DataInfo().GetName())<< "<WriteEvaluationHistosToFile> Unknown result: "
2089 << GetMethodName() << (treetype==Types::kTraining?"/kTraining":"/kTesting")
2090 << "/kMaxAnalysisType" << Endl;
2091 results->GetStorage()->Write();
2092 if (treetype==Types::kTesting) {
2093 // skipping plotting of variables if too many (default is 200)
2094 if ((int) DataInfo().GetNVariables()< gConfig().GetVariablePlotting().fMaxNumOfAllowedVariables)
2095 GetTransformationHandler().PlotVariables (GetEventCollection( Types::kTesting ), BaseDir() );
2096 else
2097 Log() << kINFO << TString::Format("Dataset[%s] : ",DataInfo().GetName())
2098 << " variable plots are not produces ! The number of variables is " << DataInfo().GetNVariables()
2099 << " , it is larger than " << gConfig().GetVariablePlotting().fMaxNumOfAllowedVariables << Endl;
2100 }
2101}
2102
2103////////////////////////////////////////////////////////////////////////////////
2104/// write special monitoring histograms to file
2105/// dummy implementation here -----------------
2106
2110
2111////////////////////////////////////////////////////////////////////////////////
2112/// reads one line from the input stream
2113/// checks for certain keywords and interprets
2114/// the line if keywords are found
2115
2116Bool_t TMVA::MethodBase::GetLine(std::istream& fin, char* buf )
2117{
2118 fin.getline(buf,512);
2119 TString line(buf);
2120 if (line.BeginsWith("TMVA Release")) {
2121 Ssiz_t start = line.First('[')+1;
2122 Ssiz_t length = line.Index("]",start)-start;
2123 TString code = line(start,length);
2124 std::stringstream s(code.Data());
2125 s >> fTMVATrainingVersion;
2126 Log() << kINFO <<Form("Dataset[%s] : ",DataInfo().GetName())<< "MVA method was trained with TMVA Version: " << GetTrainingTMVAVersionString() << Endl;
2127 }
2128 if (line.BeginsWith("ROOT Release")) {
2129 Ssiz_t start = line.First('[')+1;
2130 Ssiz_t length = line.Index("]",start)-start;
2131 TString code = line(start,length);
2132 std::stringstream s(code.Data());
2133 s >> fROOTTrainingVersion;
2134 Log() << kINFO <<Form("Dataset[%s] : ",DataInfo().GetName())<< "MVA method was trained with ROOT Version: " << GetTrainingROOTVersionString() << Endl;
2135 }
2136 if (line.BeginsWith("Analysis type")) {
2137 Ssiz_t start = line.First('[')+1;
2138 Ssiz_t length = line.Index("]",start)-start;
2139 TString code = line(start,length);
2140 std::stringstream s(code.Data());
2141 std::string analysisType;
2142 s >> analysisType;
2143 if (analysisType == "regression" || analysisType == "Regression") SetAnalysisType( Types::kRegression );
2144 else if (analysisType == "classification" || analysisType == "Classification") SetAnalysisType( Types::kClassification );
2145 else if (analysisType == "multiclass" || analysisType == "Multiclass") SetAnalysisType( Types::kMulticlass );
2146 else Log() << kFATAL << "Analysis type " << analysisType << " from weight-file not known!" << std::endl;
2147
2148 Log() << kINFO <<Form("Dataset[%s] : ",DataInfo().GetName())<< "Method was trained for "
2149 << (GetAnalysisType() == Types::kRegression ? "Regression" :
2150 (GetAnalysisType() == Types::kMulticlass ? "Multiclass" : "Classification")) << Endl;
2151 }
2152
2153 return true;
2154}
2155
2156////////////////////////////////////////////////////////////////////////////////
2157/// Create PDFs of the MVA output variables
2158
2160{
2161 Data()->SetCurrentType(Types::kTraining);
2162
2163 // the PDF's are stored as results ONLY if the corresponding "results" are booked,
2164 // otherwise they will be only used 'online'
2165 ResultsClassification * mvaRes = dynamic_cast<ResultsClassification*>
2166 ( Data()->GetResults(GetMethodName(), Types::kTraining, Types::kClassification) );
2167
2168 if (mvaRes==0 || mvaRes->GetSize()==0) {
2169 Log() << kERROR<<Form("Dataset[%s] : ",DataInfo().GetName())<< "<CreateMVAPdfs> No result of classifier testing available" << Endl;
2170 }
2171
2172 Double_t minVal = *std::min_element(mvaRes->GetValueVector()->begin(),mvaRes->GetValueVector()->end());
2173 Double_t maxVal = *std::max_element(mvaRes->GetValueVector()->begin(),mvaRes->GetValueVector()->end());
2174
2175 // create histograms that serve as basis to create the MVA Pdfs
2176 TH1* histMVAPdfS = new TH1D( GetMethodTypeName() + "_tr_S", GetMethodTypeName() + "_tr_S",
2177 fMVAPdfS->GetHistNBins( mvaRes->GetSize() ), minVal, maxVal );
2178 TH1* histMVAPdfB = new TH1D( GetMethodTypeName() + "_tr_B", GetMethodTypeName() + "_tr_B",
2179 fMVAPdfB->GetHistNBins( mvaRes->GetSize() ), minVal, maxVal );
2180
2181
2182 // compute sum of weights properly
2183 histMVAPdfS->Sumw2();
2184 histMVAPdfB->Sumw2();
2185
2186 // fill histograms
2187 for (UInt_t ievt=0; ievt<mvaRes->GetSize(); ievt++) {
2188 Double_t theVal = mvaRes->GetValueVector()->at(ievt);
2189 Double_t theWeight = Data()->GetEvent(ievt)->GetWeight();
2190
2191 if (DataInfo().IsSignal(Data()->GetEvent(ievt))) histMVAPdfS->Fill( theVal, theWeight );
2192 else histMVAPdfB->Fill( theVal, theWeight );
2193 }
2194
2195 gTools().NormHist( histMVAPdfS );
2196 gTools().NormHist( histMVAPdfB );
2197
2198 // momentary hack for ROOT problem
2199 if(!IsSilentFile())
2200 {
2201 histMVAPdfS->Write();
2202 histMVAPdfB->Write();
2203 }
2204 // create PDFs
2205 fMVAPdfS->BuildPDF ( histMVAPdfS );
2206 fMVAPdfB->BuildPDF ( histMVAPdfB );
2207 fMVAPdfS->ValidatePDF( histMVAPdfS );
2208 fMVAPdfB->ValidatePDF( histMVAPdfB );
2209
2210 if (DataInfo().GetNClasses() == 2) { // TODO: this is an ugly hack.. adapt this to new framework
2211 Log() << kINFO<<Form("Dataset[%s] : ",DataInfo().GetName())
2212 << TString::Format( "<CreateMVAPdfs> Separation from histogram (PDF): %1.3f (%1.3f)",
2213 GetSeparation( histMVAPdfS, histMVAPdfB ), GetSeparation( fMVAPdfS, fMVAPdfB ) )
2214 << Endl;
2215 }
2216
2217 delete histMVAPdfS;
2218 delete histMVAPdfB;
2219}
2220
2222 // the simple one, automatically calculates the mvaVal and uses the
2223 // SAME sig/bkg ratio as given in the training sample (typically 50/50
2224 // .. (NormMode=EqualNumEvents) but can be different)
2225 if (!fMVAPdfS || !fMVAPdfB) {
2226 Log() << kINFO<<Form("Dataset[%s] : ",DataInfo().GetName()) << "<GetProba> MVA PDFs for Signal and Background don't exist yet, we'll create them on demand" << Endl;
2227 CreateMVAPdfs();
2228 }
2229 Double_t sigFraction = DataInfo().GetTrainingSumSignalWeights() / (DataInfo().GetTrainingSumSignalWeights() + DataInfo().GetTrainingSumBackgrWeights() );
2230 Double_t mvaVal = GetMvaValue(ev);
2231
2232 return GetProba(mvaVal,sigFraction);
2233
2234}
2235////////////////////////////////////////////////////////////////////////////////
2236/// compute likelihood ratio
2237
2239{
2240 if (!fMVAPdfS || !fMVAPdfB) {
2241 Log() << kWARNING <<Form("Dataset[%s] : ",DataInfo().GetName())<< "<GetProba> MVA PDFs for Signal and Background don't exist" << Endl;
2242 return -1.0;
2243 }
2244 Double_t p_s = fMVAPdfS->GetVal( mvaVal );
2245 Double_t p_b = fMVAPdfB->GetVal( mvaVal );
2246
2247 Double_t denom = p_s*ap_sig + p_b*(1 - ap_sig);
2248
2249 return (denom > 0) ? (p_s*ap_sig) / denom : -1;
2250}
2251
2252////////////////////////////////////////////////////////////////////////////////
2253/// compute rarity:
2254/// \f[
2255/// R(x) = \int_{[-\infty..x]} { PDF(x') dx' }
2256/// \f]
2257/// where PDF(x) is the PDF of the classifier's signal or background distribution
2258
2260{
2261 if ((reftype == Types::kSignal && !fMVAPdfS) || (reftype == Types::kBackground && !fMVAPdfB)) {
2262 Log() << kWARNING <<Form("Dataset[%s] : ",DataInfo().GetName())<< "<GetRarity> Required MVA PDF for Signal or Background does not exist: "
2263 << "select option \"CreateMVAPdfs\"" << Endl;
2264 return 0.0;
2265 }
2266
2267 PDF* thePdf = ((reftype == Types::kSignal) ? fMVAPdfS : fMVAPdfB);
2268
2269 return thePdf->GetIntegral( thePdf->GetXmin(), mvaVal );
2270}
2271
2272////////////////////////////////////////////////////////////////////////////////
2273/// fill background efficiency (resp. rejection) versus signal efficiency plots
2274/// returns signal efficiency at background efficiency indicated in theString
2275
2277{
2278 Data()->SetCurrentType(type);
2279 Results* results = Data()->GetResults( GetMethodName(), type, Types::kClassification );
2280 std::vector<Float_t>* mvaRes = dynamic_cast<ResultsClassification*>(results)->GetValueVector();
2281
2282 // parse input string for required background efficiency
2283 TList* list = gTools().ParseFormatLine( theString );
2284
2285 // sanity check
2286 Bool_t computeArea = kFALSE;
2287 if (!list || list->GetSize() < 2) computeArea = kTRUE; // the area is computed
2288 else if (list->GetSize() > 2) {
2289 Log() << kFATAL <<Form("Dataset[%s] : ",DataInfo().GetName())<< "<GetEfficiency> Wrong number of arguments"
2290 << " in string: " << theString
2291 << " | required format, e.g., Efficiency:0.05, or empty string" << Endl;
2292 delete list;
2293 return -1;
2294 }
2295
2296 // sanity check
2297 if ( results->GetHist("MVA_S")->GetNbinsX() != results->GetHist("MVA_B")->GetNbinsX() ||
2298 results->GetHist("MVA_HIGHBIN_S")->GetNbinsX() != results->GetHist("MVA_HIGHBIN_B")->GetNbinsX() ) {
2299 Log() << kFATAL <<Form("Dataset[%s] : ",DataInfo().GetName())<< "<GetEfficiency> Binning mismatch between signal and background histos" << Endl;
2300 delete list;
2301 return -1.0;
2302 }
2303
2304 // create histograms
2305
2306 // first, get efficiency histograms for signal and background
2307 TH1 * effhist = results->GetHist("MVA_HIGHBIN_S");
2308 Double_t xmin = effhist->GetXaxis()->GetXmin();
2309 Double_t xmax = effhist->GetXaxis()->GetXmax();
2310
2311 TTHREAD_TLS(Double_t) nevtS;
2312
2313 // first round ? --> create histograms
2314 if (results->DoesExist("MVA_EFF_S")==0) {
2315
2316 // for efficiency plot
2317 TH1* eff_s = new TH1D( GetTestvarName() + "_effS", GetTestvarName() + " (signal)", fNbinsH, xmin, xmax );
2318 TH1* eff_b = new TH1D( GetTestvarName() + "_effB", GetTestvarName() + " (background)", fNbinsH, xmin, xmax );
2319 results->Store(eff_s, "MVA_EFF_S");
2320 results->Store(eff_b, "MVA_EFF_B");
2321
2322 // sign if cut
2323 Int_t sign = (fCutOrientation == kPositive) ? +1 : -1;
2324
2325 // this method is unbinned
2326 nevtS = 0;
2327 for (UInt_t ievt=0; ievt<Data()->GetNEvents(); ievt++) {
2328
2329 // read the tree
2330 Bool_t isSignal = DataInfo().IsSignal(GetEvent(ievt));
2331 Float_t theWeight = GetEvent(ievt)->GetWeight();
2332 Float_t theVal = (*mvaRes)[ievt];
2333
2334 // select histogram depending on if sig or bgd
2335 TH1* theHist = isSignal ? eff_s : eff_b;
2336
2337 // count signal and background events in tree
2338 if (isSignal) nevtS+=theWeight;
2339
2340 TAxis* axis = theHist->GetXaxis();
2341 Int_t maxbin = Int_t((theVal - axis->GetXmin())/(axis->GetXmax() - axis->GetXmin())*fNbinsH) + 1;
2342 if (sign > 0 && maxbin > fNbinsH) continue; // can happen... event doesn't count
2343 if (sign < 0 && maxbin < 1 ) continue; // can happen... event doesn't count
2344 if (sign > 0 && maxbin < 1 ) maxbin = 1;
2345 if (sign < 0 && maxbin > fNbinsH) maxbin = fNbinsH;
2346
2347 if (sign > 0)
2348 for (Int_t ibin=1; ibin<=maxbin; ibin++) theHist->AddBinContent( ibin , theWeight);
2349 else if (sign < 0)
2350 for (Int_t ibin=maxbin+1; ibin<=fNbinsH; ibin++) theHist->AddBinContent( ibin , theWeight );
2351 else
2352 Log() << kFATAL <<Form("Dataset[%s] : ",DataInfo().GetName())<< "<GetEfficiency> Mismatch in sign" << Endl;
2353 }
2354
2355 // renormalise maximum to <=1
2356 // eff_s->Scale( 1.0/TMath::Max(1.,eff_s->GetMaximum()) );
2357 // eff_b->Scale( 1.0/TMath::Max(1.,eff_b->GetMaximum()) );
2358
2359 eff_s->Scale( 1.0/TMath::Max(std::numeric_limits<double>::epsilon(),eff_s->GetMaximum()) );
2360 eff_b->Scale( 1.0/TMath::Max(std::numeric_limits<double>::epsilon(),eff_b->GetMaximum()) );
2361
2362 // background efficiency versus signal efficiency
2363 TH1* eff_BvsS = new TH1D( GetTestvarName() + "_effBvsS", GetTestvarName() + "", fNbins, 0, 1 );
2364 results->Store(eff_BvsS, "MVA_EFF_BvsS");
2365 eff_BvsS->SetXTitle( "Signal eff" );
2366 eff_BvsS->SetYTitle( "Backgr eff" );
2367
2368 // background rejection (=1-eff.) versus signal efficiency
2369 TH1* rej_BvsS = new TH1D( GetTestvarName() + "_rejBvsS", GetTestvarName() + "", fNbins, 0, 1 );
2370 results->Store(rej_BvsS);
2371 rej_BvsS->SetXTitle( "Signal eff" );
2372 rej_BvsS->SetYTitle( "Backgr rejection (1-eff)" );
2373
2374 // inverse background eff (1/eff.) versus signal efficiency
2375 TH1* inveff_BvsS = new TH1D( GetTestvarName() + "_invBeffvsSeff",
2376 GetTestvarName(), fNbins, 0, 1 );
2377 results->Store(inveff_BvsS);
2378 inveff_BvsS->SetXTitle( "Signal eff" );
2379 inveff_BvsS->SetYTitle( "Inverse backgr. eff (1/eff)" );
2380
2381 // use root finder
2382 // spline background efficiency plot
2383 // note that there is a bin shift when going from a TH1D object to a TGraph :-(
2385 fSplRefS = new TSpline1( "spline2_signal", new TGraph( eff_s ) );
2386 fSplRefB = new TSpline1( "spline2_background", new TGraph( eff_b ) );
2387
2388 // verify spline sanity
2389 gTools().CheckSplines( eff_s, fSplRefS );
2390 gTools().CheckSplines( eff_b, fSplRefB );
2391 }
2392
2393 // make the background-vs-signal efficiency plot
2394
2395 // create root finder
2396 RootFinder rootFinder( this, fXmin, fXmax );
2397
2398 Double_t effB = 0;
2399 fEffS = eff_s; // to be set for the root finder
2400 for (Int_t bini=1; bini<=fNbins; bini++) {
2401
2402 // find cut value corresponding to a given signal efficiency
2403 Double_t effS = eff_BvsS->GetBinCenter( bini );
2404 Double_t cut = rootFinder.Root( effS );
2405
2406 // retrieve background efficiency for given cut
2407 if (Use_Splines_for_Eff_) effB = fSplRefB->Eval( cut );
2408 else effB = eff_b->GetBinContent( eff_b->FindBin( cut ) );
2409
2410 // and fill histograms
2411 eff_BvsS->SetBinContent( bini, effB );
2412 rej_BvsS->SetBinContent( bini, 1.0-effB );
2413 if (effB>std::numeric_limits<double>::epsilon())
2414 inveff_BvsS->SetBinContent( bini, 1.0/effB );
2415 }
2416
2417 // create splines for histogram
2418 fSpleffBvsS = new TSpline1( "effBvsS", new TGraph( eff_BvsS ) );
2419
2420 // search for overlap point where, when cutting on it,
2421 // one would obtain: eff_S = rej_B = 1 - eff_B
2422 Double_t effS = 0., rejB, effS_ = 0., rejB_ = 0.;
2423 Int_t nbins_ = 5000;
2424 for (Int_t bini=1; bini<=nbins_; bini++) {
2425
2426 // get corresponding signal and background efficiencies
2427 effS = (bini - 0.5)/Float_t(nbins_);
2428 rejB = 1.0 - fSpleffBvsS->Eval( effS );
2429
2430 // find signal efficiency that corresponds to required background efficiency
2431 if ((effS - rejB)*(effS_ - rejB_) < 0) break;
2432 effS_ = effS;
2433 rejB_ = rejB;
2434 }
2435
2436 // find cut that corresponds to signal efficiency and update signal-like criterion
2437 Double_t cut = rootFinder.Root( 0.5*(effS + effS_) );
2438 SetSignalReferenceCut( cut );
2439 fEffS = 0;
2440 }
2441
2442 // must exist...
2443 if (0 == fSpleffBvsS) {
2444 delete list;
2445 return 0.0;
2446 }
2447
2448 // now find signal efficiency that corresponds to required background efficiency
2449 Double_t effS = 0, effB = 0, effS_ = 0, effB_ = 0;
2450 Int_t nbins_ = 1000;
2451
2452 if (computeArea) {
2453
2454 // compute area of rej-vs-eff plot
2455 Double_t integral = 0;
2456 for (Int_t bini=1; bini<=nbins_; bini++) {
2457
2458 // get corresponding signal and background efficiencies
2459 effS = (bini - 0.5)/Float_t(nbins_);
2460 effB = fSpleffBvsS->Eval( effS );
2461 integral += (1.0 - effB);
2462 }
2463 integral /= nbins_;
2464
2465 delete list;
2466 return integral;
2467 }
2468 else {
2469
2470 // that will be the value of the efficiency retured (does not affect
2471 // the efficiency-vs-bkg plot which is done anyway.
2472 Float_t effBref = atof( ((TObjString*)list->At(1))->GetString() );
2473
2474 // find precise efficiency value
2475 for (Int_t bini=1; bini<=nbins_; bini++) {
2476
2477 // get corresponding signal and background efficiencies
2478 effS = (bini - 0.5)/Float_t(nbins_);
2479 effB = fSpleffBvsS->Eval( effS );
2480
2481 // find signal efficiency that corresponds to required background efficiency
2482 if ((effB - effBref)*(effB_ - effBref) <= 0) break;
2483 effS_ = effS;
2484 effB_ = effB;
2485 }
2486
2487 // take mean between bin above and bin below
2488 effS = 0.5*(effS + effS_);
2489
2490 effSerr = 0;
2491 if (nevtS > 0) effSerr = TMath::Sqrt( effS*(1.0 - effS)/nevtS );
2492
2493 delete list;
2494 return effS;
2495 }
2496
2497 return -1;
2498}
2499
2500////////////////////////////////////////////////////////////////////////////////
2501
2503{
2504 Data()->SetCurrentType(Types::kTraining);
2505
2506 Results* results = Data()->GetResults(GetMethodName(), Types::kTesting, Types::kNoAnalysisType);
2507
2508 // fill background efficiency (resp. rejection) versus signal efficiency plots
2509 // returns signal efficiency at background efficiency indicated in theString
2510
2511 // parse input string for required background efficiency
2512 TList* list = gTools().ParseFormatLine( theString );
2513 // sanity check
2514
2515 if (list->GetSize() != 2) {
2516 Log() << kFATAL <<Form("Dataset[%s] : ",DataInfo().GetName())<< "<GetTrainingEfficiency> Wrong number of arguments"
2517 << " in string: " << theString
2518 << " | required format, e.g., Efficiency:0.05" << Endl;
2519 delete list;
2520 return -1;
2521 }
2522 // that will be the value of the efficiency retured (does not affect
2523 // the efficiency-vs-bkg plot which is done anyway.
2524 Float_t effBref = atof( ((TObjString*)list->At(1))->GetString() );
2525
2526 delete list;
2527
2528 // sanity check
2529 if (results->GetHist("MVA_S")->GetNbinsX() != results->GetHist("MVA_B")->GetNbinsX() ||
2530 results->GetHist("MVA_HIGHBIN_S")->GetNbinsX() != results->GetHist("MVA_HIGHBIN_B")->GetNbinsX() ) {
2531 Log() << kFATAL <<Form("Dataset[%s] : ",DataInfo().GetName())<< "<GetTrainingEfficiency> Binning mismatch between signal and background histos"
2532 << Endl;
2533 return -1.0;
2534 }
2535
2536 // create histogram
2537
2538 // first, get efficiency histograms for signal and background
2539 TH1 * effhist = results->GetHist("MVA_HIGHBIN_S");
2540 Double_t xmin = effhist->GetXaxis()->GetXmin();
2541 Double_t xmax = effhist->GetXaxis()->GetXmax();
2542
2543 // first round ? --> create and fill histograms
2544 if (results->DoesExist("MVA_TRAIN_S")==0) {
2545
2546 // classifier response distributions for test sample
2547 Double_t sxmax = fXmax+0.00001;
2548
2549 // MVA plots on the training sample (check for overtraining)
2550 TH1* mva_s_tr = new TH1D( GetTestvarName() + "_Train_S",GetTestvarName() + "_Train_S", fNbinsMVAoutput, fXmin, sxmax );
2551 TH1* mva_b_tr = new TH1D( GetTestvarName() + "_Train_B",GetTestvarName() + "_Train_B", fNbinsMVAoutput, fXmin, sxmax );
2552 results->Store(mva_s_tr, "MVA_TRAIN_S");
2553 results->Store(mva_b_tr, "MVA_TRAIN_B");
2554 mva_s_tr->Sumw2();
2555 mva_b_tr->Sumw2();
2556
2557 // Training efficiency plots
2558 TH1* mva_eff_tr_s = new TH1D( GetTestvarName() + "_trainingEffS", GetTestvarName() + " (signal)",
2559 fNbinsH, xmin, xmax );
2560 TH1* mva_eff_tr_b = new TH1D( GetTestvarName() + "_trainingEffB", GetTestvarName() + " (background)",
2561 fNbinsH, xmin, xmax );
2562 results->Store(mva_eff_tr_s, "MVA_TRAINEFF_S");
2563 results->Store(mva_eff_tr_b, "MVA_TRAINEFF_B");
2564
2565 // sign if cut
2566 Int_t sign = (fCutOrientation == kPositive) ? +1 : -1;
2567
2568 std::vector<Double_t> mvaValues = GetMvaValues(0,Data()->GetNEvents());
2569 assert( (Long64_t) mvaValues.size() == Data()->GetNEvents());
2570
2571 // this method is unbinned
2572 for (Int_t ievt=0; ievt<Data()->GetNEvents(); ievt++) {
2573
2574 Data()->SetCurrentEvent(ievt);
2575 const Event* ev = GetEvent();
2576
2577 Double_t theVal = mvaValues[ievt];
2578 Double_t theWeight = ev->GetWeight();
2579
2580 TH1* theEffHist = DataInfo().IsSignal(ev) ? mva_eff_tr_s : mva_eff_tr_b;
2581 TH1* theClsHist = DataInfo().IsSignal(ev) ? mva_s_tr : mva_b_tr;
2582
2583 theClsHist->Fill( theVal, theWeight );
2584
2585 TAxis* axis = theEffHist->GetXaxis();
2586 Int_t maxbin = Int_t((theVal - axis->GetXmin())/(axis->GetXmax() - axis->GetXmin())*fNbinsH) + 1;
2587 if (sign > 0 && maxbin > fNbinsH) continue; // can happen... event doesn't count
2588 if (sign < 0 && maxbin < 1 ) continue; // can happen... event doesn't count
2589 if (sign > 0 && maxbin < 1 ) maxbin = 1;
2590 if (sign < 0 && maxbin > fNbinsH) maxbin = fNbinsH;
2591
2592 if (sign > 0) for (Int_t ibin=1; ibin<=maxbin; ibin++) theEffHist->AddBinContent( ibin , theWeight );
2593 else for (Int_t ibin=maxbin+1; ibin<=fNbinsH; ibin++) theEffHist->AddBinContent( ibin , theWeight );
2594 }
2595
2596 // normalise output distributions
2597 // uncomment those (and several others if you want unnormalized output
2598 gTools().NormHist( mva_s_tr );
2599 gTools().NormHist( mva_b_tr );
2600
2601 // renormalise to maximum
2602 mva_eff_tr_s->Scale( 1.0/TMath::Max(std::numeric_limits<double>::epsilon(), mva_eff_tr_s->GetMaximum()) );
2603 mva_eff_tr_b->Scale( 1.0/TMath::Max(std::numeric_limits<double>::epsilon(), mva_eff_tr_b->GetMaximum()) );
2604
2605 // Training background efficiency versus signal efficiency
2606 TH1* eff_bvss = new TH1D( GetTestvarName() + "_trainingEffBvsS", GetTestvarName() + "", fNbins, 0, 1 );
2607 // Training background rejection (=1-eff.) versus signal efficiency
2608 TH1* rej_bvss = new TH1D( GetTestvarName() + "_trainingRejBvsS", GetTestvarName() + "", fNbins, 0, 1 );
2609 results->Store(eff_bvss, "EFF_BVSS_TR");
2610 results->Store(rej_bvss, "REJ_BVSS_TR");
2611
2612 // use root finder
2613 // spline background efficiency plot
2614 // note that there is a bin shift when going from a TH1D object to a TGraph :-(
2616 if (fSplTrainRefS) delete fSplTrainRefS;
2617 if (fSplTrainRefB) delete fSplTrainRefB;
2618 fSplTrainRefS = new TSpline1( "spline2_signal", new TGraph( mva_eff_tr_s ) );
2619 fSplTrainRefB = new TSpline1( "spline2_background", new TGraph( mva_eff_tr_b ) );
2620
2621 // verify spline sanity
2622 gTools().CheckSplines( mva_eff_tr_s, fSplTrainRefS );
2623 gTools().CheckSplines( mva_eff_tr_b, fSplTrainRefB );
2624 }
2625
2626 // make the background-vs-signal efficiency plot
2627
2628 // create root finder
2629 RootFinder rootFinder(this, fXmin, fXmax );
2630
2631 Double_t effB = 0;
2632 fEffS = results->GetHist("MVA_TRAINEFF_S");
2633 for (Int_t bini=1; bini<=fNbins; bini++) {
2634
2635 // find cut value corresponding to a given signal efficiency
2636 Double_t effS = eff_bvss->GetBinCenter( bini );
2637
2638 Double_t cut = rootFinder.Root( effS );
2639
2640 // retrieve background efficiency for given cut
2641 if (Use_Splines_for_Eff_) effB = fSplTrainRefB->Eval( cut );
2642 else effB = mva_eff_tr_b->GetBinContent( mva_eff_tr_b->FindBin( cut ) );
2643
2644 // and fill histograms
2645 eff_bvss->SetBinContent( bini, effB );
2646 rej_bvss->SetBinContent( bini, 1.0-effB );
2647 }
2648 fEffS = 0;
2649
2650 // create splines for histogram
2651 fSplTrainEffBvsS = new TSpline1( "effBvsS", new TGraph( eff_bvss ) );
2652 }
2653
2654 // must exist...
2655 if (0 == fSplTrainEffBvsS) return 0.0;
2656
2657 // now find signal efficiency that corresponds to required background efficiency
2658 Double_t effS = 0., effB, effS_ = 0., effB_ = 0.;
2659 Int_t nbins_ = 1000;
2660 for (Int_t bini=1; bini<=nbins_; bini++) {
2661
2662 // get corresponding signal and background efficiencies
2663 effS = (bini - 0.5)/Float_t(nbins_);
2664 effB = fSplTrainEffBvsS->Eval( effS );
2665
2666 // find signal efficiency that corresponds to required background efficiency
2667 if ((effB - effBref)*(effB_ - effBref) <= 0) break;
2668 effS_ = effS;
2669 effB_ = effB;
2670 }
2671
2672 return 0.5*(effS + effS_); // the mean between bin above and bin below
2673}
2674
2675////////////////////////////////////////////////////////////////////////////////
2676
2677std::vector<Float_t> TMVA::MethodBase::GetMulticlassEfficiency(std::vector<std::vector<Float_t> >& purity)
2678{
2679 Data()->SetCurrentType(Types::kTesting);
2680 ResultsMulticlass* resMulticlass = dynamic_cast<ResultsMulticlass*>(Data()->GetResults(GetMethodName(), Types::kTesting, Types::kMulticlass));
2681 if (!resMulticlass) Log() << kFATAL<<Form("Dataset[%s] : ",DataInfo().GetName())<< "unable to create pointer in GetMulticlassEfficiency, exiting."<<Endl;
2682
2683 purity.push_back(resMulticlass->GetAchievablePur());
2684 return resMulticlass->GetAchievableEff();
2685}
2686
2687////////////////////////////////////////////////////////////////////////////////
2688
2689std::vector<Float_t> TMVA::MethodBase::GetMulticlassTrainingEfficiency(std::vector<std::vector<Float_t> >& purity)
2690{
2691 Data()->SetCurrentType(Types::kTraining);
2692 ResultsMulticlass* resMulticlass = dynamic_cast<ResultsMulticlass*>(Data()->GetResults(GetMethodName(), Types::kTraining, Types::kMulticlass));
2693 if (!resMulticlass) Log() << kFATAL<< "unable to create pointer in GetMulticlassTrainingEfficiency, exiting."<<Endl;
2694
2695 Log() << kINFO <<Form("Dataset[%s] : ",DataInfo().GetName())<< "Determine optimal multiclass cuts for training data..." << Endl;
2696 for (UInt_t icls = 0; icls<DataInfo().GetNClasses(); ++icls) {
2697 resMulticlass->GetBestMultiClassCuts(icls);
2698 }
2699
2700 purity.push_back(resMulticlass->GetAchievablePur());
2701 return resMulticlass->GetAchievableEff();
2702}
2703
2704////////////////////////////////////////////////////////////////////////////////
2705/// Construct a confusion matrix for a multiclass classifier. The confusion
2706/// matrix compares, in turn, each class agaist all other classes in a pair-wise
2707/// fashion. In rows with index \f$ k_r = 0 ... K \f$, \f$ k_r \f$ is
2708/// considered signal for the sake of comparison and for each column
2709/// \f$ k_c = 0 ... K \f$ the corresponding class is considered background.
2710///
2711/// Note that the diagonal elements will be returned as NaN since this will
2712/// compare a class against itself.
2713///
2714/// \see TMVA::ResultsMulticlass::GetConfusionMatrix
2715///
2716/// \param[in] effB The background efficiency for which to evaluate.
2717/// \param[in] type The data set on which to evaluate (training, testing ...).
2718///
2719/// \return A matrix containing signal efficiencies for the given background
2720/// efficiency. The diagonal elements are NaN since this measure is
2721/// meaningless (comparing a class against itself).
2722///
2723
2725{
2726 if (GetAnalysisType() != Types::kMulticlass) {
2727 Log() << kFATAL << "Cannot get confusion matrix for non-multiclass analysis." << std::endl;
2728 return TMatrixD(0, 0);
2729 }
2730
2731 Data()->SetCurrentType(type);
2732 ResultsMulticlass *resMulticlass =
2733 dynamic_cast<ResultsMulticlass *>(Data()->GetResults(GetMethodName(), type, Types::kMulticlass));
2734
2735 if (resMulticlass == nullptr) {
2736 Log() << kFATAL << Form("Dataset[%s] : ", DataInfo().GetName())
2737 << "unable to create pointer in GetMulticlassEfficiency, exiting." << Endl;
2738 return TMatrixD(0, 0);
2739 }
2740
2741 return resMulticlass->GetConfusionMatrix(effB);
2742}
2743
2744////////////////////////////////////////////////////////////////////////////////
2745/// compute significance of mean difference
2746/// \f[
2747/// significance = \frac{|<S> - <B>|}{\sqrt{RMS_{S2} + RMS_{B2}}}
2748/// \f]
2749
2751{
2752 Double_t rms = sqrt( fRmsS*fRmsS + fRmsB*fRmsB );
2753
2754 return (rms > 0) ? TMath::Abs(fMeanS - fMeanB)/rms : 0;
2755}
2756
2757////////////////////////////////////////////////////////////////////////////////
2758/// compute "separation" defined as
2759/// \f[
2760/// <s2> = \frac{1}{2} \int_{-\infty}^{+\infty} { \frac{(S(x) - B(x))^2}{(S(x) + B(x))} dx }
2761/// \f]
2762
2764{
2765 return gTools().GetSeparation( histoS, histoB );
2766}
2767
2768////////////////////////////////////////////////////////////////////////////////
2769/// compute "separation" defined as
2770/// \f[
2771/// <s2> = \frac{1}{2} \int_{-\infty}^{+\infty} { \frac{(S(x) - B(x))^2}{(S(x) + B(x))} dx }
2772/// \f]
2773
2775{
2776 // note, if zero pointers given, use internal pdf
2777 // sanity check first
2778 if ((!pdfS && pdfB) || (pdfS && !pdfB))
2779 Log() << kFATAL <<Form("Dataset[%s] : ",DataInfo().GetName())<< "<GetSeparation> Mismatch in pdfs" << Endl;
2780 if (!pdfS) pdfS = fSplS;
2781 if (!pdfB) pdfB = fSplB;
2782
2783 if (!fSplS || !fSplB) {
2784 Log()<<kDEBUG<<Form("[%s] : ",DataInfo().GetName())<< "could not calculate the separation, distributions"
2785 << " fSplS or fSplB are not yet filled" << Endl;
2786 return 0;
2787 }else{
2788 return gTools().GetSeparation( *pdfS, *pdfB );
2789 }
2790}
2791
2792////////////////////////////////////////////////////////////////////////////////
2793/// calculate the area (integral) under the ROC curve as a
2794/// overall quality measure of the classification
2795
2797{
2798 // note, if zero pointers given, use internal pdf
2799 // sanity check first
2800 if ((!histS && histB) || (histS && !histB))
2801 Log() << kFATAL <<Form("Dataset[%s] : ",DataInfo().GetName())<< "<GetROCIntegral(TH1D*, TH1D*)> Mismatch in hists" << Endl;
2802
2803 if (histS==0 || histB==0) return 0.;
2804
2805 TMVA::PDF *pdfS = new TMVA::PDF( " PDF Sig", histS, TMVA::PDF::kSpline3 );
2806 TMVA::PDF *pdfB = new TMVA::PDF( " PDF Bkg", histB, TMVA::PDF::kSpline3 );
2807
2808
2809 Double_t xmin = TMath::Min(pdfS->GetXmin(), pdfB->GetXmin());
2810 Double_t xmax = TMath::Max(pdfS->GetXmax(), pdfB->GetXmax());
2811
2812 Double_t integral = 0;
2813 UInt_t nsteps = 1000;
2814 Double_t step = (xmax-xmin)/Double_t(nsteps);
2815 Double_t cut = xmin;
2816 for (UInt_t i=0; i<nsteps; i++) {
2817 integral += (1-pdfB->GetIntegral(cut,xmax)) * pdfS->GetVal(cut);
2818 cut+=step;
2819 }
2820 delete pdfS;
2821 delete pdfB;
2822 return integral*step;
2823}
2824
2825
2826////////////////////////////////////////////////////////////////////////////////
2827/// calculate the area (integral) under the ROC curve as a
2828/// overall quality measure of the classification
2829
2831{
2832 // note, if zero pointers given, use internal pdf
2833 // sanity check first
2834 if ((!pdfS && pdfB) || (pdfS && !pdfB))
2835 Log() << kFATAL <<Form("Dataset[%s] : ",DataInfo().GetName())<< "<GetSeparation> Mismatch in pdfs" << Endl;
2836 if (!pdfS) pdfS = fSplS;
2837 if (!pdfB) pdfB = fSplB;
2838
2839 if (pdfS==0 || pdfB==0) return 0.;
2840
2841 Double_t xmin = TMath::Min(pdfS->GetXmin(), pdfB->GetXmin());
2842 Double_t xmax = TMath::Max(pdfS->GetXmax(), pdfB->GetXmax());
2843
2844 Double_t integral = 0;
2845 UInt_t nsteps = 1000;
2846 Double_t step = (xmax-xmin)/Double_t(nsteps);
2847 Double_t cut = xmin;
2848 for (UInt_t i=0; i<nsteps; i++) {
2849 integral += (1-pdfB->GetIntegral(cut,xmax)) * pdfS->GetVal(cut);
2850 cut+=step;
2851 }
2852 return integral*step;
2853}
2854
2855////////////////////////////////////////////////////////////////////////////////
2856/// plot significance, \f$ \frac{S}{\sqrt{S^2 + B^2}} \f$, curve for given number
2857/// of signal and background events; returns cut for maximum significance
2858/// also returned via reference is the maximum significance
2859
2861 Double_t BackgroundEvents,
2862 Double_t& max_significance_value ) const
2863{
2864 Results* results = Data()->GetResults( GetMethodName(), Types::kTesting, Types::kMaxAnalysisType );
2865
2866 Double_t max_significance(0);
2867 Double_t effS(0),effB(0),significance(0);
2868 TH1D *temp_histogram = new TH1D("temp", "temp", fNbinsH, fXmin, fXmax );
2869
2870 if (SignalEvents <= 0 || BackgroundEvents <= 0) {
2871 Log() << kFATAL <<Form("Dataset[%s] : ",DataInfo().GetName())<< "<GetMaximumSignificance> "
2872 << "Number of signal or background events is <= 0 ==> abort"
2873 << Endl;
2874 }
2875
2876 Log() << kINFO <<Form("Dataset[%s] : ",DataInfo().GetName())<< "Using ratio SignalEvents/BackgroundEvents = "
2877 << SignalEvents/BackgroundEvents << Endl;
2878
2879 TH1* eff_s = results->GetHist("MVA_EFF_S");
2880 TH1* eff_b = results->GetHist("MVA_EFF_B");
2881
2882 if ( (eff_s==0) || (eff_b==0) ) {
2883 Log() << kWARNING <<Form("Dataset[%s] : ",DataInfo().GetName())<< "Efficiency histograms empty !" << Endl;
2884 Log() << kWARNING <<Form("Dataset[%s] : ",DataInfo().GetName())<< "no maximum cut found, return 0" << Endl;
2885 return 0;
2886 }
2887
2888 for (Int_t bin=1; bin<=fNbinsH; bin++) {
2889 effS = eff_s->GetBinContent( bin );
2890 effB = eff_b->GetBinContent( bin );
2891
2892 // put significance into a histogram
2893 significance = sqrt(SignalEvents)*( effS )/sqrt( effS + ( BackgroundEvents / SignalEvents) * effB );
2894
2895 temp_histogram->SetBinContent(bin,significance);
2896 }
2897
2898 // find maximum in histogram
2899 max_significance = temp_histogram->GetBinCenter( temp_histogram->GetMaximumBin() );
2900 max_significance_value = temp_histogram->GetBinContent( temp_histogram->GetMaximumBin() );
2901
2902 // delete
2903 delete temp_histogram;
2904
2905 Log() << kINFO <<Form("Dataset[%s] : ",DataInfo().GetName())<< "Optimal cut at : " << max_significance << Endl;
2906 Log() << kINFO<<Form("Dataset[%s] : ",DataInfo().GetName()) << "Maximum significance: " << max_significance_value << Endl;
2907
2908 return max_significance;
2909}
2910
2911////////////////////////////////////////////////////////////////////////////////
2912/// calculates rms,mean, xmin, xmax of the event variable
2913/// this can be either done for the variables as they are or for
2914/// normalised variables (in the range of 0-1) if "norm" is set to kTRUE
2915
2917 Double_t& meanS, Double_t& meanB,
2918 Double_t& rmsS, Double_t& rmsB,
2920{
2921 Types::ETreeType previousTreeType = Data()->GetCurrentType();
2922 Data()->SetCurrentType(treeType);
2923
2924 Long64_t entries = Data()->GetNEvents();
2925
2926 // sanity check
2927 if (entries <=0)
2928 Log() << kFATAL <<Form("Dataset[%s] : ",DataInfo().GetName())<< "<CalculateEstimator> Wrong tree type: " << treeType << Endl;
2929
2930 // index of the wanted variable
2931 UInt_t varIndex = DataInfo().FindVarIndex( theVarName );
2932
2933 // first fill signal and background in arrays before analysis
2934 xmin = +DBL_MAX;
2935 xmax = -DBL_MAX;
2936
2937 // take into account event weights
2938 meanS = 0;
2939 meanB = 0;
2940 rmsS = 0;
2941 rmsB = 0;
2942 Double_t sumwS = 0, sumwB = 0;
2943
2944 // loop over all training events
2945 for (Int_t ievt = 0; ievt < entries; ievt++) {
2946
2947 const Event* ev = GetEvent(ievt);
2948
2949 Double_t theVar = ev->GetValue(varIndex);
2950 Double_t weight = ev->GetWeight();
2951
2952 if (DataInfo().IsSignal(ev)) {
2953 sumwS += weight;
2954 meanS += weight*theVar;
2955 rmsS += weight*theVar*theVar;
2956 }
2957 else {
2958 sumwB += weight;
2959 meanB += weight*theVar;
2960 rmsB += weight*theVar*theVar;
2961 }
2962 xmin = TMath::Min( xmin, theVar );
2963 xmax = TMath::Max( xmax, theVar );
2964 }
2965
2966 meanS = meanS/sumwS;
2967 meanB = meanB/sumwB;
2968 rmsS = TMath::Sqrt( rmsS/sumwS - meanS*meanS );
2969 rmsB = TMath::Sqrt( rmsB/sumwB - meanB*meanB );
2970
2971 Data()->SetCurrentType(previousTreeType);
2972}
2973
2974////////////////////////////////////////////////////////////////////////////////
2975/// create reader class for method (classification only at present)
2976
2977void TMVA::MethodBase::MakeClass( const TString& theClassFileName ) const
2978{
2979 // the default consists of
2980 TString classFileName = "";
2981 if (theClassFileName == "")
2982 classFileName = GetWeightFileDir() + "/" + GetJobName() + "_" + GetMethodName() + ".class.C";
2983 else
2984 classFileName = theClassFileName;
2985
2986 TString className = TString("Read") + GetMethodName();
2987
2988 TString tfname( classFileName );
2989 Log() << kINFO //<<Form("Dataset[%s] : ",DataInfo().GetName())
2990 << "Creating standalone class: "
2991 << gTools().Color("lightblue") << classFileName << gTools().Color("reset") << Endl;
2992
2993 std::ofstream fout( classFileName );
2994 if (!fout.good()) { // file could not be opened --> Error
2995 Log() << kFATAL << "<MakeClass> Unable to open file: " << classFileName << Endl;
2996 }
2997
2998 // now create the class
2999 // preamble
3000 fout << "// Class: " << className << std::endl;
3001 fout << "// Automatically generated by MethodBase::MakeClass" << std::endl << "//" << std::endl;
3002
3003 // print general information and configuration state
3004 fout << std::endl;
3005 fout << "/* configuration options =====================================================" << std::endl << std::endl;
3006 WriteStateToStream( fout );
3007 fout << std::endl;
3008 fout << "============================================================================ */" << std::endl;
3009
3010 // generate the class
3011 fout << "" << std::endl;
3012 fout << "#include <array>" << std::endl;
3013 fout << "#include <vector>" << std::endl;
3014 fout << "#include <cmath>" << std::endl;
3015 fout << "#include <string>" << std::endl;
3016 fout << "#include <iostream>" << std::endl;
3017 fout << "" << std::endl;
3018 // now if the classifier needs to write some additional classes for its response implementation
3019 // this code goes here: (at least the header declarations need to come before the main class
3020 this->MakeClassSpecificHeader( fout, className );
3021
3022 fout << "#ifndef IClassifierReader__def" << std::endl;
3023 fout << "#define IClassifierReader__def" << std::endl;
3024 fout << std::endl;
3025 fout << "class IClassifierReader {" << std::endl;
3026 fout << std::endl;
3027 fout << " public:" << std::endl;
3028 fout << std::endl;
3029 fout << " // constructor" << std::endl;
3030 fout << " IClassifierReader() : fStatusIsClean( true ) {}" << std::endl;
3031 fout << " virtual ~IClassifierReader() {}" << std::endl;
3032 fout << std::endl;
3033 fout << " // return classifier response" << std::endl;
3034 if(GetAnalysisType() == Types::kMulticlass) {
3035 fout << " virtual std::vector<double> GetMulticlassValues( const std::vector<double>& inputValues ) const = 0;" << std::endl;
3036 } else {
3037 fout << " virtual double GetMvaValue( const std::vector<double>& inputValues ) const = 0;" << std::endl;
3038 }
3039 fout << std::endl;
3040 fout << " // returns classifier status" << std::endl;
3041 fout << " bool IsStatusClean() const { return fStatusIsClean; }" << std::endl;
3042 fout << std::endl;
3043 fout << " protected:" << std::endl;
3044 fout << std::endl;
3045 fout << " bool fStatusIsClean;" << std::endl;
3046 fout << "};" << std::endl;
3047 fout << std::endl;
3048 fout << "#endif" << std::endl;
3049 fout << std::endl;
3050 fout << "class " << className << " : public IClassifierReader {" << std::endl;
3051 fout << std::endl;
3052 fout << " public:" << std::endl;
3053 fout << std::endl;
3054 fout << " // constructor" << std::endl;
3055 fout << " " << className << "( std::vector<std::string>& theInputVars )" << std::endl;
3056 fout << " : IClassifierReader()," << std::endl;
3057 fout << " fClassName( \"" << className << "\" )," << std::endl;
3058 fout << " fNvars( " << GetNvar() << " )" << std::endl;
3059 fout << " {" << std::endl;
3060 fout << " // the training input variables" << std::endl;
3061 fout << " const char* inputVars[] = { ";
3062 for (UInt_t ivar=0; ivar<GetNvar(); ivar++) {
3063 fout << "\"" << GetOriginalVarName(ivar) << "\"";
3064 if (ivar<GetNvar()-1) fout << ", ";
3065 }
3066 fout << " };" << std::endl;
3067 fout << std::endl;
3068 fout << " // sanity checks" << std::endl;
3069 fout << " if (theInputVars.size() <= 0) {" << std::endl;
3070 fout << " std::cout << \"Problem in class \\\"\" << fClassName << \"\\\": empty input vector\" << std::endl;" << std::endl;
3071 fout << " fStatusIsClean = false;" << std::endl;
3072 fout << " }" << std::endl;
3073 fout << std::endl;
3074 fout << " if (theInputVars.size() != fNvars) {" << std::endl;
3075 fout << " std::cout << \"Problem in class \\\"\" << fClassName << \"\\\": mismatch in number of input values: \"" << std::endl;
3076 fout << " << theInputVars.size() << \" != \" << fNvars << std::endl;" << std::endl;
3077 fout << " fStatusIsClean = false;" << std::endl;
3078 fout << " }" << std::endl;
3079 fout << std::endl;
3080 fout << " // validate input variables" << std::endl;
3081 fout << " for (size_t ivar = 0; ivar < theInputVars.size(); ivar++) {" << std::endl;
3082 fout << " if (theInputVars[ivar] != inputVars[ivar]) {" << std::endl;
3083 fout << " std::cout << \"Problem in class \\\"\" << fClassName << \"\\\": mismatch in input variable names\" << std::endl" << std::endl;
3084 fout << " << \" for variable [\" << ivar << \"]: \" << theInputVars[ivar].c_str() << \" != \" << inputVars[ivar] << std::endl;" << std::endl;
3085 fout << " fStatusIsClean = false;" << std::endl;
3086 fout << " }" << std::endl;
3087 fout << " }" << std::endl;
3088 fout << std::endl;
3089 fout << " // initialize min and max vectors (for normalisation)" << std::endl;
3090 for (UInt_t ivar = 0; ivar < GetNvar(); ivar++) {
3091 fout << " fVmin[" << ivar << "] = " << std::setprecision(15) << GetXmin( ivar ) << ";" << std::endl;
3092 fout << " fVmax[" << ivar << "] = " << std::setprecision(15) << GetXmax( ivar ) << ";" << std::endl;
3093 }
3094 fout << std::endl;
3095 fout << " // initialize input variable types" << std::endl;
3096 for (UInt_t ivar=0; ivar<GetNvar(); ivar++) {
3097 fout << " fType[" << ivar << "] = \'" << DataInfo().GetVariableInfo(ivar).GetVarType() << "\';" << std::endl;
3098 }
3099 fout << std::endl;
3100 fout << " // initialize constants" << std::endl;
3101 fout << " Initialize();" << std::endl;
3102 fout << std::endl;
3103 if (GetTransformationHandler().GetTransformationList().GetSize() != 0) {
3104 fout << " // initialize transformation" << std::endl;
3105 fout << " InitTransform();" << std::endl;
3106 }
3107 fout << " }" << std::endl;
3108 fout << std::endl;
3109 fout << " // destructor" << std::endl;
3110 fout << " virtual ~" << className << "() {" << std::endl;
3111 fout << " Clear(); // method-specific" << std::endl;
3112 fout << " }" << std::endl;
3113 fout << std::endl;
3114 fout << " // the classifier response" << std::endl;
3115 fout << " // \"inputValues\" is a vector of input values in the same order as the" << std::endl;
3116 fout << " // variables given to the constructor" << std::endl;
3117 if(GetAnalysisType() == Types::kMulticlass) {
3118 fout << " std::vector<double> GetMulticlassValues( const std::vector<double>& inputValues ) const override;" << std::endl;
3119 } else {
3120 fout << " double GetMvaValue( const std::vector<double>& inputValues ) const override;" << std::endl;
3121 }
3122 fout << std::endl;
3123 fout << " private:" << std::endl;
3124 fout << std::endl;
3125 fout << " // method-specific destructor" << std::endl;
3126 fout << " void Clear();" << std::endl;
3127 fout << std::endl;
3128 if (GetTransformationHandler().GetTransformationList().GetSize()!=0) {
3129 fout << " // input variable transformation" << std::endl;
3130 GetTransformationHandler().MakeFunction(fout, className,1);
3131 fout << " void InitTransform();" << std::endl;
3132 fout << " void Transform( std::vector<double> & iv, int sigOrBgd ) const;" << std::endl;
3133 fout << std::endl;
3134 }
3135 fout << " // common member variables" << std::endl;
3136 fout << " const char* fClassName;" << std::endl;
3137 fout << std::endl;
3138 fout << " const size_t fNvars;" << std::endl;
3139 fout << " size_t GetNvar() const { return fNvars; }" << std::endl;
3140 fout << " char GetType( int ivar ) const { return fType[ivar]; }" << std::endl;
3141 fout << std::endl;
3142 fout << " // normalisation of input variables" << std::endl;
3143 fout << " double fVmin[" << GetNvar() << "];" << std::endl;
3144 fout << " double fVmax[" << GetNvar() << "];" << std::endl;
3145 fout << " double NormVariable( double x, double xmin, double xmax ) const {" << std::endl;
3146 fout << " // normalise to output range: [-1, 1]" << std::endl;
3147 fout << " return 2*(x - xmin)/(xmax - xmin) - 1.0;" << std::endl;
3148 fout << " }" << std::endl;
3149 fout << std::endl;
3150 fout << " // type of input variable: 'F' or 'I'" << std::endl;
3151 fout << " char fType[" << GetNvar() << "];" << std::endl;
3152 fout << std::endl;
3153 fout << " // initialize internal variables" << std::endl;
3154 fout << " void Initialize();" << std::endl;
3155 if(GetAnalysisType() == Types::kMulticlass) {
3156 fout << " std::vector<double> GetMulticlassValues__( const std::vector<double>& inputValues ) const;" << std::endl;
3157 } else {
3158 fout << " double GetMvaValue__( const std::vector<double>& inputValues ) const;" << std::endl;
3159 }
3160 fout << "" << std::endl;
3161 fout << " // private members (method specific)" << std::endl;
3162
3163 // call the classifier specific output (the classifier must close the class !)
3164 MakeClassSpecific( fout, className );
3165
3166 if(GetAnalysisType() == Types::kMulticlass) {
3167 fout << "inline std::vector<double> " << className << "::GetMulticlassValues( const std::vector<double>& inputValues ) const" << std::endl;
3168 } else {
3169 fout << "inline double " << className << "::GetMvaValue( const std::vector<double>& inputValues ) const" << std::endl;
3170 }
3171 fout << "{" << std::endl;
3172 fout << " // classifier response value" << std::endl;
3173 if(GetAnalysisType() == Types::kMulticlass) {
3174 fout << " std::vector<double> retval;" << std::endl;
3175 } else {
3176 fout << " double retval = 0;" << std::endl;
3177 }
3178 fout << std::endl;
3179 fout << " // classifier response, sanity check first" << std::endl;
3180 fout << " if (!IsStatusClean()) {" << std::endl;
3181 fout << " std::cout << \"Problem in class \\\"\" << fClassName << \"\\\": cannot return classifier response\"" << std::endl;
3182 fout << " << \" because status is dirty\" << std::endl;" << std::endl;
3183 fout << " }" << std::endl;
3184 fout << " else {" << std::endl;
3185 if (IsNormalised()) {
3186 fout << " // normalise variables" << std::endl;
3187 fout << " std::vector<double> iV;" << std::endl;
3188 fout << " iV.reserve(inputValues.size());" << std::endl;
3189 fout << " int ivar = 0;" << std::endl;
3190 fout << " for (std::vector<double>::const_iterator varIt = inputValues.begin();" << std::endl;
3191 fout << " varIt != inputValues.end(); varIt++, ivar++) {" << std::endl;
3192 fout << " iV.push_back(NormVariable( *varIt, fVmin[ivar], fVmax[ivar] ));" << std::endl;
3193 fout << " }" << std::endl;
3194 if (GetTransformationHandler().GetTransformationList().GetSize() != 0 && GetMethodType() != Types::kLikelihood &&
3195 GetMethodType() != Types::kHMatrix) {
3196 fout << " Transform( iV, -1 );" << std::endl;
3197 }
3198
3199 if(GetAnalysisType() == Types::kMulticlass) {
3200 fout << " retval = GetMulticlassValues__( iV );" << std::endl;
3201 } else {
3202 fout << " retval = GetMvaValue__( iV );" << std::endl;
3203 }
3204 } else {
3205 if (GetTransformationHandler().GetTransformationList().GetSize() != 0 && GetMethodType() != Types::kLikelihood &&
3206 GetMethodType() != Types::kHMatrix) {
3207 fout << " std::vector<double> iV(inputValues);" << std::endl;
3208 fout << " Transform( iV, -1 );" << std::endl;
3209 if(GetAnalysisType() == Types::kMulticlass) {
3210 fout << " retval = GetMulticlassValues__( iV );" << std::endl;
3211 } else {
3212 fout << " retval = GetMvaValue__( iV );" << std::endl;
3213 }
3214 } else {
3215 if(GetAnalysisType() == Types::kMulticlass) {
3216 fout << " retval = GetMulticlassValues__( inputValues );" << std::endl;
3217 } else {
3218 fout << " retval = GetMvaValue__( inputValues );" << std::endl;
3219 }
3220 }
3221 }
3222 fout << " }" << std::endl;
3223 fout << std::endl;
3224 fout << " return retval;" << std::endl;
3225 fout << "}" << std::endl;
3226
3227 // create output for transformation - if any
3228 if (GetTransformationHandler().GetTransformationList().GetSize()!=0)
3229 GetTransformationHandler().MakeFunction(fout, className,2);
3230
3231 // close the file
3232 fout.close();
3233}
3234
3235////////////////////////////////////////////////////////////////////////////////
3236/// prints out method-specific help method
3237
3239{
3240 // if options are written to reference file, also append help info
3241 std::streambuf* cout_sbuf = std::cout.rdbuf(); // save original sbuf
3242 std::ofstream* o = 0;
3243 if (gConfig().WriteOptionsReference()) {
3244 Log() << kINFO << "Print Help message for class " << GetName() << " into file: " << GetReferenceFile() << Endl;
3245 o = new std::ofstream( GetReferenceFile(), std::ios::app );
3246 if (!o->good()) { // file could not be opened --> Error
3247 Log() << kFATAL << "<PrintHelpMessage> Unable to append to output file: " << GetReferenceFile() << Endl;
3248 }
3249 std::cout.rdbuf( o->rdbuf() ); // redirect 'std::cout' to file
3250 }
3251
3252 // "|--------------------------------------------------------------|"
3253 if (!o) {
3254 Log() << kINFO << Endl;
3255 Log() << gTools().Color("bold")
3256 << "================================================================"
3257 << gTools().Color( "reset" )
3258 << Endl;
3259 Log() << gTools().Color("bold")
3260 << "H e l p f o r M V A m e t h o d [ " << GetName() << " ] :"
3261 << gTools().Color( "reset" )
3262 << Endl;
3263 }
3264 else {
3265 Log() << "Help for MVA method [ " << GetName() << " ] :" << Endl;
3266 }
3267
3268 // print method-specific help message
3269 GetHelpMessage();
3270
3271 if (!o) {
3272 Log() << Endl;
3273 Log() << "<Suppress this message by specifying \"!H\" in the booking option>" << Endl;
3274 Log() << gTools().Color("bold")
3275 << "================================================================"
3276 << gTools().Color( "reset" )
3277 << Endl;
3278 Log() << Endl;
3279 }
3280 else {
3281 // indicate END
3282 Log() << "# End of Message___" << Endl;
3283 }
3284
3285 std::cout.rdbuf( cout_sbuf ); // restore the original stream buffer
3286 if (o) o->close();
3287}
3288
3289// ----------------------- r o o t f i n d i n g ----------------------------
3290
3291////////////////////////////////////////////////////////////////////////////////
3292/// returns efficiency as function of cut
3293
3295{
3296 Double_t retval=0;
3297
3298 // retrieve the class object
3300 retval = fSplRefS->Eval( theCut );
3301 }
3302 else retval = fEffS->GetBinContent( fEffS->FindBin( theCut ) );
3303
3304 // caution: here we take some "forbidden" action to hide a problem:
3305 // in some cases, in particular for likelihood, the binned efficiency distributions
3306 // do not equal 1, at xmin, and 0 at xmax; of course, in principle we have the
3307 // unbinned information available in the trees, but the unbinned minimization is
3308 // too slow, and we don't need to do a precision measurement here. Hence, we force
3309 // this property.
3310 Double_t eps = 1.0e-5;
3311 if (theCut-fXmin < eps) retval = (GetCutOrientation() == kPositive) ? 1.0 : 0.0;
3312 else if (fXmax-theCut < eps) retval = (GetCutOrientation() == kPositive) ? 0.0 : 1.0;
3313
3314 return retval;
3315}
3316
3317////////////////////////////////////////////////////////////////////////////////
3318/// returns the event collection (i.e. the dataset) TRANSFORMED using the
3319/// classifiers specific Variable Transformation (e.g. Decorr or Decorr:Gauss:Decorr)
3320
3322{
3323 // if there's no variable transformation for this classifier, just hand back the
3324 // event collection of the data set
3325 if (GetTransformationHandler().GetTransformationList().GetEntries() <= 0) {
3326 return (Data()->GetEventCollection(type));
3327 }
3328
3329 // otherwise, transform ALL the events and hand back the vector of the pointers to the
3330 // transformed events. If the pointer is already != 0, i.e. the whole thing has been
3331 // done before, I don't need to do it again, but just "hand over" the pointer to those events.
3332 Int_t idx = Data()->TreeIndex(type); //index indicating Training,Testing,... events/datasets
3333 if (fEventCollections.at(idx) == 0) {
3334 fEventCollections.at(idx) = &(Data()->GetEventCollection(type));
3335 fEventCollections.at(idx) = GetTransformationHandler().CalcTransformations(*(fEventCollections.at(idx)),kTRUE);
3336 }
3337 return *(fEventCollections.at(idx));
3338}
3339
3340////////////////////////////////////////////////////////////////////////////////
3341/// calculates the TMVA version string from the training version code on the fly
3342
3344{
3345 UInt_t a = GetTrainingTMVAVersionCode() & 0xff0000; a>>=16;
3346 UInt_t b = GetTrainingTMVAVersionCode() & 0x00ff00; b>>=8;
3347 UInt_t c = GetTrainingTMVAVersionCode() & 0x0000ff;
3348
3349 return TString::Format("%i.%i.%i",a,b,c);
3350}
3351
3352////////////////////////////////////////////////////////////////////////////////
3353/// calculates the ROOT version string from the training version code on the fly
3354
3356{
3357 UInt_t a = GetTrainingROOTVersionCode() & 0xff0000; a>>=16;
3358 UInt_t b = GetTrainingROOTVersionCode() & 0x00ff00; b>>=8;
3359 UInt_t c = GetTrainingROOTVersionCode() & 0x0000ff;
3360
3361 return TString::Format("%i.%02i/%02i",a,b,c);
3362}
3363
3364////////////////////////////////////////////////////////////////////////////////
3365
3367 ResultsClassification* mvaRes = dynamic_cast<ResultsClassification*>
3368 ( Data()->GetResults(GetMethodName(),Types::kTesting, Types::kClassification) );
3369
3370 if (mvaRes != NULL) {
3371 TH1D *mva_s = dynamic_cast<TH1D*> (mvaRes->GetHist("MVA_S"));
3372 TH1D *mva_b = dynamic_cast<TH1D*> (mvaRes->GetHist("MVA_B"));
3373 TH1D *mva_s_tr = dynamic_cast<TH1D*> (mvaRes->GetHist("MVA_TRAIN_S"));
3374 TH1D *mva_b_tr = dynamic_cast<TH1D*> (mvaRes->GetHist("MVA_TRAIN_B"));
3375
3376 if ( !mva_s || !mva_b || !mva_s_tr || !mva_b_tr) return -1;
3377
3378 if (SorB == 's' || SorB == 'S')
3379 return mva_s->KolmogorovTest( mva_s_tr, opt.Data() );
3380 else
3381 return mva_b->KolmogorovTest( mva_b_tr, opt.Data() );
3382 }
3383 return -1;
3384}
const Bool_t Use_Splines_for_Eff_
const Int_t NBIN_HIST_HIGH
#define d(i)
Definition RSha256.hxx:102
#define b(i)
Definition RSha256.hxx:100
#define c(i)
Definition RSha256.hxx:101
#define a(i)
Definition RSha256.hxx:99
#define s1(x)
Definition RSha256.hxx:91
#define ROOT_VERSION_CODE
Definition RVersion.hxx:27
bool Bool_t
Boolean (0=false, 1=true) (bool)
Definition RtypesCore.h:78
int Int_t
Signed integer 4 bytes (int)
Definition RtypesCore.h:60
char Char_t
Character 1 byte (char)
Definition RtypesCore.h:52
float Float_t
Float 4 bytes (float)
Definition RtypesCore.h:72
constexpr Bool_t kFALSE
Definition RtypesCore.h:109
double Double_t
Double 8 bytes.
Definition RtypesCore.h:74
long long Long64_t
Portable signed long integer 8 bytes.
Definition RtypesCore.h:84
constexpr Bool_t kTRUE
Definition RtypesCore.h:108
Option_t Option_t TPoint TPoint const char GetTextMagnitude GetFillStyle GetLineColor GetLineWidth GetMarkerStyle GetTextAlign GetTextColor GetTextSize void data
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 r
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
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 Int_t Int_t Window_t TString Int_t GCValues_t GetPrimarySelectionOwner GetDisplay GetScreen GetColormap GetNativeEvent const char const char dpyName wid window const char font_name cursor keysym reg const char only_if_exist regb h Point_t winding char text const char depth char const char Int_t count const char ColorStruct_t color const char Pixmap_t Pixmap_t PictureAttributes_t attr const char char ret_data h unsigned char height h length
Option_t Option_t TPoint TPoint const char GetTextMagnitude GetFillStyle GetLineColor GetLineWidth GetMarkerStyle GetTextAlign GetTextColor GetTextSize void value
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 Int_t Int_t Window_t TString Int_t GCValues_t GetPrimarySelectionOwner GetDisplay GetScreen GetColormap GetNativeEvent const char const char dpyName wid window const char font_name cursor keysym reg const char only_if_exist regb h Point_t winding char text const char depth char const char Int_t count const char ColorStruct_t color const char Pixmap_t Pixmap_t PictureAttributes_t attr const char char ret_data h unsigned char height h Atom_t Int_t ULong_t ULong_t unsigned char prop_list Atom_t Atom_t Atom_t Time_t type
char name[80]
Definition TGX11.cxx:142
float xmin
float xmax
TMatrixT< Double_t > TMatrixD
Definition TMatrixDfwd.h:23
char * Form(const char *fmt,...)
Formats a string in a circular formatting buffer.
Definition TString.cxx:2571
R__EXTERN TSystem * gSystem
Definition TSystem.h:582
#define TMVA_VERSION_CODE
Definition Version.h:47
Class to manage histogram axis.
Definition TAxis.h:32
Double_t GetXmax() const
Definition TAxis.h:142
Double_t GetXmin() const
Definition TAxis.h:141
Int_t Write(const char *name=nullptr, Int_t option=0, Int_t bufsize=0) override
Write all objects in this collection.
This class stores the date and time with a precision of one second in an unsigned 32 bit word (950130...
Definition TDatime.h:37
const char * AsString() const
Return the date & time as a string (ctime() format).
Definition TDatime.cxx:98
TObject * Get(const char *namecycle) override
Return pointer to object identified by namecycle.
TDirectory::TContext keeps track and restore the current directory.
Definition TDirectory.h:89
Describe directory structure in memory.
Definition TDirectory.h:45
virtual TDirectory * GetDirectory(const char *namecycle, Bool_t printError=false, const char *funcname="GetDirectory")
Find a directory using apath.
virtual Bool_t cd()
Change current directory to "this" directory.
virtual TDirectory * mkdir(const char *name, const char *title="", Bool_t returnExistingDirectory=kFALSE)
Create a sub-directory "a" or a hierarchy of sub-directories "a/b/c/...".
A file, usually with extension .root, that stores data and code in the form of serialized objects in ...
Definition TFile.h:130
static TFile * Open(const char *name, Option_t *option="", const char *ftitle="", Int_t compress=ROOT::RCompressionSetting::EDefaults::kUseCompiledDefault, Int_t netopt=0)
Create / open a file.
Definition TFile.cxx:3802
void Close(Option_t *option="") override
Close a file.
Definition TFile.cxx:992
A TGraph is an object made of two arrays X and Y with npoints each.
Definition TGraph.h:41
1-D histogram with a double per channel (see TH1 documentation)
Definition TH1.h:926
1-D histogram with a float per channel (see TH1 documentation)
Definition TH1.h:878
TH1 is the base class of all histogram classes in ROOT.
Definition TH1.h:109
virtual Double_t GetBinCenter(Int_t bin) const
Return bin center for 1D histogram.
Definition TH1.cxx:9371
virtual Double_t GetMean(Int_t axis=1) const
For axis = 1,2 or 3 returns the mean value of the histogram along X,Y or Z axis.
Definition TH1.cxx:7744
virtual void SetXTitle(const char *title)
Definition TH1.h:667
TAxis * GetXaxis()
Definition TH1.h:571
virtual Double_t GetMaximum(Double_t maxval=FLT_MAX) const
Return maximum value smaller than maxval of bins in the range, unless the value has been overridden b...
Definition TH1.cxx:8778
virtual Int_t GetNbinsX() const
Definition TH1.h:541
virtual Int_t Fill(Double_t x)
Increment bin with abscissa X by 1.
Definition TH1.cxx:3489
virtual void SetBinContent(Int_t bin, Double_t content)
Set bin content see convention for numbering bins in TH1::GetBin In case the bin number is greater th...
Definition TH1.cxx:9452
virtual Int_t GetMaximumBin() const
Return location of bin with maximum value in the range.
Definition TH1.cxx:8810
virtual Double_t GetBinContent(Int_t bin) const
Return content of bin number bin.
Definition TH1.cxx:5239
virtual void SetYTitle(const char *title)
Definition TH1.h:668
virtual void Scale(Double_t c1=1, Option_t *option="")
Multiply this histogram by a constant c1.
Definition TH1.cxx:6815
virtual Int_t FindBin(Double_t x, Double_t y=0, Double_t z=0)
Return Global bin number corresponding to x,y,z.
Definition TH1.cxx:3823
virtual Int_t GetQuantiles(Int_t n, Double_t *xp, const Double_t *p=nullptr)
Compute Quantiles for this histogram.
Definition TH1.cxx:4766
virtual void AddBinContent(Int_t bin)=0
Increment bin content by 1.
virtual Double_t KolmogorovTest(const TH1 *h2, Option_t *option="") const
Statistical test of compatibility in shape between this histogram and h2, using Kolmogorov test.
Definition TH1.cxx:8407
virtual void Sumw2(Bool_t flag=kTRUE)
Create structure to store sum of squares of weights.
Definition TH1.cxx:9253
2-D histogram with a float per channel (see TH1 documentation)
Definition TH2.h:345
Int_t Fill(Double_t) override
Invalid Fill method.
Definition TH2.cxx:368
A doubly linked list.
Definition TList.h:38
Class that contains all the information of a class.
Definition ClassInfo.h:49
UInt_t GetNumber() const
Definition ClassInfo.h:65
TString fWeightFileExtension
Definition Config.h:125
VariablePlotting & GetVariablePlotting()
Definition Config.h:97
class TMVA::Config::VariablePlotting fVariablePlotting
IONames & GetIONames()
Definition Config.h:98
MsgLogger * fLogger
! message logger
Class that contains all the data information.
Definition DataSetInfo.h:62
Class that contains all the data information.
Definition DataSet.h:58
Float_t GetValue(UInt_t ivar) const
return value of i'th variable
Definition Event.cxx:236
Double_t GetWeight() const
return the event weight - depending on whether the flag IgnoreNegWeightsInTraining is or not.
Definition Event.cxx:389
static void SetIsTraining(Bool_t)
when this static function is called, it sets the flag whether events with negative event weight shoul...
Definition Event.cxx:399
Float_t GetTarget(UInt_t itgt) const
Definition Event.h:102
static void SetIgnoreNegWeightsInTraining(Bool_t)
when this static function is called, it sets the flag whether events with negative event weight shoul...
Definition Event.cxx:408
Interface for all concrete MVA method implementations.
Definition IMethod.h:53
Virtual base Class for all MVA method.
Definition MethodBase.h:79
TDirectory * MethodBaseDir() const
returns the ROOT directory where all instances of the corresponding MVA method are stored
virtual Double_t GetKSTrainingVsTest(Char_t SorB, TString opt="X")
MethodBase(const TString &jobName, Types::EMVA methodType, const TString &methodTitle, DataSetInfo &dsi, const TString &theOption="")
standard constructor
void PrintHelpMessage() const override
prints out method-specific help method
virtual std::vector< Float_t > GetAllMulticlassValues()
Get all multi-class values.
virtual Double_t GetSeparation(TH1 *, TH1 *) const
compute "separation" defined as
const char * GetName() const override
Definition MethodBase.h:304
void ReadClassesFromXML(void *clsnode)
read number of classes from XML
void SetWeightFileDir(TString fileDir)
set directory of weight file
void WriteStateToXML(void *parent) const
general method used in writing the header of the weight files where the used variables,...
void DeclareBaseOptions()
define the options (their key words) that can be set in the option string here the options valid for ...
virtual void TestRegression(Double_t &bias, Double_t &biasT, Double_t &dev, Double_t &devT, Double_t &rms, Double_t &rmsT, Double_t &mInf, Double_t &mInfT, Double_t &corr, Types::ETreeType type)
calculate <sum-of-deviation-squared> of regression output versus "true" value from test sample
virtual void DeclareCompatibilityOptions()
options that are used ONLY for the READER to ensure backward compatibility they are hence without any...
virtual Double_t GetSignificance() const
compute significance of mean difference
virtual Double_t GetProba(const Event *ev)
virtual TMatrixD GetMulticlassConfusionMatrix(Double_t effB, Types::ETreeType type)
Construct a confusion matrix for a multiclass classifier.
virtual void WriteEvaluationHistosToFile(Types::ETreeType treetype)
writes all MVA evaluation histograms to file
virtual void TestMulticlass()
test multiclass classification
const std::vector< TMVA::Event * > & GetEventCollection(Types::ETreeType type)
returns the event collection (i.e.
virtual std::vector< Double_t > GetDataMvaValues(DataSet *data=nullptr, Long64_t firstEvt=0, Long64_t lastEvt=-1, Bool_t logProgress=false)
get all the MVA values for the events of the given Data type
void SetupMethod()
setup of methods
TDirectory * BaseDir() const
returns the ROOT directory where info/histograms etc of the corresponding MVA method instance are sto...
virtual std::vector< Float_t > GetMulticlassEfficiency(std::vector< std::vector< Float_t > > &purity)
void AddInfoItem(void *gi, const TString &name, const TString &value) const
xml writing
virtual void AddClassifierOutputProb(Types::ETreeType type)
prepare tree branch with the method's discriminating variable
virtual Double_t GetEfficiency(const TString &, Types::ETreeType, Double_t &err)
fill background efficiency (resp.
TString GetTrainingTMVAVersionString() const
calculates the TMVA version string from the training version code on the fly
void Statistics(Types::ETreeType treeType, const TString &theVarName, Double_t &, Double_t &, Double_t &, Double_t &, Double_t &, Double_t &)
calculates rms,mean, xmin, xmax of the event variable this can be either done for the variables as th...
Bool_t GetLine(std::istream &fin, char *buf)
reads one line from the input stream checks for certain keywords and interprets the line if keywords ...
void ProcessSetup()
process all options the "CheckForUnusedOptions" is done in an independent call, since it may be overr...
virtual std::vector< Double_t > GetMvaValues(Long64_t firstEvt=0, Long64_t lastEvt=-1, Bool_t logProgress=false)
get all the MVA values for the events of the current Data type
virtual Bool_t IsSignalLike()
uses a pre-set cut on the MVA output (SetSignalReferenceCut and SetSignalReferenceCutOrientation) for...
virtual ~MethodBase()
destructor
void WriteMonitoringHistosToFile() const override
write special monitoring histograms to file dummy implementation here --------------—
virtual Double_t GetMaximumSignificance(Double_t SignalEvents, Double_t BackgroundEvents, Double_t &optimal_significance_value) const
plot significance, , curve for given number of signal and background events; returns cut for maximum ...
virtual Double_t GetTrainingEfficiency(const TString &)
void SetWeightFileName(TString)
set the weight file name (depreciated)
TString GetWeightFileName() const
retrieve weight file name
virtual void TestClassification()
initialization
void AddOutput(Types::ETreeType type, Types::EAnalysisType analysisType)
virtual void AddRegressionOutput(Types::ETreeType type)
prepare tree branch with the method's discriminating variable
void InitBase()
default initialization called by all constructors
virtual void GetRegressionDeviation(UInt_t tgtNum, Types::ETreeType type, Double_t &stddev, Double_t &stddev90Percent) const
void ReadStateFromXMLString(const char *xmlstr)
for reading from memory
void MakeClass(const TString &classFileName=TString("")) const override
create reader class for method (classification only at present)
void CreateMVAPdfs()
Create PDFs of the MVA output variables.
TString GetTrainingROOTVersionString() const
calculates the ROOT version string from the training version code on the fly
virtual Double_t GetValueForRoot(Double_t)
returns efficiency as function of cut
void ReadStateFromFile()
Function to write options and weights to file.
void WriteVarsToStream(std::ostream &tf, const TString &prefix="") const
write the list of variables (name, min, max) for a given data transformation method to the stream
void ReadVarsFromStream(std::istream &istr)
Read the variables (name, min, max) for a given data transformation method from the stream.
void ReadSpectatorsFromXML(void *specnode)
read spectator info from XML
void ReadVariablesFromXML(void *varnode)
read variable info from XML
virtual std::map< TString, Double_t > OptimizeTuningParameters(TString fomType="ROCIntegral", TString fitType="FitGA")
call the Optimizer with the set of parameters and ranges that are meant to be tuned.
virtual std::vector< Float_t > GetMulticlassTrainingEfficiency(std::vector< std::vector< Float_t > > &purity)
void WriteStateToStream(std::ostream &tf) const
general method used in writing the header of the weight files where the used variables,...
virtual Double_t GetRarity(Double_t mvaVal, Types::ESBType reftype=Types::kBackground) const
compute rarity:
virtual void SetTuneParameters(std::map< TString, Double_t > tuneParameters)
set the tuning parameters according to the argument This is just a dummy .
void ReadStateFromStream(std::istream &tf)
read the header from the weight files of the different MVA methods
void AddVarsXMLTo(void *parent) const
write variable info to XML
Double_t GetMvaValue(Double_t *errLower=nullptr, Double_t *errUpper=nullptr) override=0
void AddTargetsXMLTo(void *parent) const
write target info to XML
void ReadTargetsFromXML(void *tarnode)
read target info from XML
void ProcessBaseOptions()
the option string is decoded, for available options see "DeclareOptions"
void ReadStateFromXML(void *parent)
virtual std::vector< Float_t > GetAllRegressionValues()
Get al regression values in one call.
void NoErrorCalc(Double_t *const err, Double_t *const errUpper)
void WriteStateToFile() const
write options and weights to file note that each one text file for the main configuration information...
void AddClassesXMLTo(void *parent) const
write class info to XML
virtual void AddClassifierOutput(Types::ETreeType type)
prepare tree branch with the method's discriminating variable
void AddSpectatorsXMLTo(void *parent) const
write spectator info to XML
virtual Double_t GetROCIntegral(TH1D *histS, TH1D *histB) const
calculate the area (integral) under the ROC curve as a overall quality measure of the classification
virtual void AddMulticlassOutput(Types::ETreeType type)
prepare tree branch with the method's discriminating variable
virtual void CheckSetup()
check may be overridden by derived class (sometimes, eg, fitters are used which can only be implement...
void SetSource(const std::string &source)
Definition MsgLogger.h:68
PDF wrapper for histograms; uses user-defined spline interpolation.
Definition PDF.h:63
Double_t GetXmin() const
Definition PDF.h:104
Double_t GetXmax() const
Definition PDF.h:105
Double_t GetVal(Double_t x) const
returns value PDF(x)
Definition PDF.cxx:700
@ kSpline3
Definition PDF.h:70
@ kSpline2
Definition PDF.h:70
Double_t GetIntegral(Double_t xmin, Double_t xmax)
computes PDF integral within given ranges
Definition PDF.cxx:653
Class that is the base-class for a vector of result.
std::vector< Float_t > * GetValueVector()
void SetValue(Float_t value, Int_t ievt, Bool_t type)
set MVA response
Class which takes the results of a multiclass classification.
TMatrixD GetConfusionMatrix(Double_t effB)
Returns a confusion matrix where each class is pitted against each other.
Float_t GetAchievablePur(UInt_t cls)
std::vector< Double_t > GetBestMultiClassCuts(UInt_t targetClass)
calculate the best working point (optimal cut values) for the multiclass classifier
void CreateMulticlassHistos(TString prefix, Int_t nbins, Int_t nbins_high)
this function fills the mva response histos for multiclass classification
Float_t GetAchievableEff(UInt_t cls)
void CreateMulticlassPerformanceHistos(TString prefix)
Create performance graphs for this classifier a multiclass setting.
Class that is the base-class for a vector of result.
Class that is the base-class for a vector of result.
Definition Results.h:57
Bool_t DoesExist(const TString &alias) const
Returns true if there is an object stored in the result for a given alias, false otherwise.
Definition Results.cxx:127
void Store(TObject *obj, const char *alias=nullptr)
Definition Results.cxx:86
TH1 * GetHist(const TString &alias) const
Definition Results.cxx:136
TList * GetStorage() const
Definition Results.h:72
Root finding using Brents algorithm (translated from CERNLIB function RZERO)
Definition RootFinder.h:48
Double_t Root(Double_t refValue)
Root finding using Brents algorithm; taken from CERNLIB function RZERO.
Linear interpolation of TGraph.
Definition TSpline1.h:43
Timing information for training and evaluation of MVA methods.
Definition Timer.h:58
Double_t ElapsedSeconds(void)
computes elapsed tim in seconds
Definition Timer.cxx:136
TString GetElapsedTime(Bool_t Scientific=kTRUE)
returns pretty string with elapsed time
Definition Timer.cxx:145
void DrawProgressBar(Int_t, const TString &comment="")
draws progress bar in color or B&W caution:
Definition Timer.cxx:201
void ComputeStat(const std::vector< TMVA::Event * > &, std::vector< Float_t > *, Double_t &, Double_t &, Double_t &, Double_t &, Double_t &, Double_t &, Int_t signalClass, Bool_t norm=kFALSE)
sanity check
Definition Tools.cxx:203
TList * ParseFormatLine(TString theString, const char *sep=":")
Parse the string and cut into labels separated by ":".
Definition Tools.cxx:376
Double_t GetSeparation(TH1 *S, TH1 *B) const
compute "separation" defined as
Definition Tools.cxx:122
Double_t GetMutualInformation(const TH2F &)
Mutual Information method for non-linear correlations estimates in 2D histogram Author: Moritz Backes...
Definition Tools.cxx:564
const TString & Color(const TString &)
human readable color strings
Definition Tools.cxx:803
TXMLEngine & xmlengine()
Definition Tools.h:262
Bool_t CheckSplines(const TH1 *, const TSpline *)
check quality of splining by comparing splines and histograms in each bin
Definition Tools.cxx:454
void ReadAttr(void *node, const char *, T &value)
read attribute from xml
Definition Tools.h:329
void * GetChild(void *parent, const char *childname=nullptr)
get child node
Definition Tools.cxx:1125
void AddAttr(void *node, const char *, const T &value, Int_t precision=16)
add attribute to xml
Definition Tools.h:347
Double_t NormHist(TH1 *theHist, Double_t norm=1.0)
normalises histogram
Definition Tools.cxx:358
void * AddChild(void *parent, const char *childname, const char *content=nullptr, bool isRootNode=false)
add child node
Definition Tools.cxx:1099
void * GetNextChild(void *prevchild, const char *childname=nullptr)
XML helpers.
Definition Tools.cxx:1137
Singleton class for Global types used by TMVA.
Definition Types.h:71
@ kSignal
Never change this number - it is elsewhere assumed to be zero !
Definition Types.h:135
@ kBackground
Definition Types.h:136
@ kLikelihood
Definition Types.h:79
@ kHMatrix
Definition Types.h:81
@ kMulticlass
Definition Types.h:129
@ kNoAnalysisType
Definition Types.h:130
@ kClassification
Definition Types.h:127
@ kMaxAnalysisType
Definition Types.h:131
@ kRegression
Definition Types.h:128
@ kTraining
Definition Types.h:143
Linear interpolation class.
Gaussian Transformation of input variables.
Class for type info of MVA input variable.
void ReadFromXML(void *varnode)
read VariableInfo from stream
const TString & GetExpression() const
char GetVarType() const
void ReadFromStream(std::istream &istr)
read VariableInfo from stream
void AddToXML(void *varnode)
write class to XML
void SetExternalLink(void *p)
void * GetExternalLink() const
void BuildTransformationFromVarInfo(const std::vector< TMVA::VariableInfo > &var)
this method is only used when building a normalization transformation from old text files in this cas...
Linear interpolation class.
Linear interpolation class.
virtual void ReadTransformationFromStream(std::istream &istr, const TString &classname="")=0
const char * GetName() const override
Returns name of object.
Definition TNamed.h:49
Collectable string class.
Definition TObjString.h:28
virtual Int_t Write(const char *name=nullptr, Int_t option=0, Int_t bufsize=0)
Write this object to the current directory.
Definition TObject.cxx:987
Basic string class.
Definition TString.h:137
Ssiz_t Length() const
Definition TString.h:426
void ToLower()
Change string to lower-case.
Definition TString.cxx:1190
Int_t Atoi() const
Return integer value of string.
Definition TString.cxx:2069
Bool_t EndsWith(const char *pat, ECaseCompare cmp=kExact) const
Return true if string ends with the specified string.
Definition TString.cxx:2325
TSubString Strip(EStripType s=kTrailing, char c=' ') const
Return a substring of self stripped at beginning and/or end.
Definition TString.cxx:1171
const char * Data() const
Definition TString.h:385
TString & ReplaceAll(const TString &s1, const TString &s2)
Definition TString.h:714
@ kLeading
Definition TString.h:283
Ssiz_t Last(char c) const
Find last occurrence of a character c.
Definition TString.cxx:939
Bool_t IsNull() const
Definition TString.h:423
static TString Format(const char *fmt,...)
Static method which formats a string using a printf style format descriptor and return a TString.
Definition TString.cxx:2460
Ssiz_t Index(const char *pat, Ssiz_t i=0, ECaseCompare cmp=kExact) const
Definition TString.h:661
virtual const char * GetBuildNode() const
Return the build node name.
Definition TSystem.cxx:3997
virtual int mkdir(const char *name, Bool_t recursive=kFALSE)
Make a file system directory.
Definition TSystem.cxx:921
virtual const char * WorkingDirectory()
Return working directory.
Definition TSystem.cxx:886
virtual UserGroup_t * GetUserInfo(Int_t uid)
Returns all user info in the UserGroup_t structure.
Definition TSystem.cxx:1623
void SaveDoc(XMLDocPointer_t xmldoc, const char *filename, Int_t layout=1)
store document content to file if layout<=0, no any spaces or newlines will be placed between xmlnode...
void FreeDoc(XMLDocPointer_t xmldoc)
frees allocated document data and deletes document itself
XMLNodePointer_t DocGetRootElement(XMLDocPointer_t xmldoc)
returns root node of document
XMLDocPointer_t NewDoc(const char *version="1.0")
creates new xml document with provided version
XMLDocPointer_t ParseFile(const char *filename, Int_t maxbuf=100000)
Parses content of file and tries to produce xml structures.
XMLDocPointer_t ParseString(const char *xmlstring)
parses content of string and tries to produce xml structures
void DocSetRootElement(XMLDocPointer_t xmldoc, XMLNodePointer_t xmlnode)
set main (root) node for document
TLine * line
TH1F * h1
Definition legend1.C:5
Config & gConfig()
Tools & gTools()
void CreateVariableTransforms(const TString &trafoDefinition, TMVA::DataSetInfo &dataInfo, TMVA::TransformationHandler &transformationHandler, TMVA::MsgLogger &log)
MsgLogger & Endl(MsgLogger &ml)
Definition MsgLogger.h:148
Short_t Max(Short_t a, Short_t b)
Returns the largest of a and b.
Definition TMathBase.h:249
Double_t Sqrt(Double_t x)
Returns the square root of x.
Definition TMath.h:675
Short_t Min(Short_t a, Short_t b)
Returns the smallest of a and b.
Definition TMathBase.h:197
Short_t Abs(Short_t d)
Returns the absolute value of parameter Short_t d.
Definition TMathBase.h:122
TString fUser
Definition TSystem.h:149