Logo ROOT  
Reference Guide
 
Loading...
Searching...
No Matches
ROperator_Identity.hxx
Go to the documentation of this file.
1#ifndef TMVA_SOFIE_ROPERATOR_IDENTITY
2#define TMVA_SOFIE_ROPERATOR_IDENTITY
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 bool fIsOutputInitialized = false; // the output is the same weight as the input
20 std::string fNX;
21 std::string fNY;
22 std::vector<Dim> fShape;
23
24public:
26 ROperator_Identity(std::string nameX, std::string nameY):
27 fNX(UTILITY::Clean_name(nameX)), fNY(UTILITY::Clean_name(nameY)){
30 }
31
32 std::vector<ETensorType> TypeInference(std::vector<ETensorType> input) override {
33 return input;
34 }
35
36 std::vector<std::vector<size_t>> ShapeInference(std::vector<std::vector<size_t>> input) override {
37 auto ret = input; //suggest copy to compiler
38 return ret;
39 }
40
41 void Initialize(RModel& model) override {
42 //input must be a graph input, or already initialized intermediate tensor
43 if (model.CheckIfTensorAlreadyExist(fNX) == false){
44 throw std::runtime_error("TMVA SOFIE Identity Op Input Tensor is not found in model");
45 }
46 fShape = model.GetDimTensorShape(fNX);
47 if (model.IsInitializedTensor(fNX)) {
48 // we need to check if is a weight (initialized) or a constant tensor: in both cases the
49 // output is registered directly and no code is generated at run time
50 if (model.IsConstantTensor(fNX)) {
51 auto inputData = static_cast<T*>(model.GetInitializedTensorData(fNX).get());
52 model.AddConstantTensor<T>(fNY, model.GetTensorShape(fNX), inputData);
53 fIsOutputConstant = true;
54 } else {
55 // the output is the same weight under another name (exporters emit this for a
56 // shared parameter); registering it as an initialized tensor keeps it resolvable
57 // while the code is generated, as BatchNormalization needs its scale to be.
58 // Note that the generated code and the weight file then hold the values twice,
59 // once under each name.
61 model.AddInitializedTensor(fNY, model.GetTensorType(fNX), model.GetTensorShape(fNX),
62 model.GetInitializedTensorData(fNX));
63 }
64 } else {
65 model.AddIntermediateTensor(fNY, model.GetTensorType(fNX), fShape);
66 }
67 }
68
69 std::string Generate(std::string OpName) override {
71 return "";
72 OpName = "op_" + OpName;
73 if (fShape.empty()) {
74 throw std::runtime_error("TMVA SOFIE Operator Identity called to Generate without being initialized first");
75 }
76 std::stringstream out;
77 out << "\n//------ IDENTITY\n";
78 out << SP << "std::copy(tensor_" << fNX << ", tensor_" << fNX << " + " << ConvertDimShapeToLength(fShape)
79 << ", tensor_" << fNY << ");\n";
80 return out.str();
81 }
82
83};
84
85}//SOFIE
86}//Experimental
87}//TMVA
88
89
90#endif //TMVA_SOFIE_ROPERATOR_IDENTITY
ROOT::Detail::TRangeCast< T, true > TRangeDynCast
TRangeDynCast is an adapter class that allows the typed iteration through a TCollection.
Option_t Option_t TPoint TPoint const char GetTextMagnitude GetFillStyle GetLineColor GetLineWidth GetMarkerStyle GetTextAlign GetTextColor GetTextSize void input
std::vector< ETensorType > TypeInference(std::vector< ETensorType > input) override
std::vector< std::vector< size_t > > ShapeInference(std::vector< std::vector< size_t > > input) override
ROperator_Identity(std::string nameX, std::string nameY)
std::string Generate(std::string OpName) override
std::vector< std::string_view > fInputTensorNames
Definition ROperator.hxx:47
bool fIsOutputConstant
flag to identify if operator has a constant output (no need to generate code)
Definition ROperator.hxx:44
const std::string SP
space used to correctly indent the generated C++ code
Definition ROperator.hxx:42
std::vector< std::string_view > fOutputTensorNames
Definition ROperator.hxx:48
std::string ConvertDimShapeToLength(const std::vector< Dim > &shape)
create variable transformations