File indexing completed on 2026-10-04 09:24:21
0001 #ifndef TMVA_SOFIE_ROPERATOR_TRANSPOSE
0002 #define TMVA_SOFIE_ROPERATOR_TRANSPOSE
0003
0004 #include "TMVA/SOFIE_common.hxx"
0005 #include "TMVA/ROperator.hxx"
0006 #include "TMVA/RModel.hxx"
0007
0008 #include <sstream>
0009 #include <cassert>
0010
0011 namespace TMVA{
0012 namespace Experimental{
0013 namespace SOFIE{
0014
0015
0016
0017 class ROperator_Transpose final : public ROperator
0018 {
0019
0020 private:
0021
0022 std::vector<int64_t> fAttrPerm;
0023
0024 std::string fNX;
0025 std::string fNY;
0026 std::vector<Dim> fShapeX;
0027 std::vector<Dim> fShapeY;
0028
0029 public:
0030
0031 ROperator_Transpose(){}
0032 ROperator_Transpose(std::vector<int64_t> attr_perm, std::string nameData, std::string nameOutput):
0033 fAttrPerm(attr_perm), fNX(UTILITY::Clean_name(nameData)), fNY(UTILITY::Clean_name(nameOutput)) {
0034 fInputTensorNames = { fNX };
0035 fOutputTensorNames = { fNY };
0036 }
0037
0038
0039 std::vector<ETensorType> TypeInference(std::vector<ETensorType> input) override {
0040 return input;
0041 }
0042
0043 std::vector<std::vector<size_t>> ShapeInference(std::vector<std::vector<size_t>> input) override {
0044 if (input.size() > 1) throw std::runtime_error("TMVA SOFIE Tranpose Op Shape Inference only need 1 input tensor");
0045 auto& data = input[0];
0046 if (fAttrPerm.size() != data.size() )
0047 throw std::runtime_error("TMVA SOFIE Tranpose Op - Invalid axes attributes");
0048
0049 std::vector<size_t> output_shape(fAttrPerm.size());
0050 for (size_t i = 0; i < fAttrPerm.size(); i++){
0051 output_shape[i] = data[fAttrPerm[i]];
0052 }
0053 std::vector<std::vector<size_t>> ret;
0054 ret.push_back(output_shape);
0055 return ret;
0056 }
0057
0058 template<class T>
0059 void ProcessInitializedTensor(RModel& model) {
0060
0061
0062 auto shapeX = ConvertShapeToInt(fShapeX);
0063 auto shapeY = ConvertShapeToInt(fShapeY);
0064 fIsOutputConstant = true;
0065
0066 auto inStrides = UTILITY::ComputeStrideFromShape(shapeX);
0067 auto outStrides = UTILITY::ComputeStrideFromShape(shapeY);
0068 size_t length = ConvertShapeToLength(shapeY);
0069 auto inputData = static_cast<T *>(model.GetInitializedTensorData(fNX).get());
0070 size_t dim = fShapeX.size();
0071 std::vector<size_t> outputIdx(dim);
0072 std::vector<T> outputData(length);
0073 for (size_t i = 0; i < length; i++) {
0074 outputIdx[0] = i / outStrides[0];
0075 for (size_t j = 1; j < dim; j++) {
0076 outputIdx[j] = (i % outStrides[j - 1]) / outStrides[j];
0077 }
0078
0079 size_t inputIndex = 0;
0080 for (size_t j = 0; j < dim; j++) {
0081
0082 int k = std::find(fAttrPerm.begin(), fAttrPerm.end(), j) - fAttrPerm.begin();
0083 inputIndex += outputIdx[k] * inStrides[j];
0084 }
0085 outputData[i] = inputData[inputIndex];
0086 }
0087 model.AddConstantTensor<T>(fNY, shapeY, outputData.data());
0088 if (model.Verbose()) {
0089 std::cout << "Transpose: output is a constant tensor " << ConvertShapeToString(shapeY) << " : "
0090 << ConvertValuesToString(outputData) << std::endl;
0091 }
0092 }
0093
0094 void Initialize(RModel& model) override {
0095 if (model.CheckIfTensorAlreadyExist(fNX) == false){
0096 std::cout<<"Input tensor for transpose: "<<fNX<<'\n';
0097 throw std::runtime_error("TMVA SOFIE Tranpose Op Input Tensor is not found in model");
0098 }
0099 fShapeX = model.GetDimTensorShape(fNX);
0100 if (fAttrPerm.empty()){
0101 fAttrPerm.reserve(fShapeX.size());
0102 for (int i = fShapeX.size() - 1; i >= 0; i--){
0103 fAttrPerm.push_back(i);
0104 }
0105 }
0106
0107
0108 if (fAttrPerm.size() != fShapeX.size() )
0109 throw std::runtime_error("TMVA SOFIE Tranpose Op - Invalid axes attributes");
0110
0111 fShapeY.resize(fAttrPerm.size());
0112 for (size_t i = 0; i < fAttrPerm.size(); i++){
0113 fShapeY[i] = fShapeX[fAttrPerm[i]];
0114 }
0115
0116 if (model.IsInitializedTensor(fNX) ) {
0117 auto type = model.GetTensorType(fNX);
0118 switch(type) {
0119 case ETensorType::FLOAT:
0120 ProcessInitializedTensor<float>(model);
0121 break;
0122 case ETensorType::INT64:
0123 ProcessInitializedTensor<int64_t>(model);
0124 break;
0125 case ETensorType::BOOL:
0126 ProcessInitializedTensor<uint8_t>(model);
0127 break;
0128 case ETensorType::UINT8:
0129 ProcessInitializedTensor<uint8_t>(model);
0130 break;
0131 default:
0132 std::cout << "Transpose - no support for initialized tensor of type " << ConvertTypeToString(type) << std::endl;
0133 }
0134 return;
0135 }
0136
0137 model.AddIntermediateTensor(fNY, model.GetTensorType(fNX), fShapeY);
0138 if (model.Verbose()) {
0139 std::cout << "Transpose ---> " << fNY << " " << ConvertDimShapeToString(fShapeY) << std::endl;
0140 }
0141 }
0142
0143 std::string Generate(std::string opName) override {
0144 if (fIsOutputConstant) return "";
0145 opName = "op_" + opName;
0146 if (fShapeX.empty() || fShapeY.empty()){
0147 throw std::runtime_error("TMVA SOFIE Transpose Op called to Generate without being initialized first");
0148 }
0149 auto stridesX = UTILITY::ComputeStrideFromShape(fShapeX);
0150 auto stridesY = UTILITY::ComputeStrideFromShape(fShapeY);
0151
0152 auto intShapeX = ConvertShapeToInt(fShapeX);
0153 size_t rank = fShapeX.size();
0154 bool isDynamic = (intShapeX.empty() && rank > 0);
0155
0156 std::string constQualifier = (isDynamic) ? "const" : "constexpr";
0157
0158 std::stringstream out;
0159
0160 out << SP << "///------- Transpose operator " << opName << ConvertDimShapeToString(fShapeX)
0161 << " --> " << ConvertDimShapeToString(fShapeY) << std::endl;
0162
0163
0164
0165
0166
0167
0168 out << SP << "{\n";
0169 out << SP << SP << "// Pre-baked input strides (row-major)\n";
0170 out << SP << SP << constQualifier << " size_t " << opName << "_strX[] = {";
0171 for (size_t i = 0; i < rank; ++i)
0172 out << stridesX[i] << (i + 1 < rank ? ", " : "");
0173 out << "};\n";
0174
0175 out << SP << SP << "// Pre-baked output strides (row-major)\n";
0176 out << SP << SP << constQualifier << " size_t " << opName << "_strY[] = {";
0177 for (size_t i = 0; i < rank; ++i)
0178 out << stridesY[i] << (i + 1 < rank ? ", " : "");
0179 out << "};\n\n";
0180
0181
0182 bool innerContiguous = (fAttrPerm.back() == (int64_t) (rank - 1));
0183 size_t outerRank = innerContiguous ? rank - 1 : rank;
0184 size_t innerSize = innerContiguous ? (isDynamic ? 0 : intShapeX[fAttrPerm[rank - 1]])
0185 : 1;
0186
0187 if (innerContiguous && !isDynamic && innerSize > 1) {
0188
0189 out << SP << SP
0190 << "// Fast path: last permuted axis is contiguous in source\n";
0191 out << SP << SP
0192 << "// Inner " << innerSize << " elements copied with pointer arithmetic\n";
0193
0194
0195 EmitNestedLoops(out, outerRank, fShapeY);
0196
0197
0198 out << SP << SP << SP << "size_t src_off = ";
0199 for (size_t i = 0; i < outerRank; ++i) {
0200 out << "idx_" << i << " * " << opName << "_strX["
0201 << fAttrPerm[i] << "]";
0202 if (i + 1 < outerRank) out << " + ";
0203 }
0204 out << ";\n";
0205
0206 out << SP << SP << SP << "size_t dst_off = ";
0207 for (size_t i = 0; i < outerRank; ++i) {
0208 out << "idx_" << i << " * " << opName << "_strY[" << i << "]";
0209 if (i + 1 < outerRank) out << " + ";
0210 }
0211 out << ";\n";
0212
0213
0214 out << SP << SP << SP
0215 << "std::copy(tensor_" << fNX << " + src_off, "
0216 << "tensor_" << fNX << " + src_off + " << innerSize << ", "
0217 << "tensor_" << fNY << " + dst_off);\n";
0218
0219 CloseNestedLoops(out, outerRank);
0220
0221 } else {
0222
0223
0224 out << SP << SP << "// General N-D transpose\n";
0225
0226 EmitNestedLoops(out, rank, fShapeY);
0227
0228
0229 out << SP << SP << SP << "size_t src_idx = ";
0230 for (size_t i = 0; i < rank; ++i) {
0231 out << "idx_" << i << " * " << opName << "_strX[" << fAttrPerm[i] << "]";
0232 if (i + 1 < rank) out << " + ";
0233 }
0234 out << ";\n";
0235
0236
0237 out << SP << SP << SP << "size_t dst_idx = ";
0238 for (size_t i = 0; i < rank; ++i) {
0239 out << "idx_" << i << " * " << opName << "_strY[" << i << "]";
0240 if (i + 1 < rank) out << " + ";
0241 }
0242 out << ";\n";
0243
0244 out << SP << SP << SP
0245 << "tensor_" << fNY << "[dst_idx] = "
0246 << "tensor_" << fNX << "[src_idx];\n";
0247
0248 CloseNestedLoops(out, rank);
0249
0250 }
0251
0252 out << SP << "}\n";
0253 return out.str();
0254 }
0255
0256
0257 };
0258
0259 }
0260 }
0261 }
0262
0263
0264 #endif