Logo ROOT  
Reference Guide
 
Loading...
Searching...
No Matches
ROperator_Constant.hxx
Go to the documentation of this file.
1#ifndef TMVA_SOFIE_ROPERATOR_Constant
2#define TMVA_SOFIE_ROPERATOR_Constant
3
5#include "TMVA/ROperator.hxx"
6#include "TMVA/RModel.hxx"
7
8#include <sstream>
9
10namespace TMVA{
11namespace Experimental{
12namespace SOFIE{
13
14template<typename T>
16{
17
18private:
19
20 std::string fNX;
21 std::string fNY;
22 std::vector<size_t> fShape;
23 std::vector<Dim> fDimShape;
24 std::vector<Dim> fDimOutputShape;
25 std::vector<T> fValues;
26 std::string fAttrType;
27 bool fIsConstantOfShape = false;
29
30public:
32
33 ROperator_Constant(const std::string & type, const std::vector<T> & values, const std::vector<size_t> & shape, std::string nameX, std::string nameY):
34 fNX(UTILITY::Clean_name(nameX)),
35 fNY(UTILITY::Clean_name(nameY)),
36 fShape(shape),
37 fValues(values),
39 {
42 }
43
44 void Initialize(RModel& model) override {
45 //input must be a graph input, or already initialized intermediate tensor
46 size_t length = 1;
47 /// ConstantOfShape-------------
48 if (!fNX.empty()) {
49 // case of ConstantOfShape (since no inputs in case of Constant operator)
50 fIsConstantOfShape = true;
51 if (model.CheckIfTensorAlreadyExist(fNX) == false){
52 throw std::runtime_error("TMVA SOFIE ConstantOfShape Op Input Tensor is not found in model");
53 }
54 // get output shape from input values:
55 // can work only if input is a constant or initialized tensor (or dynamic one)
56 if (model.IsConstantTensor(fNX)) {
57 fIsOutputConstant = true;
58 auto dptr = model.GetInitializedTensorData(fNX);
59 auto input_tensor = static_cast<int64_t *>(dptr.get());
60 auto input_shape = model.GetTensorShape(fNX);
61 if (input_shape.size() > 1 )
62 throw std::runtime_error("TMVA SOFIE ConstantOfShape Op Input Tensor has invalid shape");
63 if (input_tensor != nullptr && !input_shape.empty()) {
64 fShape = std::vector<size_t> (input_shape[0]);
65 for (size_t i = 0; i < fShape.size(); i++)
66 fShape[i] = input_tensor[i];
67 } else
68 fShape = {1}; // scalar case
69
71 if (fValues.size() != 1)
72 throw std::runtime_error("TMVA SOFIE ConstantOfShape Op value Tensor has invalid size " + std::to_string(fValues.size()));
73
74 T value = fValues[0];
75 fValues = std::vector<T>(length, value);
76 }
77 else if (model.IsShapeTensor(fNX)) {
78 // case tensor values representing output shapes are known
79 fDimOutputShape = model.GetShapeTensorValues(fNX);
80 } else {
81 // case of not known shape tensors- we need to do at run time
82 // not sure if we ever encounter this case
84 fDimShape = model.GetDimTensorShape(fNX);
85 if (fDimShape.size() > 1 )
86 throw std::runtime_error("TMVA SOFIE ConstantOfShape Op Input Tensor has invalid shape");
87 if (!fDimShape[0].isParam) {
88 fDimOutputShape.resize(fDimShape[0].dim);
89 for (size_t i = 0; i < fDimShape[0].dim; i++) {
90 fDimOutputShape[i] = Dim{ std::string("s_") + fNY + "_" + std::to_string(i)};
91 }
92 }
93 else {
94 throw std::runtime_error("TMVA SOFIE ConstantOfShape Op Input Tensor has not defied shape");
95 }
96 }
97
98 } else {
99 // case of constant operator
100 // in case of standard constant the shape is provided as input
101 fIsOutputConstant = true;
103 if (length != fValues.size())
104 throw std::runtime_error("TMVA SOFIE Constant Op has invalid shape : " + ConvertShapeToString(fShape) +
105 " with " + std::to_string(fValues.size()) + " values");
106 }
107
108 // we need to create an initialized tensor of type constant to flag to not save it in a weight file
109 // but keep its initialization in the generated code. The values might also be needed in initializing the
110 // following operators using as input Constant or ConstantOfShape
111 // resize fValues to shape length
112 if (fIsOutputConstant) {
113 model.AddConstantTensor(fNY, fShape, fValues);
114 if (model.Verbose()) {
115 std::cout << "adding constant tensor " << fNY << " with shape " << ConvertShapeToString(fShape)
116 << " and values [";
117 if (!fIsConstantOfShape) {
118 ConvertValuesToString(fValues, 10); // add maximum printing values
119 } else { // for constant of shape is enough to print one value
120 std::cout << "... " << fValues[0] << " ....]" << std::endl;
121 }
122 }
123 } else {
124 model.AddIntermediateTensor(fNY, ConvertStringToType(TensorType<T>::Name()), fDimOutputShape);
125 fOutputTensorNames.emplace_back(fNY);
126 }
127 }
128
129 std::string Generate(std::string opName) override {
130 // no code to generate here. Tensor are defined in Session constructor
131 std::stringstream out;
132 if (fIsOutputConstant) {
133 if (fNX.empty())
134 out << "// ---- Constant (no-op) " << opName << " --> " << fNY << " " << ConvertDimShapeToString(fDimOutputShape) << "\n";
135 else
136 out << "// ---- ConstantOfShape (no-op) " << opName << " --> " << fNY << " " << ConvertDimShapeToString(fDimOutputShape) << "\n";
137 return out.str();
138 }
139 // Only ConstantOfShape might require generation code
140 // generate constant tensor according to input
141
142 out << "\n//--------- ConstantOfShape " << opName << " --> " << ConvertDimShapeToString(fDimOutputShape) << "\n";
143 // set shape values if needed
145 for (size_t i = 0; i < fDimOutputShape.size(); i++) {
146 out << SP << "size_t " << fDimOutputShape[i].param << " = " << "tensor_" << fNX << "[" << i << "];\n";
147 }
148 }
150 // vector is already allocated- fill with values
151 out << SP << "std::fill(tensor_" << fNY << ", tensor_" << fNY << " + " << length << ", " << fValues[0] << ");\n";
152 return out.str();
153 }
154};
155
156}//SOFIE
157}//Experimental
158}//TMVA
159
160
161#endif //TMVA_SOFIE_ROPERATOR_Constant
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
ROperator_Constant(const std::string &type, const std::vector< T > &values, const std::vector< size_t > &shape, std::string nameX, std::string nameY)
std::string Generate(std::string opName) override
std::vector< std::string_view > fInputTensorNames
Definition ROperator.hxx:44
bool fIsOutputConstant
flag to identify if operator has a constant output (no need to generate code)
Definition ROperator.hxx:41
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::string ConvertDimShapeToString(const std::vector< Dim > &shape)
std::size_t ConvertShapeToLength(const std::vector< size_t > &shape)
std::string ConvertValuesToString(size_t n, const T *data, size_t maxprint=-1)
ETensorType ConvertStringToType(std::string type)
std::string ConvertDimShapeToLength(const std::vector< Dim > &shape)
std::string ConvertShapeToString(const std::vector< size_t > &shape)
create variable transformations