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>
15class ROperator_Identity final : public ROperator
16{
17
18private:
19 bool fIsOutputInitialized = false; // the output is the same weight as the input
20 bool fIsAlias = false; // the output shares the memory of the input
21 std::string fNX;
22 std::string fNY;
23 std::vector<Dim> fShape;
24
25public:
27 ROperator_Identity(std::string nameX, std::string nameY):
28 fNX(UTILITY::Clean_name(nameX)), fNY(UTILITY::Clean_name(nameY)){
31 }
32
33 void Initialize(RModel& model) override {
34 //input must be a graph input, or already initialized intermediate tensor
35 if (model.CheckIfTensorAlreadyExist(fNX) == false){
36 throw std::runtime_error("TMVA SOFIE Identity Op Input Tensor is not found in model");
37 }
38 fShape = model.GetDimTensorShape(fNX);
39 if (model.IsInitializedTensor(fNX)) {
40 // we need to check if is a weight (initialized) or a constant tensor: in both cases the
41 // output is registered directly and no code is generated at run time
42 if (model.IsConstantTensor(fNX)) {
43 auto inputData = static_cast<T*>(model.GetInitializedTensorData(fNX).get());
44 model.AddConstantTensor<T>(fNY, model.GetTensorShape(fNX), inputData);
45 fIsOutputConstant = true;
46 } else {
47 // the output is the same weight under another name (exporters emit this for a
48 // shared parameter); registering it as an initialized tensor keeps it resolvable
49 // while the code is generated, as BatchNormalization needs its scale to be.
50 // Note that the generated code and the weight file then hold the values twice,
51 // once under each name.
53 model.AddInitializedTensor(fNY, model.GetTensorType(fNX), model.GetTensorShape(fNX),
54 model.GetInitializedTensorData(fNX));
55 }
56 } else {
57 model.AddIntermediateTensor(fNY, model.GetTensorType(fNX), fShape);
58 fIsAlias = model.AddAliasTensor(fNY, fNX);
59 }
60 }
61
62 std::string Generate(std::string OpName) override {
64 return "";
65 OpName = "op_" + OpName;
66 if (fShape.empty()) {
67 throw std::runtime_error("TMVA SOFIE Operator Identity called to Generate without being initialized first");
68 }
69 std::stringstream out;
70 out << "\n//------ IDENTITY\n";
71 if (fIsAlias) {
72 out << SP << "auto * tensor_" << fNY << " = tensor_" << fNX << ";\n";
73 } else {
74 out << SP << "std::copy(tensor_" << fNX << ", tensor_" << fNX << " + " << ConvertDimShapeToLength(fShape)
75 << ", tensor_" << fNY << ");\n";
76 }
77 return out.str();
78 }
79
80};
81
82}//SOFIE
83}//Experimental
84}//TMVA
85
86
87#endif //TMVA_SOFIE_ROPERATOR_IDENTITY
ROperator_Identity(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 ConvertDimShapeToLength(const std::vector< Dim > &shape)
create variable transformations