Logo ROOT  
Reference Guide
 
Loading...
Searching...
No Matches
ROperator_BatchNormalization.hxx
Go to the documentation of this file.
1#ifndef TMVA_SOFIE_ROPERATOR_BatchNormalization
2#define TMVA_SOFIE_ROPERATOR_BatchNormalization
3
4#include "SOFIE_common.hxx"
5#include "ROperator.hxx"
6#include "RModel.hxx"
7
8
9#include <cmath>
10#include <sstream>
11
12namespace TMVA{
13namespace Experimental{
14namespace SOFIE{
15
16template <typename T>
18{
19
20private:
21
22 /* Attributes */
23 float fepsilon = 1e-05;
24 float fmomentum = 0.9;
25 std::size_t ftraining_mode = 0;
26
27 std::string fNX;
28 std::string fNScale;
29 std::string fNB;
30 std::string fNMean;
31 std::string fNVar;
32 std::string fNY;
34 std::string fNFusedScale;
35
36 std::vector<Dim> fShapeX;
37 std::vector<Dim> fShapeY;
38
39 std::string fType;
40
41public:
43
44 /* Constructor */
45 ROperator_BatchNormalization( float epsilon, float momentum, std::size_t training_mode,
46 std::string nameX, std::string nameScale, std::string nameB,
47 std::string nameMean, std::string nameVar, std::string nameY, EActivationType activation=EActivationType::UNDEFINED):
48 fepsilon(epsilon), fmomentum(momentum), ftraining_mode(training_mode),
49 fNX(UTILITY::Clean_name(nameX)), fNScale(UTILITY::Clean_name(nameScale)),
50 fNB(UTILITY::Clean_name(nameB)), fNMean(UTILITY::Clean_name(nameMean)),
51 fNVar(UTILITY::Clean_name(nameVar)), fNY(UTILITY::Clean_name(nameY)), fActivation(activation)
52 {
55 fNFusedScale = fNScale + "_fused_inv_std_dev";
56
57 if(std::is_same<T, float>::value){
58 fType = "float";
59 }
60 else{
61 throw
62 std::runtime_error("TMVA SOFIE Encountered unsupported type parsing a BatchNormalization operator");
63 }
64 }
65
66
67 void Initialize(RModel& model) override {
68 if (!model.CheckIfTensorAlreadyExist(fNX)) {
69 throw
70 std::runtime_error("TMVA SOFIE BatchNormalization op Input Tensor " + fNX + " fnx is not found in model");
71 }
72 if (!model.CheckIfTensorAlreadyExist(fNScale)) {
73 throw
74 std::runtime_error("TMVA SOFIE BatchNormalization op Input Tensor " + fNScale + " fns is not found in model");
75 }
76 if (!model.CheckIfTensorAlreadyExist(fNB)) {
77 throw
78 std::runtime_error("TMVA SOFIE BatchNormalization op Input Tensor " + fNB + " fnb is not found in model");
79 }
80 if (!model.CheckIfTensorAlreadyExist(fNMean)) {
81 throw
82 std::runtime_error("TMVA SOFIE BatchNormalization op Input Tensor " + fNMean + " fnm is not found in model");
83 }
84 if (!model.CheckIfTensorAlreadyExist(fNVar)) {
85 throw
86 std::runtime_error("TMVA SOFIE BatchNormalization op Input Tensor " + fNVar + " fnv is not found in model");
87 }
88
89 fShapeX = model.GetDimTensorShape(fNX);
90
91 if (fShapeX.size() < 2 || fShapeX.size() > 4) {
92 throw
93 std::runtime_error("TMVA SOFIE BatchNormalization Op input tensor " + fNX + " fnx has wrong shape : " + ConvertDimShapeToString(fShapeX));
94 }
95
97 model.AddIntermediateTensor(fNY, model.GetTensorType(fNX), fShapeY);
98
99 auto original_S = model.GetInitializedTensorData(fNScale);
100 auto original_V = model.GetInitializedTensorData(fNVar);
101
102 auto shape_S = model.GetTensorShape(fNScale);
103 if (shape_S.size() != 1) {
104 throw std::runtime_error("TMVA SOFIE BatchNormalization 'scale' tensor must be 1D (per-channel).");
105 }
106 size_t channels = shape_S[0];
107
108 if (fType == "float") {
109 float *original_scale_ptr = static_cast<float *>(original_S.get());
110 float *original_var_ptr = static_cast<float *>(original_V.get());
111 float *fused_scale_data = new float[channels];
112
113 for (size_t i = 0; i < channels; i++) {
114 // Calculate scale * (1 / sqrt(variance + epsilon))
116 }
117
118 std::shared_ptr<void> fused_scale_ptr(fused_scale_data, std::default_delete<float[]>());
119 model.AddInitializedTensor(fNFusedScale, model.GetTensorType(fNScale), {channels}, fused_scale_ptr);
120 }
121 }
122
123 std::string Generate(std::string opName) override {
124 opName = "op_" + opName;
125 if (fShapeX.empty()){
126 throw std::runtime_error("TMVA SOFIE Batch Normalization called to Generate without being initialized first");
127 }
128
129 std::stringstream out;
130 //// Batch Norm op
131 auto batchSize = fShapeX[0].GetVal();
132 auto channels = fShapeX[1].GetVal();
133 std::string spatial_dim = "1";
134 if (fShapeX.size() > 2) {
135 auto spatialShape = fShapeX;
138 }
139
140 out << "\n\n//---- BatchNorm" << (fActivation == EActivationType::RELU ? " + ReLU " : " ") << opName << "\n";
141 out << SP << "{\n";
142 out << SP << " size_t i = 0;\n";
143 out << SP << " for (size_t n = 0; n < " << batchSize << "; ++n) {\n";
144 out << SP << " for (size_t c = 0; c < " << channels << "; ++c) {\n";
145 out << SP << " const float mean_val = tensor_" << fNMean << "[c];\n";
146 out << SP << " const float fused_scale_val = tensor_" << fNFusedScale << "[c];\n";
147 out << SP << " const float bias_val = tensor_" << fNB << "[c];\n";
148 out << SP << " for (size_t sp = 0; sp < " << spatial_dim << "; ++sp) {\n";
149 out << SP << " float val = (tensor_" << fNX << "[i] - mean_val) * fused_scale_val + bias_val;\n";
150
152 out << SP << " tensor_" << fNY << "[i] = (val > 0.0f) ? val : 0.0f;\n";
153 } else {
154 out << SP << " tensor_" << fNY << "[i] = val;\n";
155 }
156 out << SP << " i++;\n";
157 out << SP << " }\n";
158 out << SP << " }\n";
159 out << SP << " }\n";
160 out << SP << "}\n";
161
162 return out.str();
163 }
164
165 std::vector<std::string> GetBlasRoutines() override { return {}; }
166};
167
168}//SOFIE
169}//Experimental
170}//TMVA
171
172
173#endif //TMVA_SOFIE_ROPERATOR_BatchNormalization
#define e(i)
Definition RSha256.hxx:103
ROOT::Detail::TRangeCast< T, true > TRangeDynCast
TRangeDynCast is an adapter class that allows the typed iteration through a TCollection.
const_iterator begin() const
ROperator_BatchNormalization(float epsilon, float momentum, std::size_t training_mode, std::string nameX, std::string nameScale, std::string nameB, std::string nameMean, std::string nameVar, std::string nameY, EActivationType activation=EActivationType::UNDEFINED)
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::string ConvertDimShapeToString(const std::vector< Dim > &shape)
std::string ConvertDimShapeToLength(const std::vector< Dim > &shape)
create variable transformations