Logo ROOT  
Reference Guide
 
Loading...
Searching...
No Matches
ROperator_Where.hxx
Go to the documentation of this file.
1#ifndef TMVA_SOFIE_ROperator_Where
2#define TMVA_SOFIE_ROperator_Where
3
5#include "TMVA/ROperator.hxx"
6#include "TMVA/RModel.hxx"
7
8#include <algorithm> // for std::all_of
9#include <sstream>
10
11namespace TMVA{
12namespace Experimental{
13namespace SOFIE{
14
15
16
17template<typename T>
19private:
20
21 bool fIsInputBoolTensor = false;
22
23
24 std::string fNX;
25 std::string fNY;
26 std::string fNC;
27 std::string fNBroadcastedX;
28 std::string fNBroadcastedY;
29 std::string fNBroadcastedC;
30 std::string fNZ;
31
32
33
34 // static shapes (used when tensors are not dynamic) )
35 std::vector<size_t> fShapeX;
36 std::vector<size_t> fShapeY;
37 std::vector<size_t> fShapeC;
38 std::vector<size_t> fShapeZ;
39
40 // Dynamic generic shapes
41 std::vector<Dim> fDimShapeC;
42 std::vector<Dim> fDimShapeX;
43 std::vector<Dim> fDimShapeY;
44 std::vector<Dim> fDimShapeZ;
45
46 // Broadcast flag: mirrors convention of BasicBinary
47 // bit 0: broadcast Y->X (Y needs expanding)
48 // bit 1: broadcast X->Y (X needs expanding)
49 // bit 2: broadcast C->Z (C needs expanding)
50 // bit 4: shapes may differ at runtime (dynamic)
52
53public:
55 ROperator_Where(const std::string & nameC, const std::string & nameX, const std::string & nameY, const std::string & nameZ):
56 fNX(UTILITY::Clean_name(nameX)), fNY(UTILITY::Clean_name(nameY)), fNC(UTILITY::Clean_name(nameC)), fNZ(UTILITY::Clean_name(nameZ)){
59 }
60
61 void Initialize(RModel& model) override {
62 // input must be a graph input, or already initialized intermediate tensor
63 if (!model.CheckIfTensorAlreadyExist(fNX)){
64 throw std::runtime_error(std::string("TMVA SOFIE Where Op Input Tensor ") + fNX + "is not found in model");
65 }
66 if (!model.CheckIfTensorAlreadyExist(fNY)) {
67 throw std::runtime_error(std::string("TMVA SOFIE Where Op Input Tensor ") + fNY + "is not found in model");
68 }
69 if (!model.CheckIfTensorAlreadyExist(fNC)) {
70 throw std::runtime_error(std::string("TMVA SOFIE Where Op Input Tensor ") + fNC + "is not found in model");
71 }
72 // check if fNC input tensor is boolean
73 if (model.IsReadyInputTensor(fNC))
74 fIsInputBoolTensor = true;
75
76 // ---------------------------------------------------------------- //
77 // Collect shapes – dynamic or static
78 // ---------------------------------------------------------------- //
79 int dynamicInputs = 0; // bitmask: bit0=C, bit1=X, bit2=Y
80
81 if (model.IsDynamicTensor(fNC)) {
82 fDimShapeC = model.GetDynamicTensorShape(fNC);
83 dynamicInputs |= 1;
84 } else {
85 fShapeC = model.GetTensorShape(fNC);
87 }
88 if (model.IsDynamicTensor(fNX)) {
89 fDimShapeX = model.GetDynamicTensorShape(fNX);
90 dynamicInputs |= 2;
91 } else {
92 fShapeX = model.GetTensorShape(fNX);
94 }
95 if (model.IsDynamicTensor(fNY)) {
96 fDimShapeY = model.GetDynamicTensorShape(fNY);
97 dynamicInputs |= 4;
98 } else {
99 fShapeY = model.GetTensorShape(fNY);
101 }
102
103
104 if (model.Verbose()) {
105 if (dynamicInputs & 1)
106 std::cout << "Where : condition " << fNC << " is dynamic " << ConvertDimShapeToString(fDimShapeC) << "\n";
107 if (dynamicInputs & 2)
108 std::cout << "Where : " << fNX << " is dynamic " << ConvertDimShapeToString(fDimShapeX) << "\n";
109 if (dynamicInputs & 4)
110 std::cout << "Where : Y " << fNZ << " is dynamic " << ConvertDimShapeToString(fDimShapeZ) << "\n";
111 }
112
113 // ---------------------------------------------------------------- //
114 // Static path: all shapes known at code-gen time
115 // ---------------------------------------------------------------- //
116 if (dynamicInputs == 0) {
117
119 if (broadcast) {
120 // the broadcast output can be larger than every input, or inputs can have the same
121 // number of elements but different shapes.
123
124 // MultidirectionalBroadcastShape takes its inputs by value, so fShapeX, fShapeY and
125 // fShapeC keep their original rank: prepend the missing unit dimensions so the
126 // per-input broadcast checks below compare equal-rank shapes.
127 auto padToRank = [&](std::vector<size_t> &shape) {
128 if (shape.size() < fShapeZ.size()) {
129 size_t nPrepend = fShapeZ.size() - shape.size();
130 shape.insert(shape.begin(), nPrepend, 1);
131 }
132 };
136
140
141 // Broadcast X to Z
142 if (broadcastX) {
143 fNBroadcastedX = "BC_" + fNX + "_to_" + fNZ;
144 if (model.IsInitializedTensor(fNX)) {
145 auto data = model.GetInitializedTensorData(fNX);
146 std::shared_ptr<void> broadcastedData(
147 UTILITY::UnidirectionalBroadcast(static_cast<T *>(data.get()), fShapeX, fShapeZ),
148 std::default_delete<T[]>());
149 // Update the data and the shape of X
150 model.AddConstantTensor(fNBroadcastedX, model.GetTensorType(fNX), fShapeZ, broadcastedData);
152 }
153 }
154 // Broadcast Y to Z
155 if (broadcastY) {
156 fNBroadcastedY = "BC_" + fNY + "_to_" + fNZ;
157 if (model.IsInitializedTensor(fNY)) {
158 auto data = model.GetInitializedTensorData(fNY);
159 std::shared_ptr<void> broadcastedData(
160 UTILITY::UnidirectionalBroadcast(static_cast<T *>(data.get()), fShapeY, fShapeZ),
161 std::default_delete<T[]>());
162 // do not update tensor B but add broadcasted one (since it can be input to some other operators)
163 model.AddConstantTensor(fNBroadcastedY, model.GetTensorType(fNY), fShapeZ, broadcastedData);
165 }
166 }
167 // Broadcast C to Z
168 if (broadcastC) {
169 fNBroadcastedC = "BC_" + fNC + "_to_" + fNZ;
170 if (model.IsInitializedTensor(fNC)) {
171 auto data = model.GetInitializedTensorData(fNC);
172 std::shared_ptr<void> broadcastedData(
173 UTILITY::UnidirectionalBroadcast(static_cast<T *>(data.get()), fShapeC, fShapeZ),
174 std::default_delete<T[]>());
175 // do not update tensor C but add broadcasted one (since it can be input to some other operators)
176 model.AddConstantTensor(fNBroadcastedC, model.GetTensorType(fNC), fShapeZ, broadcastedData);
178 }
179 }
180 } else {
182 }
183 // check case of constant output (if all inputs are defined)
184 if (model.IsInitializedTensor(fNC)) {
185 std::string nameC = fNBroadcastedC.empty() ? fNC : fNBroadcastedC;
186 auto dataC = static_cast<bool *>(model.GetInitializedTensorData(nameC).get());
187 model.SetNotWritableInitializedTensor(nameC);
188 T *dataX = nullptr;
189 T *dataY = nullptr;
190 std::vector<Dim> shapeDataX;
191 std::vector<Dim> shapeDataY;
192 if (model.IsInitializedTensor(fNX)) {
193 std::string nameX = fNBroadcastedX.empty() ? fNX : fNBroadcastedX;
194 dataX = static_cast<T *>(model.GetInitializedTensorData(nameX).get());
195 // flag tensors to not be written in a file
196 model.SetNotWritableInitializedTensor(nameX);
197 } else if (model.IsShapeTensor(fNX)) {
198 shapeDataX = model.GetShapeTensorValues(fNX);
199 }
200 if (model.IsInitializedTensor(fNY)) {
201 std::string nameY = fNBroadcastedY.empty() ? fNY : fNBroadcastedY;
202 dataY = static_cast<T *>(model.GetInitializedTensorData(nameY).get());
203 model.SetNotWritableInitializedTensor(nameY);
204 } else if (model.IsShapeTensor(fNY)) {
205 shapeDataY = model.GetShapeTensorValues(fNY);
206 }
207 std::vector<T> dataZ; // used in case output is constant tensor
208 std::vector<Dim> shapeDataZ; // used in case output is a shape tensor (can be also constant if all
209 // dimensions are not parametric)
210 // if fNC (condition) is initialized we know the output is a shape or a constant tensor,
211 // so we can compute it at initialization and add it as a constant tensor to the model
212 // (and not add the operator output as intermediate tensor to the model)
213 bool isOutputConstantTensor = true;
214 if (dataX && dataY) {
216 for (size_t i = 0; i < dataZ.size(); i++)
217 dataZ[i] = (dataC[i]) ? dataX[i] : dataY[i];
218 if (model.Verbose())
219 std::cout << "data A and B : dataZ constant: " << ConvertValuesToString(dataZ) << std::endl;
220 } else if (dataX && shapeDataY.size() > 0) {
222 for (size_t i = 0; i < shapeDataZ.size(); i++) {
223 shapeDataZ[i] = (dataC[i]) ? Dim{size_t(dataX[i])} : shapeDataY[i];
224 isOutputConstantTensor &= !shapeDataZ[i].isParam;
225 }
226 if (model.Verbose())
227 std::cout << "data A but shapeB " << ConvertDimShapeToString(shapeDataY) << " "
228 << isOutputConstantTensor << std::endl;
229 } else if (dataY && shapeDataX.size() > 0) {
231 for (size_t i = 0; i < shapeDataZ.size(); i++) {
232 shapeDataZ[i] = (dataC[i]) ? shapeDataY[i] : Dim{size_t(dataY[i])};
233 isOutputConstantTensor &= !shapeDataZ[i].isParam;
234 }
235 if (model.Verbose())
236 std::cout << "data B but shapeA " << ConvertDimShapeToString(shapeDataX) << " "
237 << isOutputConstantTensor << std::endl;
238 } else if (shapeDataY.size() > 0 && shapeDataX.size() > 0) {
240 for (size_t i = 0; i < shapeDataZ.size(); i++) {
241 shapeDataZ[i] = (dataC[i]) ? shapeDataX[i] : shapeDataY[i];
242 isOutputConstantTensor &= !shapeDataZ[i].isParam;
243 }
244 if (model.Verbose())
245 std::cout << " shapeA and B " << ConvertDimShapeToString(shapeDataX) << " shapeB "
247 }
248 fIsOutputConstant = true;
249 // add as constant or shape tensor depending on the case
250 if (dataZ.size() > 0)
251 model.AddConstantTensor<T>(fNZ, fShapeZ, dataZ.data());
252 else if (shapeDataZ.size() > 0)
253 model.AddShapeTensor(fNZ, shapeDataZ, fShapeZ.size() == 0);
254 else {
255 fIsOutputConstant = false;
256 }
257 if (fIsOutputConstant && model.Verbose())
258 std::cout << "Where op ---> " << fNZ << " " << ConvertShapeToString(fShapeZ) << " : "
260 << ((dataZ.size() > 0) ? " (constant)" : " (shape)") << std::endl;
261
262 // output is a constant tensor
264 fOutputTensorNames.pop_back();
265 }
266 if (!fIsOutputConstant) {
267
269 model.AddIntermediateTensor(fNZ, model.GetTensorType(fNX), fShapeZ);
270 if (model.Verbose())
271 std::cout << "Where : condition : " << fNC << " " << ConvertShapeToString(fShapeC) << " X "
272 << fNX << " " << ConvertShapeToString(fShapeX) << " Y " << fNY << " "
273 << ConvertShapeToString(fShapeY) << " ---> " << fNZ << " " << ConvertShapeToString(fShapeZ)
274 << std::endl;
275 }
276 } else {
277 // ---------------------------------------------------------------- //
278 // Dynamic path: at least one input has a parametric shape
279 // Need to use BroadcastShape to find output shape
280 // ---------------------------------------------------------------- //
282 fBroadcastFlag = retXY.first;
283 fDimShapeZ = retXY.second;
285 fBroadcastFlag |= retCZ.first;
286 fDimShapeZ = retCZ.second;
287
288 // Resolve std::max params to actual input dim params (same logic as BasicBinary)
289 if (fBroadcastFlag & 4) {
290 auto IsInputDimParam = [&](const std::string &p) {
291 for (auto &input : model.GetInputTensorNames())
292 for (auto &s : model.GetDimTensorShape(input))
293 if (s.isParam && s.param == p) return true;
294 return false;
295 };
296 for (size_t i = 0; i < fDimShapeZ.size(); i++) {
297 auto &s = fDimShapeZ[i];
298 if (s.isParam && s.param.find("std::max") != std::string::npos) {
299 // prefer A dim over B dim
300 if (i < fDimShapeX.size() && IsInputDimParam(fDimShapeX[i].param)) {
301 s = (fDimShapeX[i].dim != 1) ? fDimShapeX[i] : fDimShapeY[i];
302 } else if (i < fDimShapeY.size() && IsInputDimParam(fDimShapeY[i].param)) {
303 s = (fDimShapeY[i].dim != 1) ? fDimShapeY[i] : fDimShapeX[i];
304 }
305 }
306 }
307 }
308 // I need to prepend to shape of X,Y,C the extra dimensions added for broadcasting to Z
309 if (fDimShapeX.size() < fDimShapeZ.size()) {
310 size_t nPrepend = fDimShapeZ.size() - fDimShapeX.size();
311 fDimShapeX.insert(fDimShapeX.begin(), nPrepend, Dim{1});
312 }
313 if (fDimShapeY.size() < fDimShapeZ.size()) {
314 size_t nPrepend = fDimShapeZ.size() - fDimShapeY.size();
315 fDimShapeY.insert(fDimShapeY.begin(), nPrepend, Dim{1});
316 }
317 if (fDimShapeC.size() < fDimShapeZ.size()) {
318 size_t nPrepend = fDimShapeZ.size() - fDimShapeC.size();
319 fDimShapeC.insert(fDimShapeC.begin(), nPrepend, Dim{1});
320 }
321
322 model.AddIntermediateTensor(fNZ, model.GetTensorType(fNX), fDimShapeZ);
323
324 if (model.Verbose())
325 std::cout << "Where (dynamic) : C=" << ConvertDimShapeToString(fDimShapeC)
328 << " --> Y=" << ConvertDimShapeToString(fDimShapeZ) << "\n";
329 }
330 }
331
332 std::string GenerateInitCode() override {
333 std::stringstream out;
334 return out.str();
335 }
336
337 std::string Generate(std::string opName) override {
338
339 opName = "op_" + opName;
340 std::stringstream out;
341 out << SP << "\n//------ WHERE " << opName << " --> " << ConvertDimShapeToString(fDimShapeZ) << "\n";
342 if (fIsOutputConstant) return out.str();
343
344
345 // ---------------------------------------------------------------- //
346 // Runtime broadcast validation (dynamic shapes, flag bit 4)
347 // ---------------------------------------------------------------- //
348 if (fBroadcastFlag & 4) {
352 out << SP << "if (" << lengthX << " != " << lengthY << " || "
353 << lengthX << " != " << lengthC << ") {\n";
354 for (size_t i = 0; i < fDimShapeZ.size(); i++) {
355 // validate X vs Z
356 if (i < fDimShapeX.size() && fDimShapeX[i].isParam) {
357 out << SP << SP << "if (" << fDimShapeX[i] << " != 1 && "
358 << fDimShapeX[i] << " != " << fDimShapeZ[i] << ")\n";
359 out << SP << SP << SP
360 << "throw std::runtime_error(\"SOFIE Where: cannot broadcast A dim " << i << " in " << opName << "\");\n";
361 }
362 // validate Y vs Z
363 if (i < fDimShapeY.size() && fDimShapeY[i].isParam) {
364 out << SP << SP << "if (" << fDimShapeY[i] << " != 1 && "
365 << fDimShapeY[i] << " != " << fDimShapeZ[i] << ")\n";
366 out << SP << SP << SP
367 << "throw std::runtime_error(\"SOFIE Where: cannot broadcast B dim " << i << " in " << opName << "\");\n";
368 }
369 // validate C vs Z
370 if (i < fDimShapeC.size() && fDimShapeC[i].isParam) {
371 out << SP << SP << "if (" << fDimShapeC[i] << " != 1 && "
372 << fDimShapeC[i] << " != " << fDimShapeZ[i] << ")\n";
373 out << SP << SP << SP
374 << "throw std::runtime_error(\"SOFIE Where: cannot broadcast C dim " << i << " in " << opName << "\");\n";
375 }
376 }
377 out << SP << "}\n";
378 }
379 // implement now where using teh strides and looping on the different dimensions
380 // ---------------------------------------------------------------- //
381 // Generate loop(s) with per-dimension stride-based index arithmetic
382 // ---------------------------------------------------------------- //
387
388 auto buildIdxExpr = [&](const std::vector<Dim> &dimShape,
389 const std::vector<Dim> &strides,
390 size_t rankZ) -> std::string {
391 if (dimShape.empty() ||
392 std::all_of(dimShape.begin(), dimShape.end(),
393 [](Dim d) { return d.dim == 1 || d.GetVal() == "1"; }))
394 return "0";
395 std::string expr;
396 size_t offset = rankZ - dimShape.size();
397 for (size_t i = 0; i < dimShape.size(); ++i) {
398 if (dimShape[i].dim == 1 || dimShape[i].GetVal() == "1") continue;
399 expr += "idx_" + std::to_string(i + offset);
400 if (strides[i].GetVal() != "1")
401 expr += " * " + strides[i].GetVal();
402 expr += " + ";
403 }
404 if (expr.size() >= 3)
405 for (int j = 0; j < 3; j++) expr.pop_back(); // remove trailing " + "
406 return expr.empty() ? "0" : expr;
407 };
408
409 std::string idxX = buildIdxExpr(fDimShapeX, stridesX, fDimShapeZ.size());
410 std::string idxY = buildIdxExpr(fDimShapeY, stridesY, fDimShapeZ.size());
411 std::string idxC = buildIdxExpr(fDimShapeC, stridesC, fDimShapeZ.size());
412
413 // Emit nested loops over output shape
414 int nloop = 0;
415 std::string idxZ;
416 // case Z is a scalar (all dimensions are 1) or Z has no dimension
417 if (fDimShapeZ.empty() ||
418 std::all_of(fDimShapeZ.begin(), fDimShapeZ.end(),
419 [](Dim d) { return d.dim == 1 || d.GetVal() == "1"; })) {
420 idxZ = "0";
421 } else {
422 for (size_t i = 0; i < fDimShapeZ.size(); ++i) {
423 if (fDimShapeZ[i].dim != 1 && fDimShapeZ[i].GetVal() != "1") {
424 nloop++;
425 for (int j = 0; j < nloop; j++) out << SP;
426 out << "for (size_t idx_" << i << " = 0; idx_" << i
427 << " < " << fDimShapeZ[i] << "; ++idx_" << i << ") {\n";
428 idxZ += "idx_" + std::to_string(i);
429 if (stridesZ[i].GetVal() != "1")
430 idxZ += " * " + stridesZ[i].GetVal();
431 idxZ += " + ";
432 }
433 }
434 if (idxZ.size() >= 3)
435 for (int j = 0; j < 3; j++) idxZ.pop_back();
436 }
437
438 // Inner assignment
439 for (int j = 0; j < nloop + 1; j++) out << SP;
440 out << "tensor_" << fNZ << "[" << idxZ << "] = "
441 << "tensor_" << fNC << "[" << idxC << "] ? "
442 << "tensor_" << fNX << "[" << idxX << "] : "
443 << "tensor_" << fNY << "[" << idxY << "];\n";
444
445 // Close loops
446 for (int i = nloop; i > 0; i--) {
447 for (int j = 0; j < i; j++) out << SP;
448 out << "}\n";
449 }
450
451 return out.str();
452 }
453
454
455};
456
457}//SOFIE
458}//Experimental
459}//TMVA
460
461
462#endif //TMVA_SOFIE_ROperator_Where
#define d(i)
Definition RSha256.hxx:102
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 data
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 offset
const_iterator begin() const
const_iterator end() const
std::string Generate(std::string opName) override
ROperator_Where(const std::string &nameC, const std::string &nameX, const std::string &nameY, const std::string &nameZ)
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
bool AreSameShape(const std::vector< size_t > &, const std::vector< size_t > &)
std::vector< size_t > MultidirectionalBroadcastShape(std::vector< std::vector< size_t > >)
T * UnidirectionalBroadcast(const T *data, const std::vector< size_t > &shape, const std::vector< size_t > &targetShape)
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 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< Dim > ConvertShapeToDim(const std::vector< size_t > &shape)
Convert shape from integer format to dynamic one (based on Dim)
std::string ConvertDimShapeToLength(const std::vector< Dim > &shape)
std::string ConvertShapeToString(const std::vector< size_t > &shape)
create variable transformations