File indexing completed on 2026-10-02 09:22:26
0001 #ifndef TMVA_SOFIE_ROPERATOR_POOL
0002 #define TMVA_SOFIE_ROPERATOR_POOL
0003
0004 #include "TMVA/SOFIE_common.hxx"
0005 #include "TMVA/ROperator.hxx"
0006 #include "TMVA/RModel.hxx"
0007
0008 #include <memory>
0009 #include <sstream>
0010 #include <algorithm>
0011 #include <stdexcept>
0012 #include <vector>
0013 #include <cassert>
0014
0015 namespace TMVA {
0016 namespace Experimental {
0017 namespace SOFIE {
0018
0019 struct RAttributes_Pool {
0020
0021 std::string auto_pad = "NOTSET";
0022 int ceil_mode = 0;
0023 int count_include_pad = 0;
0024 int storage_order = 0;
0025 std::vector<size_t> dilations;
0026 std::vector<size_t> kernel_shape;
0027 std::vector<size_t> pads;
0028 std::vector<size_t> strides;
0029 };
0030
0031 enum PoolOpMode { InvalidPool, MaxPool, AveragePool, GlobalAveragePool };
0032
0033 template<typename T>
0034 class ROperator_Pool final : public ROperator
0035 {
0036
0037 private:
0038
0039 PoolOpMode fPoolMode;
0040
0041 size_t fAttrCeilMode;
0042 size_t fAttrCountIncludePad;
0043 size_t fAttrStorageOrder;
0044 std::string fAttrAutopad;
0045 std::vector<size_t> fAttrDilations;
0046 std::vector<size_t> fAttrKernelShape;
0047 std::vector<size_t> fAttrPads;
0048 std::vector<size_t> fAttrStrides;
0049
0050 std::string fNX;
0051 std::string fNY;
0052
0053 std::vector<size_t> fShapeX;
0054 std::vector<size_t> fShapeY;
0055
0056 std::string fType;
0057
0058 size_t fDim;
0059 bool fUseSession = false;
0060
0061 public:
0062
0063 std::string Name() {
0064 if (fPoolMode == AveragePool) return "AveragePool";
0065 if (fPoolMode == MaxPool) return "MaxPool";
0066 return "Invalid";
0067 }
0068
0069 ROperator_Pool() {}
0070
0071 ROperator_Pool(PoolOpMode mode, RAttributes_Pool attr, std::string nameX, std::string nameY)
0072 : fPoolMode(mode), fAttrCeilMode(attr.ceil_mode), fAttrCountIncludePad(attr.count_include_pad),
0073 fAttrStorageOrder(attr.storage_order), fAttrAutopad(attr.auto_pad),
0074 fAttrDilations(attr.dilations), fAttrKernelShape(attr.kernel_shape), fAttrPads(attr.pads), fAttrStrides(attr.strides),
0075 fNX(UTILITY::Clean_name(nameX)), fNY(UTILITY::Clean_name(nameY))
0076 {
0077 if(std::is_same<T, float>::value) {
0078 fType = "float";
0079 } else {
0080 throw
0081 std::runtime_error("TMVA SOFIE Encountered unsupported type parsing a Pool operator");
0082 }
0083 fInputTensorNames = { fNX };
0084 fOutputTensorNames = { fNY };
0085 }
0086
0087
0088 std::vector<ETensorType> TypeInference(std::vector<ETensorType> input) override {
0089
0090 return input;
0091 }
0092
0093
0094 std::vector<std::vector<size_t>> ShapeInference(std::vector<std::vector<size_t>> input) override {
0095
0096
0097
0098 if (input.size() != 1 ) {
0099 throw std::runtime_error("TMVA SOFIE" + Name() + "Op Shape inference need 1 input tensor");
0100 }
0101 if (input[0].size() < 3) {
0102 throw std::runtime_error("TMVA SOFIE" + Name() + "Op Shape inference only accept tensor with at least 3 dimensions");
0103 }
0104
0105 if (input[0].size() < 3 || input[0].size() > 5) {
0106 throw std::runtime_error("TMVA SOFIE" + Name() + "Op : tensors with dimension " + std::to_string(input[0].size()) + " are not yet supported");
0107 }
0108
0109 if (input[0].size() -2 != fDim) {
0110 throw
0111 std::runtime_error("TMVA SOFIE Pool Op Shape inference - invalid inputs ");
0112 }
0113
0114 size_t k1 = ((fAttrKernelShape.empty())? input[0][2] : fAttrKernelShape[0]);
0115 size_t k2 = (fDim > 1) ? ((fAttrKernelShape.empty()) ? input[0][3] : fAttrKernelShape[1]) : 1;
0116 size_t k3 = (fDim > 2) ? ((fAttrKernelShape.empty()) ? input[0][4] : fAttrKernelShape[2]) : 1;
0117
0118
0119 size_t i1 = (fDim > 1) ? ((fDim > 2) ? 3 : 2) : 1;
0120 size_t i2 = (fDim > 2) ? 4 : 3;
0121 size_t i3 = 5;
0122
0123 if (fAttrDilations.empty()) {
0124 fAttrDilations = {1, 1, 1};
0125 }
0126 fAttrDilations.resize(3);
0127 if (fDim < 3) {
0128 fAttrDilations.resize(3, 1);
0129 }
0130
0131 fAttrKernelShape = {k1 + (fAttrDilations[0] - 1) * (k1 - 1),
0132 k2 + (fAttrDilations[1] - 1) * (k2 - 1),
0133 k3 + (fAttrDilations[2] - 1) * (k3 - 1)};
0134
0135 if (fAttrStrides.empty()) {
0136 fAttrStrides = {1, 1, 1};
0137 }
0138 if (fDim < 3)
0139 fAttrStrides.resize(3, 1);
0140
0141 if (fAttrAutopad == "NOTSET") {
0142
0143 if (fAttrPads.empty()) {
0144 fAttrPads = {0, 0, 0, 0, 0, 0};
0145 }
0146 } else if (fAttrAutopad == "SAME_UPPER" || fAttrAutopad == "SAME_LOWER") {
0147
0148
0149 fAttrPads.assign(6, 0);
0150 for (size_t d = 0; d < fDim; ++d) {
0151 size_t inSize = input[0][d + 2];
0152 size_t stride_d = fAttrStrides[d];
0153 size_t outSize = (inSize + stride_d - 1) / stride_d;
0154 int totalPad = std::max(0, (int)((outSize - 1) * stride_d + fAttrKernelShape[d]) - (int)inSize);
0155 if (fAttrAutopad == "SAME_UPPER") {
0156 fAttrPads[d] = (size_t)(totalPad / 2);
0157 fAttrPads[d + fDim] = (size_t)(totalPad - totalPad / 2);
0158 } else {
0159 fAttrPads[d] = (size_t)(totalPad - totalPad / 2);
0160 fAttrPads[d + fDim] = (size_t)(totalPad / 2);
0161 }
0162 }
0163 } else if (fAttrAutopad != "VALID") {
0164 throw
0165 std::runtime_error("TMVA SOFIE" + Name() + "Op invalid Autopad value : " + fAttrAutopad);
0166 }
0167
0168 if (fDim < 3) fAttrPads.resize(6, 0);
0169
0170 size_t input1 = input[0][2];
0171 size_t input2 = (fDim > 1) ? input[0][3] : 1;
0172 size_t input3 = (fDim > 2) ? input[0][4] : 1;
0173
0174
0175 auto poolOutDim = [this](size_t in, size_t pad, size_t kern, size_t stride) -> size_t {
0176 size_t n = in + pad - kern;
0177 return (fAttrCeilMode ? (n + stride - 1) / stride : n / stride) + 1;
0178 };
0179
0180 size_t pad1 = fAttrPads[0] + fAttrPads[i1];
0181 size_t output1 = poolOutDim(input1, pad1, fAttrKernelShape[0], fAttrStrides[0]);
0182
0183 size_t batch_size = input[0][0];
0184 size_t output_channels = input[0][1];
0185
0186 std::vector<std::vector<size_t>> ret({{ batch_size, output_channels, output1 }});
0187
0188 if (fDim == 1)
0189 return ret;
0190
0191 size_t pad2 = fAttrPads[1] + fAttrPads[i2];
0192 size_t output2 = poolOutDim(input2, pad2, fAttrKernelShape[1], fAttrStrides[1]);
0193
0194 ret[0].push_back(output2);
0195 if (fDim == 2)
0196 return ret;
0197
0198 size_t pad3 = fAttrPads[2] + fAttrPads[i3];
0199 size_t output3 = poolOutDim(input3, pad3, fAttrKernelShape[2], fAttrStrides[2]);
0200
0201
0202 ret[0].push_back(output3);
0203 return ret;
0204 }
0205
0206 void Initialize(RModel& model) override {
0207
0208 fUseSession = model.UseSession();
0209
0210 if (!model.CheckIfTensorAlreadyExist(fNX)) {
0211 throw
0212 std::runtime_error("TMVA SOFIE Pool op Input Tensor " + fNX + " is not found in model");
0213 }
0214 fShapeX = model.GetTensorShape(fNX);
0215 if (fShapeX.size() < 3 || fShapeX.size() > 5) {
0216 std::cout << fNX << " : " << ConvertShapeToString(fShapeX) << std::endl;
0217 throw
0218 std::runtime_error("TMVA SOFIE Pool Op input data tensor" + fNX + " is not of 3,4 or 5 dimensions");
0219 }
0220 fDim = fShapeX.size() - 2;
0221
0222 if (fPoolMode == GlobalAveragePool) {
0223 fPoolMode = AveragePool;
0224 fAttrKernelShape.resize(3);
0225 fAttrKernelShape[0] = fShapeX[2];
0226 if (fDim > 1)
0227 fAttrKernelShape[1] = fShapeX[3];
0228 if (fDim > 2)
0229 fAttrKernelShape[2] = fShapeX[4];
0230 fAttrAutopad = "VALID";
0231 fAttrPads = {0, 0, 0, 0, 0, 0 };
0232 assert(fAttrStrides.empty());
0233 }
0234
0235 fShapeY = ShapeInference({fShapeX})[0];
0236 model.AddIntermediateTensor(fNY, model.GetTensorType(fNX), fShapeY);
0237
0238
0239 if (fPoolMode == MaxPool) model.AddNeededStdLib("cmath");
0240
0241 }
0242
0243 std::string GenerateInitCode() override {
0244 std::stringstream out;
0245 return out.str();
0246 }
0247
0248
0249 virtual std::string GenerateSessionMembersCode(std::string opName) override {
0250 opName = "op_" + opName;
0251 std::stringstream out;
0252
0253 if(fDim == 1){
0254 out << "std::vector<" << fType << "> fVec_" << opName << "_xpad = std::vector<" << fType << ">("
0255 << fShapeX[1] * (fShapeX[2] + fAttrPads[0] + fAttrPads[2]) << ");\n";
0256 }
0257 else if(fDim == 2){
0258 out << "std::vector<" << fType << "> fVec_" << opName << "_xpad = std::vector<" << fType << ">("
0259 << fShapeX[1] * (fShapeX[2] + fAttrPads[0] + fAttrPads[2]) * (fShapeX[3] + fAttrPads[1] + fAttrPads[3])
0260 << ");\n";
0261 }
0262 else{
0263 out << "std::vector<" << fType << "> fVec_" << opName << "_xpad = std::vector<" << fType << ">("
0264 << fShapeX[1] * (fShapeX[2] + fAttrPads[0] + fAttrPads[2]) * (fShapeX[3] + fAttrPads[1] + fAttrPads[3]) *
0265 (fShapeX[4] + fAttrPads[2] + fAttrPads[4]) << ");\n";
0266 }
0267
0268 return out.str();
0269 }
0270
0271 std::string Generate(std::string OpName) override {
0272 OpName = "op_" + OpName;
0273
0274 if (fShapeX.empty() || fShapeY.empty()) {
0275 throw std::runtime_error("TMVA SOFIE Pool Op called to Generate without being initialized first");
0276 }
0277
0278 std::stringstream out;
0279
0280 out << "\n//---- operator " << Name() << " " << OpName << "\n";
0281 out << "{\n";
0282
0283 assert(fShapeX[0] == fShapeY[0]);
0284 assert(fShapeX[1] == fShapeY[1]);
0285 assert(fAttrPads.size() == 6);
0286 assert(fAttrKernelShape.size() == 3);
0287
0288 int hmin = - fAttrPads[0];
0289
0290 int hmax = fShapeX[2] + fAttrPads[fDim] - fAttrKernelShape[0] + (fAttrCeilMode ? (int)fAttrStrides[0] : 1);
0291 int wmin,wmax,dmin,dmax;
0292
0293 if(fDim >= 2){
0294 wmin = -fAttrPads[1];
0295 wmax = fShapeX[3] + fAttrPads[fDim + 1] - fAttrKernelShape[1] + (fAttrCeilMode ? (int)fAttrStrides[1] : 1);
0296 }
0297 else{
0298 wmin=1;
0299 wmax=1;
0300 }
0301 if(fDim == 3){
0302 dmin = -fAttrPads[2];
0303 dmax = fShapeX[4] + fAttrPads[fDim + 2] - fAttrKernelShape[2] + (fAttrCeilMode ? (int)fAttrStrides[2] : 1);
0304 }
0305 else{
0306 dmin=1;
0307 dmax=1;
0308 }
0309 out << SP << "constexpr int hsize = " << fShapeX[2] << ";\n";
0310 out << SP << "constexpr int hmin = " << hmin << ";\n";
0311 out << SP << "constexpr int hmax = " << hmax << ";\n";
0312 out << SP << "constexpr int kh = " << fAttrKernelShape[0] << ";\n";
0313 if (fDim > 1) {
0314 size_t wsize = fShapeX[3];
0315 out << SP << "constexpr int wsize = " << wsize << ";\n";
0316 out << SP << "constexpr int wmin = " << wmin << ";\n";
0317 out << SP << "constexpr int wmax = " << wmax << ";\n";
0318 out << SP << "constexpr int kw = " << fAttrKernelShape[1] << ";\n";
0319 if (fDim > 2) {
0320 size_t dsize = fShapeX[4];
0321 out << SP << "constexpr int dsize = " << dsize << ";\n";
0322 out << SP << "constexpr int dwsize = " << dsize*wsize << ";\n";
0323 out << SP << "constexpr int dmin = " << dmin << ";\n";
0324 out << SP << "constexpr int dmax = " << dmax << ";\n";
0325 out << SP << "constexpr int kd = " << fAttrKernelShape[2] << ";\n";
0326 }
0327 }
0328
0329
0330 bool doPadding = false;
0331 for ( auto & e : fAttrPads)
0332 doPadding |= (e > 0);
0333
0334
0335 if(fDim==1){
0336
0337 out << SP << "size_t outIndex = 0;\n";
0338 out << SP << "for (size_t n = 0; n < " << fShapeX[0]*fShapeX[1] << "; n++) {\n";
0339 out << SP << SP << "size_t inputOffset = n*" << fShapeX[2] << ";\n";
0340 out << SP << SP << "for (int i = hmin; i < hmax; i+=" << fAttrStrides[0] << ") {\n";
0341
0342 if (fPoolMode == MaxPool)
0343 out << SP << SP << SP << SP << "float value = -INFINITY;\n";
0344 else if (fPoolMode == AveragePool) {
0345 out << SP << SP << SP << SP << "float value = 0;\n";
0346 if (fAttrCountIncludePad == 0 && doPadding)
0347 out << SP << SP << SP << SP << "int nsum = 0;\n";
0348 else
0349 out << SP << SP << SP << SP << "constexpr int nsum = kh;\n";
0350 }
0351
0352 out << SP << SP << SP << SP << "for (int l = i; l < i + kh; l++) {\n";
0353 out << SP << SP << SP << SP << SP << "if (l < 0 || l >= hsize) continue;\n";
0354 out << SP << SP << SP << SP << SP << SP << "int index = inputOffset + l;\n";
0355 if (fPoolMode == MaxPool) {
0356 out << SP << SP << SP << SP << SP << SP << "auto xval = tensor_" << fNX << "[index];\n";
0357 out << SP << SP << SP << SP << SP << SP << "if (xval > value) value = xval;\n";
0358 }
0359 else if (fPoolMode == AveragePool) {
0360
0361 out << SP << SP << SP << SP << SP << SP << "value += tensor_" << fNX << "[index];\n";
0362 if (fAttrCountIncludePad == 0 && doPadding)
0363
0364 out << SP << SP << SP << SP << SP << SP << "nsum++;\n";
0365 }
0366 out << SP << SP << SP << SP << SP << "}\n";
0367 if (fPoolMode == AveragePool) {
0368
0369 out << SP << SP << SP << SP << "value /= float(nsum);\n";
0370 }
0371
0372 out << SP << SP << SP << SP << "tensor_" << fNY << "[outIndex++] = value;\n";
0373
0374 out << SP << SP << "}\n";
0375 out << SP << "}\n";
0376 }
0377 else if(fDim==2){
0378
0379 out << SP << "size_t outIndex = 0;\n";
0380 out << SP << "for (size_t n = 0; n < " << fShapeX[0]*fShapeX[1] << "; n++) {\n";
0381 out << SP << SP << "size_t inputOffset = n*" << fShapeX[2]*fShapeX[3] << ";\n";
0382 out << SP << SP << "for (int i = hmin; i < hmax; i+=" << fAttrStrides[0] << ") {\n";
0383 out << SP << SP << SP << "for (int j = wmin; j < wmax; j+=" << fAttrStrides[1] << ") {\n";
0384
0385 if (fPoolMode == MaxPool)
0386 out << SP << SP << SP << SP << "float value = -INFINITY;\n";
0387 else if (fPoolMode == AveragePool) {
0388 out << SP << SP << SP << SP << "float value = 0;\n";
0389 if (fAttrCountIncludePad == 0 && doPadding)
0390 out << SP << SP << SP << SP << "int nsum = 0;\n";
0391 else
0392 out << SP << SP << SP << SP << "constexpr int nsum = kw*kh;\n";
0393 }
0394
0395 out << SP << SP << SP << SP << "for (int l = i; l < i + kh; l++) {\n";
0396 out << SP << SP << SP << SP << SP << "if (l < 0 || l >= hsize) continue;\n";
0397
0398 out << SP << SP << SP << SP << SP << "for (int m = j; m < j + kw; m++) {\n";
0399 out << SP << SP << SP << SP << SP << SP << "if (m < 0 || m >= wsize) continue;\n";
0400 out << SP << SP << SP << SP << SP << SP << SP << "int index = inputOffset + l*wsize + m;\n";
0401 if (fPoolMode == MaxPool) {
0402 out << SP << SP << SP << SP << SP << SP << SP << "auto xval = tensor_" << fNX << "[index];\n";
0403 out << SP << SP << SP << SP << SP << SP << SP << "if (xval > value) value = xval;\n";
0404 }
0405 else if (fPoolMode == AveragePool) {
0406
0407 out << SP << SP << SP << SP << SP << SP << SP << "value += tensor_" << fNX << "[index];\n";
0408 if (fAttrCountIncludePad == 0 && doPadding)
0409
0410 out << SP << SP << SP << SP << SP << SP << SP << "nsum++;\n";
0411 }
0412 out << SP << SP << SP << SP << SP << SP << "}\n";
0413 out << SP << SP << SP << SP << SP << "}\n";
0414 if (fPoolMode == AveragePool) {
0415
0416 out << SP << SP << SP << SP << "value /= float(nsum);\n";
0417 }
0418 out << SP << SP << SP << SP << "tensor_" << fNY << "[outIndex++] = value;\n";
0419 out << SP << SP << SP << "}\n";
0420 out << SP << SP << "}\n";
0421 out << SP << "}\n";
0422 }
0423 else if(fDim==3){
0424
0425 out << SP << "size_t outIndex = 0;\n";
0426 out << SP << "for (size_t n = 0; n < " << fShapeX[0]*fShapeX[1] << "; n++) {\n";
0427 out << SP << SP << "size_t inputOffset = n*" << fShapeX[2]*fShapeX[3]*fShapeX[4] << ";\n";
0428 out << SP << SP << "for (int i = hmin; i < hmax; i+=" << fAttrStrides[0] << ") {\n";
0429 out << SP << SP << SP << "for (int j = wmin; j < wmax; j+=" << fAttrStrides[1] << ") {\n";
0430 out << SP << SP << SP << SP << "for (int k = dmin; k < dmax; k+=" << fAttrStrides[2] << ") {\n";
0431
0432 if (fPoolMode == MaxPool)
0433 out << SP << SP << SP << SP << "float value = -INFINITY;\n";
0434 else if (fPoolMode == AveragePool) {
0435 out << SP << SP << SP << SP << "float value = 0;\n";
0436 if (fAttrCountIncludePad == 0 && doPadding)
0437 out << SP << SP << SP << SP << "int nsum = 0;\n";
0438 else
0439 out << SP << SP << SP << SP << "constexpr int nsum = kw*kh*kd;\n";
0440 }
0441
0442 out << SP << SP << SP << SP << "for (int l = i; l < i + kh; l++) {\n";
0443 out << SP << SP << SP << SP << SP << "if (l < 0 || l >= hsize) continue;\n";
0444
0445 out << SP << SP << SP << SP << SP << "for (int m = j; m < j + kw; m++) {\n";
0446 out << SP << SP << SP << SP << SP << SP << "if (m < 0 || m >= wsize) continue;\n";
0447
0448 out << SP << SP << SP << SP << SP << SP << "for (int p = k; p < k + kd; p++) {\n";
0449 out << SP << SP << SP << SP << SP << SP << SP << "if (p < 0 || p >= dsize) continue;\n";
0450 out << SP << SP << SP << SP << SP << SP << SP << SP << "int index = inputOffset + l*dwsize + m*dsize + p;\n";
0451
0452 if (fPoolMode == MaxPool) {
0453 out << SP << SP << SP << SP << SP << SP << SP << SP << "auto xval = tensor_" << fNX << "[index];\n";
0454 out << SP << SP << SP << SP << SP << SP << SP << SP << "if (xval > value) value = xval;\n";
0455 }
0456 else if (fPoolMode == AveragePool) {
0457
0458 out << SP << SP << SP << SP << SP << SP << SP << SP << "value += tensor_" << fNX << "[index];\n";
0459 if (fAttrCountIncludePad == 0 && doPadding)
0460
0461 out << SP << SP << SP << SP << SP << SP << SP << SP << "nsum++;\n";
0462 }
0463 out << SP << SP << SP << SP << SP << SP << "}\n";
0464 out << SP << SP << SP << SP << SP << "}\n";
0465 out << SP << SP << SP << SP << "}\n";
0466 if (fPoolMode == AveragePool) {
0467
0468 out << SP << SP << SP << SP << "value /= float(nsum);\n";
0469 }
0470
0471 out << SP << SP << SP << SP << "tensor_" << fNY << "[outIndex++] = value;\n";
0472 out << SP << SP << SP << SP << "}\n" ;
0473 out << SP << SP << SP << "}\n";
0474 out << SP << SP << "}\n";
0475 out << SP << "}\n";
0476 }
0477
0478 out << SP << "}\n";
0479
0480
0481 return out.str();
0482 }
0483 };
0484
0485 }
0486 }
0487 }
0488
0489
0490 #endif