Logo ROOT  
Reference Guide
 
Loading...
Searching...
No Matches
ROperator_Random.hxx
Go to the documentation of this file.
1#ifndef TMVA_SOFIE_ROPERATOR_Random
2#define TMVA_SOFIE_ROPERATOR_Random
3
5#include "TMVA/ROperator.hxx"
6#include "TMVA/RModel.hxx"
7
8#include <sstream>
9
10namespace TMVA{
11namespace Experimental{
12namespace SOFIE{
13
14 // Random operator for RandomUniform, RandomUniformLike,
15 // RandomNormal, RandomNormalLike
17
19{
20public:
21
22 bool fUseROOT = true; // use ROOT or std for random number generation
23private:
24
27 std::string fNX;
28 std::string fNY;
29 int fSeed;
30 std::vector<size_t> fShapeY;
31 std::map<std::string,float> fParams; // parameter for random generator (e.g. low,high or mean and scale)
32
33
34
35public:
36
38 ROperator_Random(RandomOpMode mode, ETensorType type, const std::string & nameX, const std::string & nameY, const std::vector<size_t> & shape, const std::map<std::string, float> & params, float seed) :
39 fMode(mode),
40 fType(type),
41 fNX(UTILITY::Clean_name(nameX)),
42 fNY(UTILITY::Clean_name(nameY)),
43 fSeed(seed),
44 fShapeY(shape),
45 fParams(params)
46 {
49 }
50
51
52 void Initialize(RModel& model) override {
53
54 model.AddIntermediateTensor(fNY, fType, fShapeY);
55
56 if (fUseROOT) {
57 model.AddNeededCustomHeader("TRandom3.h");
58 }
59
60 // use default values
61 if (fMode == kNormal) {
62 if (fParams.count("mean") == 0 )
63 fParams["mean"] = 0;
64 if (fParams.count("scale") == 0)
65 fParams["scale"] = 1;
66 }
67 if (fMode == kUniform) {
68 if (fParams.count("low") == 0)
69 fParams["low"] = 0;
70 if (fParams.count("high") == 0)
71 fParams["high"] = 1;
72 }
73
74 if (model.Verbose()) {
75 std::cout << "Random";
76 if (fMode == kNormal) std::cout << "Normal";
77 else if (fMode == kUniform) std::cout << "Uniform";
78 std::cout << " op -> " << fNY << " : " << ConvertShapeToString(fShapeY) << std::endl;
79 for (auto & p : fParams)
80 std::cout << p.first << " : " << p.second << std::endl;
81 }
82 }
83 // generate declaration code for random number generators
84 std::string GenerateDeclCode() override {
85 std::stringstream out;
86 out << "std::unique_ptr<TRandom> fRndmEngine; // random number engine\n";
87 return out.str();
88 }
89 // generate initialization code for random number generators
90 std::string GenerateInitCode() override {
91 std::stringstream out;
92 out << "//--- creating random number generator ----\n";
93 if (fUseROOT) {
94 // generate initialization code for creating random number generator
95 out << SP << "fRndmEngine.reset(new TRandom3(" << fSeed << "));\n";
96 }
97 else {
98 // not supported
99 }
100 return out.str();
101 }
102 std::string Generate(std::string OpName) override {
103 OpName = "op_" + OpName;
104
105 std::stringstream out;
106 out << "\n//------ Random";
107 if (fMode == kNormal) out << "Normal\n";
108 else if (fMode == kUniform) out << "Uniform\n";
109
110 // generate the random array
112 out << SP << "for (int i = 0; i < " << length << "; i++) {\n";
113 if (fUseROOT) {
114 if (fMode == kNormal) {
115 if (fParams.count("mean") == 0 || fParams.count("scale") == 0)
116 throw std::runtime_error("TMVA SOFIE RandomNormal op : no mean or scale are defined");
117 float mean = fParams["mean"];
118 float scale = fParams["scale"];
119 out << SP << SP << "tensor_" << fNY << "[i] = this->fRndmEngine->Gaus(" << mean << "," << scale << ");\n";
120 } else if (fMode == kUniform) {
121 if (fParams.count("high") == 0 || fParams.count("low") == 0)
122 throw std::runtime_error("TMVA SOFIE RandomUniform op : no low or high are defined");
123 float high = fParams["high"];
124 float low = fParams["low"];
125 out << SP << SP << "tensor_" << fNY << "[i] = this->fRndmEngine->Uniform(" << low << "," << high << ");\n";
126 }
127 }
128 out << SP << "}\n";
129
130 return out.str();
131 }
132
133 std::vector<std::string> GetStdLibs() override {
134 std::vector<std::string> ret = {"memory"}; // for unique ptr
135 return ret;
136 }
137
138};
139
140}//SOFIE
141}//Experimental
142}//TMVA
143
144
145#endif //TMVA_SOFIE_ROPERATOR_Swish
ROOT::Detail::TRangeCast< T, true > TRangeDynCast
TRangeDynCast is an adapter class that allows the typed iteration through a TCollection.
winID h TVirtualViewer3D TVirtualGLPainter p
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 mode
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
ROperator_Random(RandomOpMode mode, ETensorType type, const std::string &nameX, const std::string &nameY, const std::vector< size_t > &shape, const std::map< std::string, float > &params, float seed)
std::vector< std::string > GetStdLibs() override
std::string Generate(std::string OpName) override
std::vector< std::string_view > fInputTensorNames
Definition ROperator.hxx:44
const std::string SP
space used to correctly indent the generated C++ code
Definition ROperator.hxx:40
std::vector< std::string_view > fOutputTensorNames
Definition ROperator.hxx:45
std::size_t ConvertShapeToLength(const std::vector< size_t > &shape)
std::string ConvertShapeToString(const std::vector< size_t > &shape)
create variable transformations