Logo ROOT  
Reference Guide
 
Loading...
Searching...
No Matches
ROperator_Einsum.hxx
Go to the documentation of this file.
1#ifndef TMVA_SOFIE_ROperator_Einsum
2#define TMVA_SOFIE_ROperator_Einsum
3
5#include "TMVA/ROperator.hxx"
6#include "TMVA/RModel.hxx"
7
8#include <sstream>
9#include <cassert>
10
11namespace TMVA{
12namespace Experimental{
13namespace SOFIE{
14
15
16
17template<typename T>
19private:
20
21 bool fIsInputBoolTensor = false;
22
23
24 std::vector<std::string> fNInputs;
25 std::string fNY;
26
27 std::vector<std::string> fInputLabels;
28 std::string fOutputLabels;
29 std::string fSumLabels; // string containing the reducing labels
30 std::string fGemmType;
31
32 std::vector<int> fSumDims; // dimension of the labels we use to perform summing
33
34 std::vector<std::vector<size_t>> fShapeInputs;
35 std::vector<size_t> fShapeY;
36
37
38
39
40public:
42 ROperator_Einsum(const std::string & equation, const std::vector<std::string> & namesX, const std::string & nameY):
43 fNInputs(namesX.size()), fNY(UTILITY::Clean_name(nameY))
44 {
45 for (size_t i = 0; i < namesX.size(); i++)
47
48 // parse teh equations to find labels
50 throw std::runtime_error("TMVA SOFIE Einsum Op: Error parsing the equation " + equation);
51
52 fInputTensorNames.resize(fNInputs.size());
53 std::transform(fNInputs.begin(), fNInputs.end(), fInputTensorNames.begin(),
54 [](const std::string& s) -> std::string_view { return s; });
56 }
57
58 bool ParseEquation(const std::string & input_equation) {
59 std::string eq (input_equation);
60 // remove blank spaces
61 eq.erase(std::remove(eq.begin(), eq.end(), ' '), eq.end());
62 // look for '->' finding the first occurrence
63 std::string target("->");
64 size_t pos = eq.find(target);
65 if (pos == std::string::npos) {
66 std::cout << "'->' not found in the equation." << std::endl;
67 return false;
68 }
69 // Substring before the target
70 std::string inputStr = eq.substr(0, pos);
71 // Substring after the target
72 std::string outputStr = eq.substr(pos + target.length());
73
74 // look now for the group of labels separated by "," in the inputs
75 size_t start = 0;
76 size_t pos1 = 0;
77 // Extract labels separated by commas
78 while ((pos1 = inputStr.find(',', start)) != std::string::npos) {
79 std::string labels = inputStr.substr(start, pos1 - start);
80 fInputLabels.push_back(labels);
81 start = pos1 + 1; // Move past the comma
82 }
83 // Add the last label (after the final comma)
84 fInputLabels.push_back(inputStr.substr(start));
85
86 // check if labels are ok and do not contain alphanumeric characters
87 auto checkLabel = [](const std::string & label) {
88 for (char c : label) {
89 if (!std::isalnum(c)) {
90 std::cout << "Wrong tensor label " << label << std::endl;
91 return false;
92 }
93 }
94 // empty label is OK , is a scalar
95 return true;
96 };
97 for (auto & label : fInputLabels) {
98 if (!checkLabel(label)) return false;
99 }
100 if (!checkLabel(outputStr)) {
101 std::cout << "invalid output label" << std::endl;
102 return false;
103 }
105
106 if (fInputLabels.size() != fNInputs.size()) {
107 std::cout << "Invalid number of input labels found " << fInputLabels.size() << " for #inputs = " << fNInputs.size() << std::endl;
108 return false;
109 }
110 // ignore for the time being broadcasting, empty output label and other features
111 return true;
112 }
113
114 void Initialize(RModel& model) override {
115 // input must be a graph input, or already initialized intermediate tensor
116 size_t i = 0;
117 std::map<char, int> labelsMap;
118 for ( auto & name : fNInputs) {
119 if (!model.CheckIfTensorAlreadyExist(name))
120 throw std::runtime_error(std::string("TMVA SOFIE Einsum Op Input Tensor ") + name + "is not found in model");
121
122 // if (model.IsDynamicTensor(name) || model.IsDimInputTensor(name) ) {
123 // // not yet supported
124 // } else {
125 auto shape = model.GetTensorShape(name);
126 fShapeInputs.push_back(shape);
127 //}
128 // fill the label maps
129 std::string labels = fInputLabels[i];
130 for (size_t j = 0; j < shape.size(); j++) {
131 if (j >= labels.length()) {
132 throw std::runtime_error(std::string("TMVA SOFIE Einsum Op Input Tensor has invalid label or shape ") + labels + " " + ConvertShapeToString(shape));
133 }
134 labelsMap[labels[j]] = shape[j];
135 }
136 i++;
137 }
138 // get output shape from label maps
139 for (char l : fOutputLabels) {
140 if (labelsMap.count(l) == 0)
141 throw std::runtime_error(std::string("TMVA SOFIE Einsum Op : output label ") + std::string(&l) + " is not present in inputs");
142 fShapeY.push_back(labelsMap[l]);
143 }
144 // we need to get the labels we are going to sum
145 // these are the labels not present in the output
146 fSumLabels = "";
147 fSumDims.clear();
148 for (auto & l : labelsMap) {
149 if (fOutputLabels.find(l.first) == std::string::npos) {
150 fSumLabels += l.first;
151 fSumDims.push_back(l.second);
152 }
153 }
154
155 // check if we can use MatMul for EinSum
156 // need to have one sum labels in the last 2 and have the first in common
157 if (fNInputs.size() == 2 && fSumDims.size() == 1 && fShapeInputs[0].size() >=2 && fShapeInputs[1].size() >= 2 ) {
158 // find positions of dum labels
159 char l = fSumLabels[0];
160 size_t pos1 = fInputLabels[0].find(l);
161 size_t pos2 = fInputLabels[1].find(l);
162 // check if summing is done in the last 2 indices of tensor
163
164 if (pos1 == fInputLabels[0].length() - 1 && pos2 == fInputLabels[1].length() - 2)
165 fGemmType = "nn";
166 else if (pos1 == fInputLabels[0].length() - 2 && pos2 == fInputLabels[1].length() - 2)
167 fGemmType = "tn";
168 else if (pos1 == fInputLabels[0].length() - 1 && pos2 == fInputLabels[1].length() - 1)
169 fGemmType = "nt";
170 else if (pos1 == fInputLabels[0].length() - 2 && pos2 == fInputLabels[1].length() - 1)
171 fGemmType = "tt";
172 else
173 fGemmType = "";
174 }
175
176 model.AddIntermediateTensor(fNY, model.GetTensorType(fNInputs[0]), fShapeY);
177
178 if (model.Verbose()) {
179 std::cout << "Einsum op ";
180 for (i = 0; i < fNInputs.size(); i++) {
181 if (i > 0) std::cout << ", ";
182 std::cout << fNInputs[i] << " " << ConvertShapeToString(fShapeInputs[i]) << " " << fInputLabels[i];
183 }
184 std::cout << " --> " << fNY << " " << ConvertShapeToString(fShapeY) << " " << fOutputLabels << std::endl;
185 }
186
187 }
188
189 std::string GenerateInitCode() override {
190 std::stringstream out;
191 return out.str();
192 }
193
194 std::string Generate(std::string opName) override {
195
196 if (fIsOutputConstant) return "";
197
198 opName = "op_" + opName;
199
200 if (fShapeY.size() != fOutputLabels.length()) {
201 throw std::runtime_error("TMVA SOFIE Einsum Op called to Generate without being initialized first");
202 }
203
204 // function to write compute expression index from strides
205 auto tensorIndex = [](const std::vector<size_t> & stride, const std::string & labels) {
206 std::stringstream strst;
207 int dims = labels.length();
208 // scalar case
209 if (dims == 0) return std::string("0");
210 assert (dims == (int) stride.size());
211 for (int i = 0; i < dims-1; i++) {
212 strst << stride[i] << "*" << std::string{labels[i]} << " + ";
213 }
214 strst << std::string{labels[dims-1]};
215 return strst.str();
216 };
217
218 std::stringstream out;
219 out << SP << "\n//-------- Einsum \n";
220
222
223 // loops on the output indices i0,....iN
224 if (fGemmType.empty()) {
225 int outDims = fShapeY.size();
226 int inDims = fSumLabels.length();
227 assert(outDims == int(fOutputLabels.size()));
228 assert(inDims == int(fSumDims.size()));
229 for (int i = 0; i < outDims; i++) {
230 for (int j = 0; j < i; j++) out << SP;
231 std::string l {fOutputLabels[i]};
232 out << "for (int " << l << " = 0; " << l << " < " << fShapeY[i] << "; " << l << "++) {\n";
233 }
234 // reset to zero output tensor
236
237 for (int j = 0; j < outDims; j++) out << SP;
238 out << "tensor_" << fNY << "[" << outputIndex << "] = 0;\n";
239 // loop on remaining indices where we perform the sum
240 for (int i = 0; i < inDims; i++) {
241 for (int j = 0; j < outDims + i; j++) out << SP;
242 std::string l {fSumLabels[i]};
243 out << "for (int " << l << " = 0; " << l << " < " << fSumDims[i] << "; " << l << "++) {\n";
244 }
245 for (int j = 0; j < outDims+inDims; j++) out << SP;
246 // tensor_out[outId] += t_in_0[ind0] * t_in1[ind1] *....
247 out << "tensor_" << fNY << "[" << outputIndex << "] +=\n";
248 for (size_t k = 0; k < fNInputs.size(); k++) {
251 for (int j = 0; j < outDims+inDims; j++) out << SP;
252 out << SP << "tensor_" << fNInputs[k] << "[" << inputIndex << "]";
253 if (fNInputs.size() > 1 && k < fNInputs.size() -1) out << " *\n";
254 }
255 out << ";\n";
256
257 // end loops on all indices i0,....iN
258 for (int i = outDims+inDims-1; i >= 0; i--) {
259 for (int j = 0; j < i; j++) out << SP;
260 out << "}\n";
261 }
262
263
264 } else {
265 // case we use Gemm
266 out << SP << "// implementing Einsum using MatMul \n";
267 // note A is second input and B first one - due to transpose of Fortran rep.
268 out << SP << "char " << opName << "_transA = '" << fGemmType[0] << "';\n";
269 out << SP << "char " << opName << "_transB = '" << fGemmType[1] << "';\n";
270 // need to consider case A and B have dim > 2 (for MatMul)
271 int64_t dimA = fShapeInputs[0].size();
272 int64_t dimB = fShapeInputs[1].size();
273
274 auto m = (fGemmType[0] == 't') ? fShapeInputs[0][dimA-1] : fShapeInputs[0][dimA-2];
275 auto n = (fGemmType[1] == 't') ? fShapeInputs[1][dimB-2] : fShapeInputs[1][dimB-1];
276 auto k = (fGemmType[0] == 't') ? fShapeInputs[0][dimA-2] : fShapeInputs[0][dimA-1];
277
278 out << SP << "int " << opName << "_m = " << m << ";\n";
279 out << SP << "int " << opName << "_n = " << n << ";\n";
280 out << SP << "int " << opName << "_k = " << k << ";\n";
281 out << SP << "float " << opName << "_alpha = 1.0;\n";
282 out << SP << "float " << opName << "_beta = 0.0;\n";
283 out << SP << "int " << opName << "_lda = " << ((fGemmType[0] == 't') ? m : k) << ";\n";
284 out << SP << "int " << opName << "_ldb = " << ((fGemmType[1] == 't') ? k : n) << ";\n";
285
288
289 int stackDims = fShapeY.size()-2;
290 for (int i = 0; i < stackDims; i++) {
291 for (int j = 0; j < i; j++) out << SP;
292 std::string l {fOutputLabels[i]};
293 out << "for (int " << l << " = 0; " << l << " < " << fShapeY[i] << "; " << l << "++) {\n";
294 }
295 auto tensorOffset = [](const std::vector<size_t> & stride, const std::string & labels) {
296 std::stringstream strst;
297 int dims = labels.length()-2;
298 // scalar case
299 if (dims == 0) return std::string("0");
300 assert (dims +2 == (int) stride.size());
301 for (int i = 0; i < dims; i++) {
302 strst << stride[i] << "*" << std::string{labels[i]};
303 if (i < dims-1) strst << " + ";
304 }
305 return strst.str();
306 };
307 // only float type supported
308 out << SP << "BLAS::sgemm_(&" << opName << "_transB, &" << opName << "_transA, &" << opName
309 << "_n, &" << opName << "_m, &" << opName << "_k, &" << opName << "_alpha, "
310 << "&tensor_" << fNInputs[1] << "[" << tensorOffset(inputStrideB, fInputLabels[1])
311 << "], &" << opName << "_ldb, "
312 << "&tensor_" << fNInputs[0] << "[" << tensorOffset(inputStrideA, fInputLabels[0] ) << "], &" << opName << "_lda, &" << opName << "_beta, "
313 << "&tensor_" << fNY << "[" << tensorOffset(outputStride,fOutputLabels) << "], &" << opName << "_n);\n";
314
315
316 for (int i = stackDims-1; i >= 0; i--) {
317 for (int j = 0; j < i; j++) out << SP;
318 out << "}\n";
319 }
320
321 }
322
323
324 return out.str();
325 }
326
327 std::vector<std::string> GetBlasRoutines() override {
328 return { std::string("Gemm") };
329 }
330};
331
332}//SOFIE
333}//Experimental
334}//TMVA
335
336
337#endif //TMVA_SOFIE_ROperator_Einsum
#define c(i)
Definition RSha256.hxx:101
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.
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 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 length
char name[80]
Definition TGX11.cxx:142
const_iterator begin() const
const_iterator end() const
std::vector< std::vector< size_t > > fShapeInputs
bool ParseEquation(const std::string &input_equation)
std::string Generate(std::string opName) override
ROperator_Einsum(const std::string &equation, const std::vector< std::string > &namesX, const std::string &nameY)
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::string Clean_name(std::string input_tensor_name)
std::vector< size_t > ComputeStrideFromShape(const std::vector< size_t > &shape)
compute stride of a tensor given its shape (assume layout is row-major)
std::string ConvertShapeToString(const std::vector< size_t > &shape)
create variable transformations
TMarker m
Definition textangle.C:8
TLine l
Definition textangle.C:4