Logo ROOT  
Reference Guide
 
Loading...
Searching...
No Matches
ROperator_Reshape.hxx
Go to the documentation of this file.
1#ifndef TMVA_SOFIE_ROPERATOR_RESHAPE
2#define TMVA_SOFIE_ROPERATOR_RESHAPE
3
5#include "TMVA/ROperator.hxx"
6#include "TMVA/RModel.hxx"
7
8#include <cassert>
9#include <cctype>
10#include <sstream>
11#include <algorithm>
12
13namespace TMVA{
14namespace Experimental{
15namespace SOFIE{
16
18
19
21{
22
23private:
24
25 bool fVerbose = false;
26 bool fDimInput = false;
27 bool fDynamicShape = false;
28 bool fIsAlias = false; // output shares the memory of the input
29 ReshapeOpMode fOpMode = Reshape; // type of Reshape operator
30
31 int fAllowZero = 0; // (for Reshape) zero in tensor shape makes output shape equal to input tensor shape
32 int fAxis = 1; // (for Flatten)
33
34 std::string fNData; // input data tensor name
35 std::string fNInput2; // reshape or axes tensor name depending on operator
36 std::string fNOutput; // output tensor name
37 std::vector<Dim> fShapeInput; // input shape data
38 std::vector<Dim> fShapeOutput; // output shape data
39 std::vector<Dim> fOutputShapeData; // in case output is a shape tensor we store here the shape value data (can be parametric)
40 std::vector<int64_t> fAttrAxes; // axes attributes (provided for all version of Squeeze/Unsqueeze)
41 std::vector<int64_t> fShape; // shape tensor values provided for Reshape for int shapes4
42
43public:
44
45 std::string Name() const {
46 if (fOpMode == Reshape) return "Reshape";
47 if (fOpMode == Flatten) return "Flatten";
48 if (fOpMode == Squeeze) return "Squeeze";
49 if (fOpMode == Unsqueeze) return "Unsqueeze";
50 return "";
51 }
52
54 ROperator_Reshape(ReshapeOpMode opMode, int attr_value, std::string nameData, std::string nameInput2, std::string nameOutput)
55 : fOpMode(opMode), fNData(UTILITY::Clean_name(nameData)), fNInput2(UTILITY::Clean_name(nameInput2)),
56 fNOutput(UTILITY::Clean_name(nameOutput))
57 {
60
62 if(!fNInput2.empty()){
63 fInputTensorNames.emplace_back(fNInput2);
64 }
66 }
67
68 // for squeeze/unsqueezed operators following old ONNX version (< 10)
69 // In this cases axes are passed as attribute values
70 ROperator_Reshape(ReshapeOpMode opMode, std::vector<int64_t> attrAxes, std::string nameData, std::string nameOutput)
71 : fOpMode(opMode), fNData(UTILITY::Clean_name(nameData)), fNOutput(UTILITY::Clean_name(nameOutput)),
73 {
77 }
78
79
80 // output shape
81 std::vector<Dim> DoShapeInference(const std::vector<Dim> & input_shape, const std::vector<Dim> & target_shape) {
82 if (fOpMode == Reshape) {
83 // correct the provided shape (here we have the value) for 0 or -1
84 // the target_shape can be a scalar in case of not present shape input tensor
85 std::vector<Dim> output_shape = target_shape;
86 bool hasMinusOne = false;
87 bool hasZero = false;
88 for (size_t i = 0; i < output_shape.size(); i++) {
89 // case for zero values in given shape: in this case we take the corresponding value from input shape
90 if (!output_shape[i].isParam) {
91 if (output_shape[i].dim == 0) {
92 hasZero = true;
93 if (fAllowZero)
94 output_shape[i] = Dim{0};
95 else {
96 if (i > 0 && output_shape.size() != input_shape.size())
97 std::cout << "WARNING: TMVA Reshape Op : output shape has zero value at index " << i <<
98 " but input shape has a different rank than output shape" << std::endl;
99 if (i >= input_shape.size())
100 throw std::runtime_error("TMVA Reshape Op : output shape has zero value at index " + std::to_string(i) +
101 " but input shape does not have corresponding index");
102 }
104 } else if (output_shape[i].dim == static_cast<size_t>(-1)) {
105 hasMinusOne = true;
106 }
107 }
108 }
109 if (hasZero && hasMinusOne) {
110 throw std::runtime_error("TMVA Reshape Op : zero value in shape is not allowed when there is also a -1 in shape");
111 }
112 // now case of -1 in shape - we can infer the value of -1 from all other values
113 for (size_t i = 0; i < output_shape.size(); i++) {
114 if (output_shape[i] == static_cast<size_t>(-1) && !output_shape[i].isParam) {
115 auto tmp = output_shape;
116 tmp.erase(tmp.begin() + i); // erase -1 value to compute the length of the other dimensions
119 if (fVerbose)
120 std::cout << "reshape- try simplifying " << ConvertDimShapeToString(input_shape) << " with length "
121 << input_length << " to " << tmp_length << std::endl;
122
124 output_shape[i] = Dim{static_cast<size_t>(std::stoi(input_length) / std::stoi(tmp_length))};
125 else if (IsInteger(tmp_length) && std::stoi(tmp_length) == 1) {
126 output_shape[i] = Dim{input_length, static_cast<size_t>(-1)};
127 }
128 else {
129 //we can try simplifying expression if tmp_length is integer and part of input_length
130 // contains tmp_length
131 bool canSimplify = false;
132 std::vector <Dim> reduced_input;
133 if (IsInteger(tmp_length)) {
134
135 // try to tokenize with * the input length
136
137 std::stringstream ss(input_length);
138
139 std::string token;
140
141 // Tokenizing w.r.t. space '*'
142 while(getline(ss, token, '*'))
143 {
144 // remove any whitespace
145 token.erase(std::remove_if(token.begin(), token.end(),
146 [](unsigned char x) { return std::isspace(x); }), token.end());
147 if (token != tmp_length) {
148 if (IsInteger(token)) {
149 size_t il = static_cast<size_t>(std::stoi(input_length));
150 size_t tl = static_cast<size_t>(std::stoi(tmp_length));
151 if ((il % tl) == 0) {
152 canSimplify = true;
153 reduced_input.push_back(Dim{il / tl});
154 }
155 } else {
156 reduced_input.push_back(Dim{token});
157 }
158 } else {
159 // token is equal to tmp_length, can be not considered and is simplified
160 canSimplify = true;
161 }
162 }
163 }
164 if (canSimplify) {
165 // if length contains * we need to add some brackets
167 if (res_shape.find('*') != std::string::npos)
168 output_shape[i] = Dim{std::string("(") + res_shape + ")", static_cast<size_t>(-1)};
169 else
171 }
172 if (!canSimplify)
173 output_shape[i] = Dim{std::string("(") + input_length + " / (" + tmp_length + "))", static_cast<size_t>(-1)};
174 }
175
176 break; // cannot have more than -1
177 }
178 // throw std::runtime_error(
179 // "TMVA Reshape Op : output shape has multiple negative or zero values");
180 }
181
182 if (fVerbose)
183 std::cout << "Reshape: correct output shape to " << ConvertDimShapeToString(output_shape) << std::endl;
184
186 throw std::runtime_error("TMVA Reshape Op : Invalid shapes : " + ConvertDimShapeToString(input_shape) +
188 }
189 return output_shape;
190
191 } else if (fOpMode == Flatten) {
192 // flatten case
193 if (fAxis < 0)
194 fAxis += input_shape.size();
195 auto s1 = std::vector<Dim>(input_shape.begin(), input_shape.begin() + fAxis);
196 auto s2 = std::vector<Dim>(input_shape.begin() + fAxis, input_shape.end());
199 std::vector<Dim> newShape = {Dim{l1}, Dim{l2}};
200 return newShape;
201 } else if (fOpMode == Squeeze) {
202 // squeeze
203 // assume no axis is provided - remove all axes with value equal to 1
205 if (fAttrAxes.empty()) {
206 size_t i = 0;
207 while (i < output_shape.size()) {
208 if (output_shape[i] == Dim{1}) {
209 output_shape.erase(output_shape.begin() + i);
210 } else {
211 i++;
212 }
213 }
214 } else {
215 auto axes = fAttrAxes;
216 for (size_t i = 0; i < axes.size(); i++) {
217 if (axes[i] < 0)
218 axes[i] += input_shape.size();
219 if (!(output_shape[axes[i]] == Dim{1}))
220 throw std::runtime_error("TMVA Squeeze Op : Invalid axis value " + std::to_string(axes[i]) +
222 }
223 // for calling vector::erase we must sort axes in decreasing order to avoid
224 std::sort(axes.begin(), axes.end(), std::greater<int>());
225 for (auto & axis : axes) {
226 output_shape.erase(output_shape.begin() + axis);
227 }
228 }
229 return output_shape;
230 }
231 else if (fOpMode == Unsqueeze) {
232 // unsqueeze
233 assert(!fAttrAxes.empty());
235 auto &axes = fAttrAxes;
236 // output rank
237 int64_t r = input_shape.size() + axes.size();
238 for (auto &a : axes) {
239 int64_t i = static_cast<int64_t>(a);
240 if (i < -r || i > r - 1)
241 throw std::runtime_error("TMVA Unsqueeze Op - axes input is not in correct range");
242 if (i >= 0)
243 output_shape.insert(output_shape.begin() + i, Dim{1});
244 else
245 // negative axes
246 output_shape.insert(output_shape.end() + i + 1, Dim{1});
247 }
248 return output_shape;
249 }
250 throw std::runtime_error("TMVA Reshape Op : Invalid ReshapeOpMode");
251 return {Dim{}};
252 }
253
254 void Initialize(RModel& model) override {
255
256 fVerbose = model.Verbose();
257 if (fVerbose)
258 std::cout << "initialize reshape op type " << fOpMode << " - for input " << fNData
259 << " to shape given by " << fNInput2 << std::endl;
260
261 if (model.CheckIfTensorAlreadyExist(fNData) == false) {
262 // input must be a graph input, or already initialized intermediate tensor
263 throw std::runtime_error("TMVA Reshape Op Input Tensor " + fNData + " is not found in model");
264 }
265 fShapeInput = model.GetDimTensorShape(fNData);
266 fDimInput = model.IsDynamicTensor(fNData);
267 // check if optional tensor exists defining shape or axes
268 if (!fNInput2.empty()) {
269 if (model.CheckIfTensorAlreadyExist(fNInput2)) {
270 if (model.IsInitializedTensor(fNInput2)) {
271 // assume input shape is an initialized tensor
272 auto dptr = model.GetInitializedTensorData(fNInput2);
273 auto values = static_cast<int64_t *>(dptr.get());
274 auto vec = model.GetTensorShape(fNInput2);
275 size_t n = 1;
276 if (vec.size() > 0)
277 n = vec[0]; // size of shape input tensor
278 // copy values in fShape vector or fAttrAxes
279 if (fOpMode == Reshape)
280 fShape = std::vector<int64_t>(values, values + n);
281 else
282 fAttrAxes = std::vector<int64_t>(values, values + n);
283
284 std::vector<Dim> targetShape(fShape.begin(),fShape.end());
286 // set flag to not write tensor in weight file. Its data will be hard-coded in way model is constructed
287 model.SetNotWritableInitializedTensor(fNInput2);
288 } else if (model.IsShapeTensor(fNInput2)) {
289 auto shapeData = model.GetShapeTensorValues(fNInput2);
291 if (model.Verbose())
292 std::cout << "Reshape op - get output shape from shape tensor " << fNInput2 << " with value " << ConvertDimShapeToString(shapeData) << std::endl;
293 } else {
294 // we cannot get shape at initialization time but at run-time
295 fDynamicShape = true;
296 // size of shape output us given by size of shape input tensor
297 if (model.IsDynamicTensor(fNInput2)) {
298 throw std::runtime_error("TMVA Reshape Op 2nd input Tensor " + fNInput2 + " cannot have dynamic shape");
299 }
300 auto shapeInput2 = model.GetTensorShape(fNInput2);
301 fShapeOutput.resize(shapeInput2[0]);
302 for (size_t i = 0; i < fShapeOutput.size(); i++) {
303 fShapeOutput[i] = Dim{ std::string("s_") + fNOutput + "_" + std::to_string(i)};
304 }
305 }
306 } else {
307 throw std::runtime_error("TMVA Reshape Op 2nd input Tensor " + fNInput2 + " is not found in model");
308 }
309 } else if (!fAttrAxes.empty()) {
310 // case fNShape is empty and axes are provided as attributes (e.g. for Unsqueeze)
311 fShapeOutput = DoShapeInference(fShapeInput, std::vector<Dim>{});
312 } else if (fOpMode == Flatten || fOpMode == Squeeze) {
313 fShapeOutput = DoShapeInference(fShapeInput, std::vector<Dim>{});
314 } else {
315 throw std::runtime_error("TMVA Reshape Op : Invalid Input/Attribute data");
316 }
317 // check if output is constant or not
318 if (model.IsInitializedTensor(fNData) && model.GetTensorType(fNData) == ETensorType::INT64) {
319 fIsOutputConstant = true;
320 auto inputData = static_cast<int64_t*>(model.GetInitializedTensorData(fNData).get());
323 throw std::runtime_error("TMVA Reshape Op : Invalid Input/Output lengths");
324 model.AddConstantTensor<int64_t>(fNOutput, o_shape, inputData);
325 if (model.Verbose()) {
326 std::cout << Name() << " : " << fNData << " " << ConvertDimShapeToString(fShapeInput) << " --> " << fNOutput << " (constant) " << ConvertDimShapeToString(fShapeOutput) << " : " <<
328 }
329 }
330 // for input shape tensors we can have it if output shape is size==1 or a scalar
331 else if (model.IsShapeTensor(fNData) && fShapeOutput.size() <=1) {
332 // not sure if we ever end-up here - maybe reshaping from scalar to vector or viceversa
333 fIsOutputParamShape = true;
334 fOutputShapeData = model.GetShapeTensorValues(fNData);
335 // pass the rank through the scalar flag: a shape tensor stores only its values,
336 // so a rank-0 output would otherwise read back as rank 1
337 model.AddShapeTensor(fNOutput, fOutputShapeData, fShapeOutput.empty());
338 if (model.Verbose()) {
339 std::cout << Name() << " : " << fNData << " " << ConvertDimShapeToString(fShapeInput) << " --> " << fNOutput << " (shape) " << ConvertDimShapeToString(fShapeOutput) << " : " <<
341 }
342 }
343 else {
344 // non-constant case
345 model.AddIntermediateTensor(fNOutput, model.GetTensorType(fNData), fShapeOutput);
346 // the data are not changed, so the output can share the memory of the input
347 fIsAlias = model.AddAliasTensor(fNOutput, fNData);
348 if (model.Verbose())
349 std::cout << Name() << " : " << fNData << " " << ConvertDimShapeToString(fShapeInput) << " --> "
350 << fNOutput << " " << ConvertDimShapeToString(fShapeOutput) << (fIsAlias ? " (alias)" : "")
351 << std::endl;
352 }
353 }
354
355 std::string Generate(std::string opName) override {
356
357
358 std::stringstream out;
359 std::string opType = "Reshape";
360 if (fOpMode == Flatten)
361 opType = "Flatten";
362 else if (fOpMode == Squeeze)
363 opType = "Squeeze";
364 else if (fOpMode == Unsqueeze)
365 opType = "Unsquueze";
366
367 out << SP << "///--------" << opType << " operator " << opName << " --> " << ConvertDimShapeToString(fShapeOutput) << "\n";
368
369 if (fIsOutputConstant) return out.str(); //no op for constant tensors
370
372 // no code to generate here for param shape output. Tensor output is defined in Session constructor
373 out << "//----------------output is a shape tensor----------\n";
374 for (int i = 0; i < static_cast<int>(fShapeOutput[0].dim); i++) {
375 out << SP << "tensor_" << fNOutput << "[" << i << " ] = " << fOutputShapeData[i].GetVal() << ";\n";
376 }
377 return out.str();
378 }
379
380 // in case of dynamic output shape we need to set the shape value from input shape tensor
381 // and take case of the zero values
382 if (fDynamicShape) {
383 for (size_t i = 0; i < fShapeOutput.size(); i++) {
384 // since fNInput2 values are int64_t, should we check if they are negative?
385 out << SP << "size_t " << fShapeOutput[i].param << " = " << "tensor_" << fNInput2 << "[" << i << "];\n";
386 if (!fAllowZero)
387 out << SP << "if (tensor_" << fNInput2 << "[" << i << "] <= 0 ) "
388 << fShapeOutput[i].param << " = " << fShapeInput[i] << ";\n";
389 }
390 }
391
392 // output of reshape is same as input
395 if (lengthOut != lengthIn) {
396 // check needs to be done at run-time
397 out << SP << "if (" << lengthOut << "!=" << lengthIn << ")\n";
398 out << SP << SP << "throw std::runtime_error(\"TMVA SOFIE Reshape " << opName << " output length "
399 << lengthOut << " is different than input one " << lengthIn << "\");\n";
400 }
401
402 if (fIsAlias) {
403 out << SP << "auto * tensor_" << fNOutput << " = tensor_" << fNData << ";\n";
404 } else {
405 out << SP << "std::copy( tensor_" << fNData << ", tensor_" << fNData << " + " << lengthIn << ", " << "tensor_"
406 << fNOutput << ");\n";
407 }
408 return out.str();
409 }
410};
411
412}//SOFIE
413}//Experimental
414}//TMVA
415
416
417#endif //TMVA_SOFIE_ROPERATOR_RESHAPE
#define a(i)
Definition RSha256.hxx:99
#define s1(x)
Definition RSha256.hxx:91
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 char Point_t Rectangle_t WindowAttributes_t Float_t r
const_iterator begin() const
const_iterator end() const
ROperator_Reshape(ReshapeOpMode opMode, std::vector< int64_t > attrAxes, std::string nameData, std::string nameOutput)
ROperator_Reshape(ReshapeOpMode opMode, int attr_value, std::string nameData, std::string nameInput2, std::string nameOutput)
std::vector< Dim > DoShapeInference(const std::vector< Dim > &input_shape, const std::vector< Dim > &target_shape)
std::string Generate(std::string opName) override
std::vector< std::string_view > fInputTensorNames
Definition ROperator.hxx:44
bool fIsOutputParamShape
flag to identify of the output represents a parametric shape (can be known at compile time)
Definition ROperator.hxx:42
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
Double_t x[n]
Definition legend1.C:17
const Int_t n
Definition legend1.C:16
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)
std::vector< size_t > ConvertShapeToInt(const std::vector< Dim > &shape)
Convert shape based on Dim to integer format.
std::string ConvertDimShapeToLength(const std::vector< Dim > &shape)
bool IsInteger(const std::string &s)
create variable transformations