Back to home page

EIC code displayed by LXR

 
 

    


File indexing completed on 2026-09-19 09:36:57

0001 #ifndef TMVA_SOFIE_ROPERATOR_Softmax
0002 #define TMVA_SOFIE_ROPERATOR_Softmax
0003 
0004 #include "TMVA/SOFIE_common.hxx"
0005 #include "TMVA/ROperator.hxx"
0006 #include "TMVA/RModel.hxx"
0007 
0008 #include <sstream>
0009 
0010 namespace TMVA {
0011 namespace Experimental {
0012 namespace SOFIE {
0013 
0014 // implement Softmax and LogSoftmax
0015 class ROperator_Softmax final : public ROperator {
0016 
0017 private:
0018    bool fLogSoftmax;  // for the logsoftmax case
0019    bool fUseVDT = false;
0020    int64_t fAttrAxis;
0021 
0022    std::string fNX;
0023    std::string fNY;
0024    std::vector<Dim> fShape;
0025 
0026    std::string fType;
0027 
0028 public:
0029    ROperator_Softmax() {}
0030    ROperator_Softmax(int64_t attr_axis, std::string nameX, std::string nameY, bool logSoftmax = false)
0031       : fLogSoftmax(logSoftmax),
0032       fAttrAxis(attr_axis), fNX(UTILITY::Clean_name(nameX)), fNY(UTILITY::Clean_name(nameY))
0033 
0034    {
0035          fInputTensorNames = { fNX };
0036          fOutputTensorNames = { fNY };
0037    }
0038 
0039    std::vector<ETensorType> TypeInference(std::vector<ETensorType> input) override { return input; }
0040 
0041    std::vector<std::vector<size_t>> ShapeInference(std::vector<std::vector<size_t>> input) override {
0042       auto ret = input; // suggest copy to compiler
0043       return ret;
0044    }
0045 
0046    void Initialize(RModel& model) override {
0047       if (model.CheckIfTensorAlreadyExist(fNX) ==
0048           false) { // input must be a graph input, or already initialized intermediate tensor
0049          throw std::runtime_error("TMVA SOFIE Softmax Op Input Tensor is not found in model");
0050       }
0051       fShape = model.GetDimTensorShape(fNX);
0052       model.AddIntermediateTensor(fNY, model.GetTensorType(fNX), fShape);
0053       fType = ConvertTypeToString(model.GetTensorType(fNX));
0054       if (model.Verbose()) {
0055          std::cout << "Softmax -> " << fNY << " " << ConvertDimShapeToString(fShape) << std::endl;
0056       }
0057       fUseVDT = model.UseVDT();
0058       if (fUseVDT) {
0059          model.AddNeededCustomHeader("vdt/exp.h");
0060          if (fLogSoftmax)
0061             model.AddNeededCustomHeader("vdt/log.h");
0062       }
0063    }
0064 
0065    std::string Generate(std::string OpName) override {
0066       OpName = "op_" + OpName;
0067       if (fShape.empty()) {
0068          throw std::runtime_error("TMVA SOFIE Operator Softmax called to Generate without being initialized first");
0069       }
0070       std::stringstream out;
0071       size_t size = fShape.size();
0072       auto length_str = ConvertDimShapeToLength(fShape);
0073       size_t axis = fAttrAxis < 0 ? size + fAttrAxis : fAttrAxis;
0074 
0075       std::string expFunction = (fUseVDT) ? "vdt::fast_expf" : "std::exp";
0076       std::string logFunction = (fUseVDT) ? "vdt::fast_logf" : "std::log";
0077 
0078       // Check if this is the special case where memory is contiguous.
0079       if (axis == size - 1) {
0080          std::string axis_size = fShape[axis].GetVal();
0081          std::string num_rows;
0082          if (IsInteger(length_str) && IsInteger(axis_size)) {
0083             num_rows = std::to_string(std::stoul(length_str) / std::stoul(axis_size));
0084          } else {
0085             num_rows = "(" + length_str + ") / (" + axis_size + ")";
0086          }
0087 
0088          out << "\n" << SP << "//------ SOFTMAX - " << size << "  " << length_str << "  " << axis << "\n";
0089          out << SP << "for (int i = 0; i < " << num_rows << "; ++i) {\n";
0090          out << SP << SP << "size_t offset = i * " << axis_size << ";\n";
0091          out << SP << SP << fType << " const * x_ptr = &tensor_" << fNX << "[offset];\n";
0092          out << SP << SP << fType << " * y_ptr = &tensor_" << fNY << "[offset];\n";
0093 
0094          out << SP << SP << fType << " vmax = x_ptr[0];\n";
0095          out << SP << SP << "for (int j = 1; j < " << axis_size << "; ++j) {\n";
0096          out << SP << SP << SP << "if (x_ptr[j] > vmax) vmax = x_ptr[j];\n";
0097          out << SP << SP << "}\n";
0098 
0099          out << SP << SP << fType << " sum = 0.0;\n";
0100          out << SP << SP << "for (int j = 0; j < " << axis_size << "; ++j) {\n";
0101          out << SP << SP << SP << "y_ptr[j] = " << expFunction << "(x_ptr[j] - vmax);\n";
0102          out << SP << SP << SP << "sum += y_ptr[j];\n";
0103          out << SP << SP << "}\n";
0104 
0105          out << SP << SP << fType << " inv_sum = 1.0f / sum;\n";
0106          out << SP << SP << "for (int j = 0; j < " << axis_size << "; ++j) {\n";
0107          out << SP << SP << SP << "y_ptr[j] *= inv_sum;\n";
0108          if (fLogSoftmax)
0109             out << SP << SP << SP << "y_ptr[j] = " << logFunction << "(y_ptr[j]);\n";
0110          out << SP << SP << "}\n";
0111          out << SP << "}\n";
0112 
0113       } else {
0114          auto stride = UTILITY::ComputeStrideFromShape(fShape);
0115          size_t k = 0;
0116          std::vector<std::string> l(size);
0117          for (size_t i = 0; i < size; i++) {
0118             if (i != axis) {
0119                for (size_t j = 0; j < k; j++) out << SP;
0120                l[i] = std::string("i") + std::to_string(i);
0121                out << "for (int " << l[i] << " = 0; " << l[i] << " < " << fShape[i] << "; " << l[i] << "++) {\n";
0122                k++;
0123             }
0124          }
0125          for (size_t j = 0; j < size-1; j++) out << SP;
0126          out << fType << " sum = 0.;\n";
0127          for (size_t j = 0; j < size-1; j++) out << SP;
0128          out << "size_t index = ";
0129          bool first = true;
0130          for (size_t i = 0; i < size; i++) {
0131             if (i == axis) continue;
0132             if (!first) out << " + ";
0133             if (stride[i].GetVal() != "1")
0134                out << stride[i] << "*";
0135             out << l[i];
0136             first = false;
0137          }
0138          out << ";\n";
0139          // find maximum looping along reduced axis
0140          for (size_t j = 0; j < size-1; j++) out << SP;
0141          out << fType << " vmax = tensor_" << fNX << "[index];\n";
0142          for (size_t j = 0; j < size-1; j++) out << SP;
0143          out << "for (int i = 1; i < " << fShape[axis] << "; i++) {\n";
0144          for (size_t j = 0; j < size; j++) out << SP;
0145          out << fType << " x = tensor_" << fNX << "[index + i";
0146          if (stride[axis].GetVal() != "1") out << "*(" << stride[axis] << ")";
0147          out << "];\n";
0148          for (size_t j = 0; j < size; j++) out << SP;
0149          out << "if (x > vmax) vmax = x;\n";
0150          for (size_t j = 0; j < size-1; j++) out << SP;
0151          out << "}\n";
0152          // compute softmax
0153          for (size_t j = 0; j < size-1; j++) out << SP;
0154          out << "for (int i = 0; i < " << fShape[axis] << "; i++) {\n";
0155          for (size_t j = 0; j < size; j++) out << SP;
0156          out << "size_t id = index + i";
0157          if (stride[axis].GetVal() != "1") out << "*(" << stride[axis] << ")";
0158          out << ";\n";
0159          for (size_t j = 0; j < size; j++) out << SP;
0160          out << "tensor_" << fNY << "[id] = " << expFunction << "(tensor_" << fNX << "[id] - vmax);\n";
0161          for (size_t j = 0; j < size; j++) out << SP;
0162          out << "sum += tensor_" << fNY << "[id];\n";
0163          for (size_t j = 0; j < size-1; j++) out << SP;
0164          out << "}\n";
0165          // normalize
0166          for (size_t j = 0; j < size-1; j++) out << SP;
0167          out << "for (int i = 0; i < " << fShape[axis] << "; i++) {\n";
0168          for (size_t j = 0; j < size; j++) out << SP;
0169          out << "size_t id = index + i";
0170          if (stride[axis].GetVal() != "1") out << "*(" << stride[axis] << ");\n";
0171          for (size_t j = 0; j < size; j++) out << SP;
0172          out << "tensor_" << fNY << "[id] /= sum;\n";
0173          if (fLogSoftmax) {
0174             for (size_t j = 0; j < size; j++) out << SP;
0175             out << "tensor_" << fNY << "[id] = " << logFunction << "(tensor_" << fNY << "[id]);\n";
0176          }
0177          for (size_t j = 0; j < size-1; j++) out << SP;
0178          out << "}\n";
0179          //end loops
0180          for (int i = static_cast<int>(k) - 1; i >= 0; i--) {
0181             for (int j = 0; j < i; j++) out << SP;
0182             out << "}\n";
0183          }
0184       }
0185       return out.str();
0186    }
0187    std::vector<std::string> GetStdLibs() override { return { std::string("cmath") }; }
0188 };
0189 
0190 } // namespace SOFIE
0191 } // namespace Experimental
0192 } // namespace TMVA
0193 
0194 #endif // TMVA_SOFIE_ROPERATOR_Softmax