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
0015 class ROperator_Softmax final : public ROperator {
0016
0017 private:
0018 bool fLogSoftmax;
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;
0043 return ret;
0044 }
0045
0046 void Initialize(RModel& model) override {
0047 if (model.CheckIfTensorAlreadyExist(fNX) ==
0048 false) {
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
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
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
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
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
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 }
0191 }
0192 }
0193
0194 #endif