Belle II Software light-2609-luna
Weightfile.h
1/**************************************************************************
2 * basf2 (Belle II Analysis Software Framework) *
3 * Author: The Belle II Collaboration *
4 * *
5 * See git log for contributors and copyright holders. *
6 * This file is licensed under LGPL-3.0, see LICENSE.md. *
7 **************************************************************************/
8
9#pragma once
10#ifndef INCLUDE_GUARD_BELLE2_MVA_WEIGHTFILE_HEADER
11#define INCLUDE_GUARD_BELLE2_MVA_WEIGHTFILE_HEADER
12
13#include <mva/interface/Options.h>
14
15#include <framework/database/IntervalOfValidity.h>
16#include <framework/dataobjects/EventMetaData.h>
17
18#include <boost/property_tree/ptree.hpp>
19
20#include <vector>
21#include <string>
22#include <fstream>
23
24namespace Belle2 {
29
30 namespace MVA {
31
32 std::string makeSaveForDatabase(std::string str);
33
38 class Weightfile {
39
40 public:
45
50
55 void addFeatureImportance(const std::map<std::string, float>& importance);
56
60 std::map<std::string, float> getFeatureImportance() const;
61
66 void addOptions(const Options& options);
67
72 void getOptions(Options& options) const;
73
78 void addSignalFraction(float signal_fraction);
79
84 float getSignalFraction() const;
85
92 std::string generateFileName(const std::string& suffix = "");
93
99 void addFile(const std::string& identifier, const std::string& custom_weightfile);
100
106 void addStream(const std::string& identifier, std::istream& in);
107
113 template<class T>
114 void addElement(const std::string& identifier, const T& element)
115 {
116 m_pt.put(identifier, element);
117 }
118
124 template<class T>
125 void addVector(const std::string& identifier, const std::vector<T>& vector)
126 {
127 m_pt.put(identifier + "_size", vector.size());
128 for (unsigned int i = 0; i < vector.size(); ++i) {
129 m_pt.put(identifier + std::to_string(i), vector[i]);
130 }
131 }
132
144 void addContractVersion(int version)
145 {
147 }
148
156 {
158 }
159
167 int getContractVersion(int default_value = -1) const
168 {
169 auto version = m_pt.get_optional<int>(c_contractVersion);
170 return version ? *version : default_value;
171 }
172
178 void getFile(const std::string& identifier, const std::string& custom_weightfile);
179
184 std::string getStream(const std::string& identifier) const;
185
190 template<class T>
191 T getElement(const std::string& identifier) const
192 {
193 return m_pt.get<T>(identifier);
194 }
195
200 bool containsElement(const std::string& identifier) const
201 {
202 return m_pt.count(identifier) > 0;
203 }
204
210 template<class T>
211 T getElement(const std::string& identifier, const T& default_value) const
212 {
213 return m_pt.get<T>(identifier, default_value);
214 }
215
220 template<class T>
221 std::vector<T> getVector(const std::string& identifier) const
222 {
223 std::vector<T> vector;
224 vector.resize(m_pt.get<size_t>(identifier + "_size"));
225 for (unsigned int i = 0; i < vector.size(); ++i) {
226 vector[i] = m_pt.get<T>(identifier + std::to_string(i));
227 }
228 return vector;
229 }
230
237 static void save(Weightfile& weightfile, const std::string& filename,
238 const Belle2::IntervalOfValidity& iov = Belle2::IntervalOfValidity(0, 0, -1, -1));
239
245 static void saveToROOTFile(Weightfile& weightfile, const std::string& filename);
246
252 static void saveToXMLFile(Weightfile& weightfile, const std::string& filename);
253
259 static void saveToStream(Weightfile& weightfile, std::ostream& stream);
260
266 static Weightfile load(const std::string& filename, const Belle2::EventMetaData& emd = Belle2::EventMetaData(0, 0, 0));
267
272 static Weightfile loadFromFile(const std::string& filename);
273
278 static Weightfile loadFromROOTFile(const std::string& filename);
279
284 static Weightfile loadFromXMLFile(const std::string& filename);
285
290 static Weightfile loadFromStream(std::istream& stream);
291
298 static void saveToDatabase(Weightfile& weightfile, const std::string& identifier,
299 const Belle2::IntervalOfValidity& iov = Belle2::IntervalOfValidity(0, 0, -1, -1));
300
307 static void saveArrayToDatabase(const std::vector<Weightfile>& weightfiles, const std::string& identifier,
308 const Belle2::IntervalOfValidity& iov = Belle2::IntervalOfValidity(0, 0, -1, -1));
309
315 static Weightfile loadFromDatabase(const std::string& identifier, const Belle2::EventMetaData& emd = Belle2::EventMetaData(0, 0,
316 0));
317
322 void setRemoveTemporaryDirectories(bool remove_temporary_directories) { m_remove_temporary_directories = remove_temporary_directories; }
323
327 const boost::property_tree::ptree& getXMLTree() const { return m_pt; };
328
329 private:
331 static constexpr const char* c_contractVersion = "contract_version";
332
333 boost::property_tree::ptree m_pt;
334 std::vector<std::string> m_filenames;
336 };
337
338 }
340}
341#endif
Store event, run, and experiment numbers.
A class that describes the interval of experiments/runs for which an object in the database is valid.
Abstract base class of all Options given to the MVA interface.
Definition Options.h:34
The Weightfile class serializes all information about a training into an xml tree.
Definition Weightfile.h:38
int getContractVersion(int default_value=-1) const
Returns the contract version of this weightfile, or the default value.
Definition Weightfile.h:167
void addStream(const std::string &identifier, std::istream &in)
Add a stream to our weightfile.
void addElement(const std::string &identifier, const T &element)
Add an element to the xml tree.
Definition Weightfile.h:114
Weightfile()
Construct an empty weightfile.
Definition Weightfile.h:44
T getElement(const std::string &identifier) const
Returns a stored element from the xml tree.
Definition Weightfile.h:191
void addFile(const std::string &identifier, const std::string &custom_weightfile)
Add a file (mostly a weightfile from a MVA library) to our Weightfile.
const boost::property_tree::ptree & getXMLTree() const
Get xml tree.
Definition Weightfile.h:327
std::map< std::string, float > getFeatureImportance() const
Get feature importance.
Definition Weightfile.cc:82
static Weightfile loadFromXMLFile(const std::string &filename)
Static function which loads a Weightfile from a XML file.
bool hasContractVersion() const
Returns true if this weightfile records a contract version.
Definition Weightfile.h:155
static void save(Weightfile &weightfile, const std::string &filename, const Belle2::IntervalOfValidity &iov=Belle2::IntervalOfValidity(0, 0, -1, -1))
Static function which saves a Weightfile to a file.
~Weightfile()
Destructor (removes temporary files associated with this weightfiles)
Definition Weightfile.cc:51
void setRemoveTemporaryDirectories(bool remove_temporary_directories)
Set the deletion behaviour of the weightfile object for temporary directories For debugging it can be...
Definition Weightfile.h:322
bool m_remove_temporary_directories
remove all temporary directories in the destructor of this class
Definition Weightfile.h:335
static void saveToXMLFile(Weightfile &weightfile, const std::string &filename)
Static function which saves a Weightfile to a XML file.
static Weightfile loadFromStream(std::istream &stream)
Static function which deserializes a Weightfile from a stream.
bool containsElement(const std::string &identifier) const
Returns true if given element is stored in the property tree.
Definition Weightfile.h:200
static constexpr const char * c_contractVersion
identifier of the element holding the contract version, shared by all weightfiles
Definition Weightfile.h:331
boost::property_tree::ptree m_pt
xml tree containing all the saved information of this weightfile
Definition Weightfile.h:333
std::vector< T > getVector(const std::string &identifier) const
Returns a stored vector from the xml tree.
Definition Weightfile.h:221
void addContractVersion(int version)
Store the version of the contract between this weightfile and the code applying it.
Definition Weightfile.h:144
void addOptions(const Options &options)
Add an Option object to the xml tree.
Definition Weightfile.cc:61
static Weightfile loadFromROOTFile(const std::string &filename)
Static function which loads a Weightfile from a ROOT file.
void getOptions(Options &options) const
Fills an Option object from the xml tree.
Definition Weightfile.cc:66
static Weightfile load(const std::string &filename, const Belle2::EventMetaData &emd=Belle2::EventMetaData(0, 0, 0))
Static function which loads a Weightfile from a file or from the database.
static Weightfile loadFromDatabase(const std::string &identifier, const Belle2::EventMetaData &emd=Belle2::EventMetaData(0, 0, 0))
Static function which loads a Weightfile from the basf2 condition database.
static void saveToStream(Weightfile &weightfile, std::ostream &stream)
Static function which serializes a Weightfile to a stream.
static Weightfile loadFromFile(const std::string &filename)
Static function which loads a Weightfile from a file.
void addSignalFraction(float signal_fraction)
Saves the signal fraction in the xml tree.
Definition Weightfile.cc:94
std::vector< std::string > m_filenames
generated temporary filenames, which will be removed in the destructor of this class
Definition Weightfile.h:334
void addFeatureImportance(const std::map< std::string, float > &importance)
Add variable importance.
Definition Weightfile.cc:71
static void saveToROOTFile(Weightfile &weightfile, const std::string &filename)
Static function which saves a Weightfile to a ROOT file.
void addVector(const std::string &identifier, const std::vector< T > &vector)
Add a vector to the xml tree.
Definition Weightfile.h:125
T getElement(const std::string &identifier, const T &default_value) const
Returns a stored element from the xml tree.
Definition Weightfile.h:211
float getSignalFraction() const
Loads the signal fraction frm the xml tree.
Definition Weightfile.cc:99
static void saveArrayToDatabase(const std::vector< Weightfile > &weightfiles, const std::string &identifier, const Belle2::IntervalOfValidity &iov=Belle2::IntervalOfValidity(0, 0, -1, -1))
Static function which saves an array of Weightfile objects in the basf2 condition database.
std::string generateFileName(const std::string &suffix="")
Returns a temporary filename with the given suffix.
std::string getStream(const std::string &identifier) const
Returns the content of a stored stream as string.
static void saveToDatabase(Weightfile &weightfile, const std::string &identifier, const Belle2::IntervalOfValidity &iov=Belle2::IntervalOfValidity(0, 0, -1, -1))
Static function which saves a Weightfile in the basf2 condition database.
void getFile(const std::string &identifier, const std::string &custom_weightfile)
Creates a file from our weightfile (mostly this will be a weightfile of an MVA library)
STL class.
Abstract base class for different kinds of events.