Logo ROOT  
Reference Guide
 
Loading...
Searching...
No Matches
ROperator_Gemm.hxx
Go to the documentation of this file.
1#ifndef TMVA_SOFIE_ROPERATOR_GEMM
2#define TMVA_SOFIE_ROPERATOR_GEMM
3
4
6#include "TMVA/ROperator.hxx"
7#include "TMVA/RModel.hxx"
8
9#include <sstream>
10#include <algorithm>
11#include <iterator>
12#include <iomanip>
13#include <limits>
14#include <cassert>
15
16namespace TMVA{
17namespace Experimental{
18namespace SOFIE{
19
20
21 template <typename T>
23 {
24
25 private:
26 bool fIsDynamic = false;
27 bool fBroadcastBias = false;
28 bool fCheckBiasShapeAtRuntime = false; // flag to identify the need to do a run time check of bias shape compatibility in case of dynamic shapes and uni-directional broadcasting
29 bool fBiasBroadcastAssumed = false; // Initialize assumed a broadcast: the integer shape of Y was unknown
30
31 float fAttrAlpha = 1.0;
32 float fAttrBeta = 1.0;
35
36 std::string fNA;
37 std::string fNB;
38 std::string fNC = "";
39 std::string fNY;
40 std::string fType;
42 std::vector<Dim> fShapeA;
43 std::vector<Dim> fShapeB;
44 std::vector<size_t> fShapeC;
45 std::vector<Dim> fDimShapeC;
46 std::vector<Dim> fShapeY;
47 RModel * fModel = nullptr;
48
49 public:
50
52 ROperator_Gemm(float alpha, float beta, int_t transA, int_t transB, std::string nameA, std::string nameB, std::string nameY, EActivationType activation=EActivationType::UNDEFINED):
53 fAttrAlpha(alpha), fAttrBeta(beta), fAttrTransA(transA), fAttrTransB(transB), fNA(UTILITY::Clean_name(nameA)),
54 fNB(UTILITY::Clean_name(nameB)), fNY(UTILITY::Clean_name(nameY))
55 {
57 fType = "float";
58 static_assert(std::is_same_v<T, float>,
59 "TMVA::SOFIE - Unsupported type parsing a Gemm operator");
62 }
63
64 ROperator_Gemm(float alpha, float beta, int_t transA, int_t transB, std::string nameA, std::string nameB, std::string nameC, std::string nameY, EActivationType activation=EActivationType::UNDEFINED):
65 fAttrAlpha(alpha), fAttrBeta(beta), fAttrTransA(transA), fAttrTransB(transB), fNA(UTILITY::Clean_name(nameA)),
66 fNB(UTILITY::Clean_name(nameB)), fNC(UTILITY::Clean_name(nameC)), fNY(UTILITY::Clean_name(nameY)), fActivation(activation)
67 {
69 fType = "float";
70
73 }
74
75 template <typename U>
76 std::vector<U> DoShapeInference(const std::vector<std::vector<U>> & input){
77 if (input.size() > 3) throw std::runtime_error("TMVA SOFIE Gemm Op Shape Inference only need 2 or 3 input tensor");
78 // accept tensor with input dimensions > 2
79 // example: A = (d1,d2,...,N1,N2) B = (d1,d2,...,N2,N3) --> Y = (d1,d2,..,N1,N3)
80 for (auto& i: input){
81 if (i.size() < 2){
82 throw std::runtime_error("TMVA SOFIE Gemm Op Shape Inference only accept input tensor with >=2 dimensions");
83 }
84 }
85
86 // when there are 3 inputs shape of Y is the one of C
87 if (input.size() == 3){
88 //shape of C is shape of Y
89 return input[2];
90 }
91 // ioffset cannot be less than 2
92 int ioffset = input[0].size()-2; // in case of tensors with dim > 2
93
94 std::vector<U> s_a(input[0].begin() + ioffset, input[0].begin() + ioffset + 2);
95 std::vector<U> s_b(input[1].begin() + ioffset, input[1].begin() + ioffset + 2);
96 // reverse in case of transpose
97 if (fAttrTransA){
98 std::reverse(s_a.begin(), s_a.end());
99 }
100 if (fAttrTransB){
101 std::reverse(s_b.begin(), s_b.end());
102 }
103 std::vector<U> s_y;
104 s_y.reserve(input[0].size());
105 if (input[0].size() > 2 && input[1].size() == input[0].size()) {
106 // in case of dim > 2 first dimensions are equal to the input ones not
107 // equal to 1 (e.g. (1,2,3) * (2,3,4) -> (2,2,4))
108 // here could probably use the Broadcasting function UTILITY::MultidirectionalBroadcastShape
109 for (size_t i = 0; i < input[0].size()-2; i++) {
110 Dim valueA = input[0][i];
111 Dim valueB = input[1][i];
112 if (valueA.GetVal() != valueB.GetVal()) {
113 if (valueB.GetVal() == "1")
114 s_y.push_back(input[0][i]);
115 else if (valueA.GetVal() == "1")
116 s_y.push_back(input[1][i]);
117 else if (!valueA.isParam && !valueB.isParam)
118 throw std::runtime_error("TMVA SOFIE Gemm Op - invalid input shapes " + valueA.GetVal() + " and "
119 + valueB.GetVal());
120 else if (valueA.isParam && valueB.isParam){
121 // check which parameter is first in RModel list
122 auto & dimNames = fModel->GetDimShapeNames();
123 auto p1 = std::find(dimNames.begin(), dimNames.end(), valueA.param);
124 auto p2 = std::find(dimNames.begin(), dimNames.end(), valueB.param);
125 if (p1 < p2) s_y.push_back(input[0][i]);
126 else s_y.push_back(input[1][i]);
127 }
128 else if (!valueA.isParam)
129 s_y.push_back(input[0][i]);
130 else if (!valueB.isParam)
131 s_y.push_back(input[1][i]);
132 else
133 throw std::runtime_error("TMVA SOFIE Gemm Op - invalid input shapes " + valueA.GetVal() + " and "
134 + valueB.GetVal());
135 }
136 else
137 s_y.push_back(input[0][i]);
138 }
139 }
140
141 s_y.push_back(s_a[0]);
142 s_y.push_back(s_b[1]);
143 return s_y;
144 }
145
146 std::vector<Dim> DynamicShapeInference(const std::vector<std::vector<Dim>> & input){
148 }
149
150
151
152 void Initialize(RModel& model) override {
153 //TODO: propagate A or B as specified by ONNX standard
154 fModel = &model;
155
156 if ((model.CheckIfTensorAlreadyExist(fNA) == false) || (model.CheckIfTensorAlreadyExist(fNB) == false) ){ //input must be a graph input, or already initialized intermediate tensor
157 throw std::runtime_error("TMVA SOFIE Gemm Op Input Tensor " + fNA + " or " + fNB + " is not found in model");
158 }
159 if (fNC != ""){
160 if (model.CheckIfTensorAlreadyExist(fNC) == false){ //input must be a graph input, or already initialized intermediate tensor
161 throw std::runtime_error("TMVA SOFIE Gemm Op Input Tensor " + fNC + " is not found in model");
162 }
163 }
164 if (model.IsDynamicTensor(fNA) || model.IsDimInputTensor(fNA) ) {
165 fShapeA = model.GetDynamicTensorShape(fNA);
166 fIsDynamic = true;
167 } else {
168 auto shapeA_int = model.GetTensorShape(fNA);
170 }
171 // case A is of dim1 we prepend a 1 but we need to remove later
172 bool prependOne = false;
173 if (fShapeA.size() == 1) {
174 fShapeA.insert(fShapeA.begin(), Dim(1));
175 prependOne = true;
176 }
177
178 if (model.IsDynamicTensor(fNB) || model.IsDimInputTensor(fNB)) {
179 fShapeB = model.GetDynamicTensorShape(fNB);
180 fIsDynamic = true;
181 }
182 else {
183 auto shapeB_int = model.GetTensorShape(fNB);
185 }
186 // case B is dim1 we append a 1 but we need to remove later
187 bool appendOne = false;
188 if (fShapeB.size() == 1) {
189 fShapeB.insert(fShapeB.end(), Dim(1));
190 appendOne = true;
191 }
192 // assume if not shape is 2 that extra values are 1.
193 // implement also MatMul case where we stack matrices (see numpy.matmul)
194 if (fShapeA.size() != fShapeB.size()) {
195 // if different dimensions we prepend 1 values
196 if (fShapeA.size() < fShapeB.size()) {
197 fShapeA.insert(fShapeA.begin(), fShapeB.size()-fShapeA.size(), Dim(1));
198 } else if (fShapeB.size() < fShapeA.size()) {
199 fShapeB.insert(fShapeB.begin(), fShapeA.size()-fShapeB.size(), Dim(1));
200 }
201 }
202
204 std::vector<size_t> shapeY = ConvertShapeToInt(fShapeY);
205
206 // bias is normally not dynamic (not support it for time being)
207 if (fNC != ""){
208 if (model.IsDynamicTensor(fNC))
209 fDimShapeC = model.GetDynamicTensorShape(fNC);
210 else {
211 fShapeC = model.GetTensorShape(fNC);
213 }
214 // for dynamic outputs broadcasting is always needed
215 bool broadcast_needed = false;
216 if (fIsDynamic && shapeY.empty()) {
217 broadcast_needed = true;
219 } else
220 // consider broadcasting also if they have different length
222
223
224 if (broadcast_needed) {
225 fBroadcastBias = true;
226 // check if broadcasting is compatible and note that prepend 1 to shapeC
228 // return flag must not have bit equal to 2 since this is a unidirectional broadcast of C->Y
229 //
230 if ((r.first & 2) == 2) {
231 throw std::runtime_error("TMVA SOFIE Gemm Op - bias tensor of shape " + ConvertDimShapeToString(fDimShapeC) + " cannot be uni-directional broadcasted to " + ConvertDimShapeToString(fShapeY));
232 } else if (r.first == 4) {
233 // we need to do a run time check of bias shape if it is compatible
235 }
237 }
238 }
239
240 // remove appended or prepended value of 1 in Y
241 if (prependOne) {
242 if (fIsDynamic)
243 fShapeY.erase(fShapeY.begin());
244 else
245 shapeY.erase(shapeY.begin());
246 }
247 if (appendOne) {
248 if (fIsDynamic)
249 fShapeY.erase(fShapeY.end()-1);
250 else
251 shapeY.erase(shapeY.end()-1);
252 }
253
254 // Constant-fold Gemm/MatMul when A, B (and C) are all initializers, following the
255 // ROperator_BasicBinary pattern (compute now, skip Generate() entirely). Only full
256 // constant folding is handled; propagating just A or B (see the TODO above) would
257 // need a different mechanism than fIsOutputConstant's all-or-nothing fold.
258 bool canFold = !fIsDynamic
259 && model.IsInitializedTensor(fNA)
260 && model.IsInitializedTensor(fNB)
261 && (fNC.empty() || model.IsInitializedTensor(fNC))
262 && fShapeA.size() <= 2 // exclude stacked/batched MatMul
263 && !fBroadcastBias // exclude bias requiring run-time broadcast
265
266 if (canFold) {
269 size_t dimA = shapeA_i.size();
270 size_t dimB = shapeB_i.size();
271 size_t m = fAttrTransA ? shapeA_i[dimA - 1] : shapeA_i[dimA - 2];
272 size_t k = fAttrTransA ? shapeA_i[dimA - 2] : shapeA_i[dimA - 1];
273 size_t n = fAttrTransB ? shapeB_i[dimB - 2] : shapeB_i[dimB - 1];
274
275 auto dataA = static_cast<T *>(model.GetInitializedTensorData(fNA).get());
276 auto dataB = static_cast<T *>(model.GetInitializedTensorData(fNB).get());
277
278 // plain host-side 2D matrix multiply: Y = alpha * op(A) * op(B)
279 std::vector<T> dataY(m * n, T(0));
280 for (size_t i = 0; i < m; i++) {
281 for (size_t j = 0; j < n; j++) {
282 T sum{};
283 for (size_t p = 0; p < k; p++) {
284 T aVal = fAttrTransA ? dataA[p * m + i] : dataA[i * k + p];
285 T bVal = fAttrTransB ? dataB[j * k + p] : dataB[p * n + j];
286 sum += aVal * bVal;
287 }
288 dataY[i * n + j] = static_cast<T>(fAttrAlpha) * sum;
289 }
290 }
291 // Y += beta * C (fBroadcastBias is false here, so C already matches Y's length)
292 if (!fNC.empty()) {
293 auto dataC = static_cast<T *>(model.GetInitializedTensorData(fNC).get());
294 for (size_t idx = 0; idx < dataY.size(); idx++)
295 dataY[idx] += static_cast<T>(fAttrBeta) * dataC[idx];
296 }
297 // fuse ReLU now since Generate() will be skipped entirely for a constant output
299 for (auto &v : dataY)
300 v = std::max(v, T(0));
301 }
302
303 model.AddConstantTensor<T>(fNY, shapeY, dataY.data());
304 // flag the operand tensors to not be written in the generated code or weight file
305 model.SetNotWritableInitializedTensor(fNA);
306 model.SetNotWritableInitializedTensor(fNB);
307 if (!fNC.empty())
308 model.SetNotWritableInitializedTensor(fNC);
309 fIsOutputConstant = true;
310
311 if (model.Verbose()) {
312 std::cout << "Gemm (or MatMul) " << fNA << " , " << fNB;
313 if (!fNC.empty())
314 std::cout << " , " << fNC;
315 std::cout << " ---> " << fNY << " (constant) " << ConvertShapeToString(shapeY) << std::endl;
316 }
317 return;
318 }
319
320 if (!fIsDynamic)
321 model.AddIntermediateTensor(fNY, model.GetTensorType(fNA), shapeY);
322 else
323 model.AddDynamicTensor(fNY, model.GetTensorType(fNA), fShapeY);
324
325 if (model.Verbose()){
326 std::cout << "Gemm (or MatMul) " << " ---> " << fNY << " shape ";
327 if (fIsDynamic)
328 std::cout << ConvertDimShapeToString(fShapeY) << std::endl;
329 else
330 std::cout << ConvertShapeToString(shapeY) << std::endl;
331 }
332
333 model.AddNeededStdLib("algorithm");
334
335 // register the inference helper functions used by the generated code
336 if (fType == "float")
337 model.AddNeededHelperFunction("Gemm_Call");
338 // bias handling emits Copy / Fill, fused activation emits Relu
339 if (fNC != "") {
340 model.AddNeededHelperFunction("Copy");
341 model.AddNeededHelperFunction("Fill");
342 }
344 model.AddNeededHelperFunction("Relu");
345 }
346
347 std::string Generate(std::string opName) override {
349 return ""; // no op for constant tensors
350
351 opName = "op_" + opName;
352
353 // if (fShapeA.empty() || fShapeB.empty() || fShapeY.empty() || (fNC != "" && fShapeC.empty())) {
354 // throw std::runtime_error("TMVA SOFIE Gemm Op called to Generate without being initialized first");
355 // }
356 std::stringstream out;
357 out << "\n//--------- Gemm " << opName << " " << ConvertDimShapeToString(fShapeA) << " * " << ConvertDimShapeToString(fShapeB)
358 << " -> " << ConvertDimShapeToString(fShapeY) << "\n";
359 // need to consider case A and B have dim > 2 (for MatMul)
360 int64_t dimA = fShapeA.size();
361 int64_t dimB = fShapeB.size();
362 int64_t dimY = fShapeY.size();
363 int64_t dimC = fDimShapeC.size();
364 if (dimA != dimB || dimA != dimY || (fBroadcastBias && dimC != dimY)) {
365 std::cout << " shape A " << ConvertDimShapeToString(fShapeA)
366 << " shape B " << ConvertDimShapeToString(fShapeB)
367 << " shape C " << ConvertDimShapeToString(fDimShapeC)
368 << " shape Y " << ConvertDimShapeToString(fShapeY) << std::endl;
369 throw std::runtime_error("TMVA SOFIE Gemm(MatMul) has invalid shape for inputs or output");
370 }
371 auto m = (fAttrTransA ? fShapeA[dimA-1].GetVal() : fShapeA[dimA-2].GetVal());
372 auto n = (fAttrTransB ? fShapeB[dimB-2].GetVal() : fShapeB[dimB-1].GetVal());
373 auto k = (fAttrTransA ? fShapeA[dimA-2].GetVal() : fShapeA[dimA-1].GetVal());
374 // size of A: if (transposeA) is m*k else k*m
375 // size of B n*k
376 std::vector<Dim> sY = {fShapeY[dimY-2], fShapeY[dimY-1]};
377 // extra dimensions in case of stacked MatMul
378 std::vector<Dim> sExtraY;
379 for (int64_t i = 0; i < dimY-2; i++) {
380 sExtraY.push_back(fShapeY[i]);
381 }
382 auto lengthGemm = ConvertDimShapeToLength(sY); // size of the Gemm operation
383 auto lengthExtra_Y = ConvertDimShapeToLength(sExtraY); // extra length in case input tensors are of dim>2 (MatMul)
384 std::string lengthExtra_C;
385 std::vector<Dim> sExtraC;
386 std::vector<Dim> sC;
387 bool haveExtraC = false;
388 if (dimC > 2) {
389 sC = {fDimShapeC[dimC-2], fDimShapeC[dimC-1]};
390 for (int64_t i = 0; i < dimC-2; i++) {
391 sExtraC.push_back(fDimShapeC[i]);
392 }
394 if (lengthExtra_C != "1") haveExtraC = true;
395 } else if (dimC > 0) {
396 for (int64_t i = 0; i < dimC; i++) {
397 sC.push_back(fDimShapeC[i]);
398 }
399 }
400
401 // case bias is present
402 if (!fNC.empty()){
403 // when the 2 last dims of bias and Y are not compatible we need to perform a run time broadcast
404 if (sC != sY)
405 fBroadcastBias = true;
407 // C has exactly the shape of Y, nothing to broadcast. Only revisit the
408 // assumption Initialize had to make while the shape of Y was still unknown:
409 // a bias it did compare and found to need broadcasting keeps it.
410 fBroadcastBias = false;
411 if (!fBroadcastBias) {
412 // add a check in case broadcasting was not needed or done outside of session
413 // C should have smaller dimension of Y
414 if (!fIsDynamic) {
415 if ((std::stoi(lengthGemm) != std::stoi(ConvertDimShapeToLength(sC))) ||
416 (haveExtraC && std::stoi(lengthExtra_Y) != std::stoi(lengthExtra_C)))
417 throw std::runtime_error("TMVA SOFIE Gemm Op " + opName + " Bias tensor " + fNC +
418 " has not correct size " + ConvertShapeToString(fShapeC) +
419 " output length " + lengthGemm);
420 } else {
421 // add a dynamic check (C should not be a dynamic tensor)
422 out << SP << "assert(" << lengthGemm << " == " << ConvertDimShapeToLength(sC) << ");\n";
423 if (haveExtraC)
424 out << SP << "assert(" << lengthExtra_Y << " == " << lengthExtra_C << ");\n";
425 }
426 }
427 } else {
428 fBroadcastBias = false;
429 //in this case fAttrBeta needs to be equal to zero otherwise second time we run we will use
430 // the previous result
431 if (fAttrBeta != 0) {
432 // some model don't have bias but Beta is not zero - force it to zero
433 fAttrBeta = 0;
434 std::cout << "WARNING: TMVA SOFIE Gemm Op " + opName + " Bias tensor is not present but beta value in Gemm is not zero - force it to zero\n";
435 }
436 }
437
438 // include MatMul case where we stack the Gemm operations
439 // exclude case where we have only 1's in the additional dims
440 bool doStackMul = dimY > 2 && ( fIsDynamic || std::stoi(lengthExtra_Y) > 1);
441 // compute input offset for stack multiplications
442 std::string lengthExtra_A;
443 std::string lengthExtra_B;
444 std::string increment_A;
445 std::string increment_B;
446
447 if (doStackMul) {
448 std::vector<Dim> sA(fShapeA.begin(), fShapeA.begin()+dimA-2);
449 std::vector<Dim> sB(fShapeB.begin(), fShapeB.begin()+dimB-2);
450 std::vector<Dim> mA = {fShapeA[dimA-2], fShapeA[dimA-1]};
451 std::vector<Dim> mB = {fShapeB[dimB-2], fShapeB[dimB-1]};
454 // if A ( b, m, k) and B (b, k, n) these are the strides of A and B ( m*k for A and n*k for B )
457 }
458 bool extraA = (doStackMul && lengthExtra_A != "1");
459 bool extraB = (doStackMul && lengthExtra_B != "1");
461 // run time check for bias broadcasting
462 std::string biasShapeType = opName + "_biasShapeType";
464 // create a flag according to bias shape:
465 // = 1 for (1,Y2)
466 // = 2 for (Y1,1)
467 // = 3 for a scalar
468 out << SP << "int " << biasShapeType << " = 0;\n";
469 // case vector of columns
470 if (sC[0].GetVal() != "1" && sC[1].GetVal() != sY[1].GetVal())
471 out << SP << "if (" << sC[0] << " == 1 && " << sC[1] << " == " << sY[1] << ")\n";
472 else if (sC[0].GetVal() == "1")
473 out << SP << "if (" << sC[1] << " == " << sY[1] << ")\n";
474 else if (sC[1].GetVal() == sY[1].GetVal())
475 out << SP << "if (" << sC[0] << " == 1)\n";
476
477 out << SP << SP << biasShapeType << " = 1;\n";
478
479 // case vector of rows
480 if (sC[1].GetVal() != "1" && sC[0].GetVal() != sY[0].GetVal())
481 out << SP << "else if (" << sC[1] << " == 1 && " << sC[0] << " == " << sY[0] << ")\n";
482 else if (sC[1].GetVal() == "1")
483 out << SP << "else if (" << sC[0] << " == " << sY[0] << ")\n";
484 else if (sC[0].GetVal() == sY[0].GetVal())
485 out << SP << "else if (" << sC[1] << " == 1)\n";
486
487 out << SP << SP << biasShapeType << " = 2;\n";
488
489 // case scalar
490 if (sC[0].GetVal() != "1" && sC[1].GetVal() != "1")
491 out << SP << "else if (" << sC[0] << " == 1 && " << sC[1] << " == 1 )\n";
492 else if (sC[0].GetVal() == "1")
493 out << SP << "else if (" << sC[1] << " == 1)\n";
494 else if (sC[1].GetVal() == "1")
495 out << SP << "else if (" << sC[0] << " == 1)\n";
496 out << SP << SP << biasShapeType << " = 3;\n";
497 out << SP << "else\n";
498 out << SP << SP << "throw std::runtime_error(\"TMVA SOFIE Gemm Op - bias tensor "
499 << ConvertDimShapeToString(fDimShapeC) << " cannot be broadcasted to "
500 << ConvertDimShapeToString(fShapeY) << "\");\n";
501 }
502 auto SP2 = SP;
503 if (doStackMul) {
504 out << SP << "size_t " << opName << "_y_offset = 0;\n"; // needed if we stack the gemm operations
505 if (extraA)
506 out << SP << "size_t " << opName << "_A_offset = 0;\n";
507 if (extraB)
508 out << SP << "size_t " << opName << "_B_offset = 0;\n";
509 if (extraC)
510 out << SP << "size_t " << opName << "_C_offset = 0;\n";
511 out << SP << "for (size_t i = 0; i < " << lengthExtra_Y << "; i++){\n";
512 SP2 += SP;
513 }
514 // do the bias broadcasting at run time by
515 // initializing output Y vector with bias values
516 if (fBroadcastBias) {
517
518 fAttrBeta = 1.;
519
520 // loop on first output dimension
521 out << SP2 << "for (size_t j = 0; j < " << sY[0] << "; j++) { \n";
522 out << SP2 << SP << "size_t y_index = ";
523 if (doStackMul) // add offset in case of stack multiplications (not sure if bias is present in these cases)
524 out << opName << "_y_offset + ";
525 if (sY[1].GetVal() != "1")
526 out << sY[1] << " * j;\n";
527 else
528 out << "j;\n";
529
530 std::string prefix = SP2 + SP;
531 std::string target = "tensor_" + fNY;
532 if (sC.size() != 2) {
533 throw std::runtime_error("TMVA SOFIE Gemm Op - invalid rank for bias tensor " + ConvertDimShapeToString(fDimShapeC) + ConvertDimShapeToString(sC));
534 } if (sC[0].GetVal() == "1" && sC[1].GetVal() == sY[1].GetVal()) {
535 out << prefix << "Copy(" << target << " + y_index, tensor_" << fNC << ", " << sY[1] << ");\n";
536 } else if (sC[1].GetVal() == "1" && sC[0].GetVal() == sY[0].GetVal()) {
537 out << prefix << "Fill(" << target << " + y_index, tensor_" << fNC << "[j], " << sY[1] << ");\n";
538 } else if (sC[0].GetVal() == "1" && sC[1].GetVal() == "1") {
539 // scalar case
540 out << prefix << "Fill(" << target << " + y_index, tensor_" << fNC << "[0], " << sY[1] << ");\n";
541 } else if (fCheckBiasShapeAtRuntime) {
542 // in the generic dynamic case we check at run time that bias is compatible
543 // we check that bias[0] = 1 or equal to SY[0] and that bias[1] = 1 or equal to SY[1]
544 // tbd: this run-time check coul;d be moved outside the loop for better run time efficiency
545 out << SP2 << SP << "if (" << biasShapeType << " == 1)\n"; // case vector of columns
546 out << SP << prefix << "Copy(" << target << " + y_index, tensor_" << fNC << ", " << sY[1] << ");\n";
547 out << SP2 << SP << "else if (" << biasShapeType << " == 2)\n"; // case vector of rows
548 out << SP << prefix << "Fill(" << target << " + y_index, tensor_" << fNC << "[j], " << sY[1] << ");\n";
549 out << SP2 << SP << "else \n"; // scalar case
550 out << SP << prefix << "Fill(" << target << " + y_index, tensor_" << fNC << "[0], " << sY[1] << ");\n";
551 } else {
552 throw std::runtime_error("TMVA SOFIE Gemm Op - invalid shape for bias tensor " + ConvertDimShapeToString(fDimShapeC));
553 }
554
555 out << SP2 << "}\n";
556 }
557
558 if (fType == "float"){
559
560 out << SP2 << "Gemm_Call(" << "tensor_" << fNY;
561 if (doStackMul) out << " + " << opName << "_y_offset";
562 out << ", "
563 << (fAttrTransB ? "true, " : "false, ")
564 << (fAttrTransA ? "true, " : "false, ")
565 << n << ", " << m << ", " << k << ", ";
566 out << std::setprecision(std::numeric_limits<float>::max_digits10) << fAttrAlpha << ", tensor_" << fNB;
567 if (extraB) out << " + " << opName << "_B_offset";
568 out << ", tensor_" << fNA;
569 if (extraA) out << " + " << opName << "_A_offset";
570 out << ", " << std::setprecision(std::numeric_limits<float>::max_digits10) << fAttrBeta << ",";
571 // in the case of bias and no broadcasting needed - I need to add bias as an extra tensor in Gemm call
572 if (!fNC.empty() && !fBroadcastBias) {
573 out << "tensor_" << fNC;
574 if (extraC) {
575 out << " + " << opName << "_C_offset";
576 }
577 } else {
578 out << "nullptr";
579 }
580 out << ");\n";
581
582 }
583
584 if (doStackMul) {
585 out << SP << SP << opName << "_y_offset += " << lengthGemm << ";\n";
586 if (lengthExtra_A != "1")
587 out << SP << SP << opName << "_A_offset += " << increment_A << ";\n";
588 if (lengthExtra_B != "1")
589 out << SP << SP << opName << "_B_offset += " << increment_B << ";\n";
590 if (extraC)
591 // increment_C is lengthGEmm
592 out << SP << SP << opName << "_C_offset += " << lengthGemm << ";\n";
593 out << SP << "}\n"; // end of loop on the stacked multiplication
594 }
595
596 // fuse with Relu
598 out << SP << "//--- applying RELU to output\n";
599 std::string tnsr = "tensor_" + fNY;
601 out << SP << "Relu(" << tnsr << ", " << tnsr << ", " << reluSize << ");\n";
602 }
603
604 return out.str();
605 }
606
607 std::vector<std::string> GetBlasRoutines() override { return {"Gemm", "Gemv"}; }
608
609 };
610
611
612}//SOFIE
613}//Experimental
614}//TMVA
615
616
617#endif //TMVA_SOFIE_ROPERATOR_GEMM
size_t size(const MatrixT &matrix)
retrieve the size of a square matrix
ROOT::Detail::TRangeCast< T, true > TRangeDynCast
TRangeDynCast is an adapter class that allows the typed iteration through a TCollection.
winID h TVirtualViewer3D TVirtualGLPainter p
Option_t Option_t TPoint TPoint const char GetTextMagnitude GetFillStyle GetLineColor GetLineWidth GetMarkerStyle GetTextAlign GetTextColor GetTextSize void input
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 target
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
const std::vector< std::string > & GetDimShapeNames() const
Definition RModel.hxx:293
ROperator_Gemm(float alpha, float beta, int_t transA, int_t transB, std::string nameA, std::string nameB, std::string nameC, std::string nameY, EActivationType activation=EActivationType::UNDEFINED)
std::vector< Dim > DynamicShapeInference(const std::vector< std::vector< Dim > > &input)
ROperator_Gemm(float alpha, float beta, int_t transA, int_t transB, std::string nameA, std::string nameB, std::string nameY, EActivationType activation=EActivationType::UNDEFINED)
std::vector< U > DoShapeInference(const std::vector< std::vector< U > > &input)
std::string Generate(std::string opName) override
void Initialize(RModel &model) override
std::vector< std::string > GetBlasRoutines() 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
const Int_t n
Definition legend1.C:16
std::vector< size_t > MultidirectionalBroadcastShape(std::vector< std::vector< size_t > >)
std::string ConvertDimShapeToString(const std::vector< Dim > &shape)
std::vector< Dim > ConvertShapeToDim(const std::vector< size_t > &shape)
Convert shape from integer format to dynamic one (based on Dim)
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)
std::string ConvertShapeToString(const std::vector< size_t > &shape)
create variable transformations
TMarker m
Definition textangle.C:8
static uint64_t sum(uint64_t i)
Definition Factory.cxx:2335