已合并
feat: add SoftMax AscIR operator for V2 (A5/v35) #1222
feat: add SoftMax AscIR operator for V2 (A5/v35) #1222
已合并
WangYanMale创建于 7月7日
共 18 个文件变更+279-14
@@ -177,6 +177,7 @@ inline const std::string kDuplicate = "Duplicate";
177inline const std::string kGatherMask = "GatherMask";177inline const std::string kGatherMask = "GatherMask";
178inline const std::string kMaxs = "Maxs";178inline const std::string kMaxs = "Maxs";
179inline const std::string kMax = "Max";179inline const std::string kMax = "Max";
180+inline const std::string kSoftmax = "Softmax";
180inline const std::string kArgMax = "ArgMax";181inline const std::string kArgMax = "ArgMax";
181inline const std::string kArgMaxMultiRPhase1 = "ArgMaxMultiRPhase1";182inline const std::string kArgMaxMultiRPhase1 = "ArgMaxMultiRPhase1";
182inline const std::string kArgMaxMultiRPhase2 = "ArgMaxMultiRPhase2";183inline const std::string kArgMaxMultiRPhase2 = "ArgMaxMultiRPhase2";
@@ -237,13 +238,14 @@ inline const std::string kMul = "Mul";
237inline const std::string kNeg = "Neg";238inline const std::string kNeg = "Neg";
238inline const std::string kReciprocal = "Reciprocal";239inline const std::string kReciprocal = "Reciprocal";
239inline const std::string kRelu = "Relu";240inline const std::string kRelu = "Relu";
240-inline const std::string kReduceAll = "ReduceAll"; // All241+inline const std::string kReduceAll = "ReduceAll"; // All
241-inline const std::string kReduceAny = "ReduceAny"; // Any242+inline const std::string kReduceAny = "ReduceAny"; // Any
242-inline const std::string kReduceMax = "ReduceMax"; // Max243+inline const std::string kReduceMax = "ReduceMax"; // Max
243-inline const std::string kReduceMean = "ReduceMean"; // Mean244+inline const std::string kReduceSoftmax = "ReduceSoftmax"; // Softmax
244-inline const std::string kReduceMin = "ReduceMin"; // Min245+inline const std::string kReduceMean = "ReduceMean"; // Mean
245-inline const std::string kReduceSum = "ReduceSum"; // Sum246+inline const std::string kReduceMin = "ReduceMin"; // Min
246-inline const std::string kReduceProd = "ReduceProd"; // Prod247+inline const std::string kReduceSum = "ReduceSum"; // Sum
248+inline const std::string kReduceProd = "ReduceProd"; // Prod
247inline const std::string kAll = "All";249inline const std::string kAll = "All";
248inline const std::string kProd = "Prod";250inline const std::string kProd = "Prod";
249inline const std::string kAny = "Any";251inline const std::string kAny = "Any";
@@ -82,6 +82,7 @@ PyMODINIT_FUNC PyInit_pyautofuse(void);
82 OP(Min) \82 OP(Min) \
83 OP(Mean) \83 OP(Mean) \
84 OP(Prod) \84 OP(Prod) \
85+ OP(Softmax) \
85 OP(Ge) \86 OP(Ge) \
86 OP(Ne) \87 OP(Ne) \
87 OP(Eq) \88 OP(Eq) \
@@ -1340,6 +1340,19 @@ def ArgMax(
1340 )1340 )
1341 1341 
1342 1342 
1343+def Softmax(
1344+ owner_graph: ascir.HintGraph,
1345+ x: ascir.OpsOperatorOutput,
1346+ *,
1347+ axis: List[ascir.Axis],
1348+ size: Optional[List[ascir.SizeExpr]] = None,
1349+ stride: Optional[List[ascir.SizeExpr]] = None,
1350+) -> ascir.OpsOperatorOutput:
1351+ return _common_in_1_out_1_normal_op(
1352+ "Softmax", owner_graph, x, axis=axis, size=size, stride=stride
1353+ )
1354+ 
1355+ 
1343def ArgMaxMultiRPhase1(1356def ArgMaxMultiRPhase1(
1344 owner_graph: ascir.HintGraph,1357 owner_graph: ascir.HintGraph,
1345 x: ascir.OpsOperatorOutput,1358 x: ascir.OpsOperatorOutput,
@@ -556,9 +556,17 @@ Status TilingGroup::GenReduceTilingGroupFullLoad(af::AscNode &node, AxisGroup &a
556 std::vector<ascir::AxisId> axes;556 std::vector<ascir::AxisId> axes;
557 GE_CHK_STATUS_RET(ScheduleUtils::GetLoopAxis(node, axes), "Get loop axis failed.");557 GE_CHK_STATUS_RET(ScheduleUtils::GetLoopAxis(node, axes), "Get loop axis failed.");
558 axes_group.axes_order.resize(axes.size());558 axes_group.axes_order.resize(axes.size());
559- std::vector<ascir::SizeExpr> src_strides;559+ 
560- GE_CHK_STATUS_RET(ScheduleUtils::GetReduceInputStrides(node, src_strides), "Get loop strides failed.");560+ if (node.GetType() == "Softmax") {
561- axes_group.n_group = CalcReduceAxes(src_strides, node.outputs[0].attr.strides, axes);561+ if (!axes.empty()) {
562+ axes_group.n_group = {axes.back()};
563+ }
564+ } else {
565+ std::vector<ascir::SizeExpr> src_strides;
566+ GE_CHK_STATUS_RET(ScheduleUtils::GetReduceInputStrides(node, src_strides), "Get loop strides failed.");
567+ axes_group.n_group = CalcReduceAxes(src_strides, node.outputs[0].attr.strides, axes);
568+ }
569+ 
562 int64_t y_order_index = 0;570 int64_t y_order_index = 0;
563 int64_t r_order_index = axes.size() - axes_group.n_group.size();571 int64_t r_order_index = axes.size() - axes_group.n_group.size();
564 for (size_t i = 0; i < axes.size(); ++i) {572 for (size_t i = 0; i < axes.size(); ++i) {
@@ -740,6 +740,14 @@ Status Optimizer::GetNonContinuousAxisPairBySpecialRule(ascir::ImplGraph &impl_g
740 non_continuous_pair.emplace(attr_axis, attr_axis + 1);740 non_continuous_pair.emplace(attr_axis, attr_axis + 1);
741 }741 }
742 }742 }
743+ 
744+ if (node->GetType() == "Softmax") {
745+ auto axis_size = static_cast<int64_t>(node->inputs[0].attr.repeats.size());
746+ if (axis_size > 1) {
747+ non_continuous_pair.emplace(
748+ axis_size - 2, axis_size - 1); // Softmax沿最后一个轴做归约,最后两个轴(axis - 2、axis - 1)不能合并
749+ }
750+ }
743 }751 }
744 return af::SUCCESS;752 return af::SUCCESS;
745}753}
@@ -533,6 +533,9 @@ bool HasReduceNodeOnPath(const af::AscNodePtr &b, const af::AscNodePtr &a) {
533bool ScheduleUtils::IsLastAxisReduce(const ascir::ImplGraph &impl_graph) {533bool ScheduleUtils::IsLastAxisReduce(const ascir::ImplGraph &impl_graph) {
534 for (const auto &node : impl_graph.GetAllNodes()) {534 for (const auto &node : impl_graph.GetAllNodes()) {
535 if (ScheduleUtils::IsReduce(node)) {535 if (ScheduleUtils::IsReduce(node)) {
536+ if (node->GetType() == "Softmax") {
537+ return true;
538+ }
536 std::vector<ascir::SizeExpr> src_strides;539 std::vector<ascir::SizeExpr> src_strides;
537 ScheduleUtils::GetReduceInputStrides(*node, src_strides);540 ScheduleUtils::GetReduceInputStrides(*node, src_strides);
538 const std::vector<ascir::SizeExpr> &dst_strides = node->outputs[0].attr.strides;541 const std::vector<ascir::SizeExpr> &dst_strides = node->outputs[0].attr.strides;
@@ -641,6 +644,11 @@ bool ScheduleUtils::IsReduceArFullLoad(const ascir::ImplGraph &implGraph) {
641 continue;644 continue;
642 }645 }
643 646 
647+ if (node->GetType() == "Softmax") {
648+ GELOGD("Reduce node %s is Softmax, force all load.", node->GetName().c_str());
649+ return true;
650+ }
651+ 
644 if (HasBroadcastDescendantNode(node)) {652 if (HasBroadcastDescendantNode(node)) {
645 GELOGD("There is a broadcast node behind the reduced node %s.", node->GetName().c_str());653 GELOGD("There is a broadcast node behind the reduced node %s.", node->GetName().c_str());
646 return true;654 return true;
@@ -79,7 +79,8 @@ const std::unordered_map<std::string, std::function<ReduceType(const char *)>> r
79 {"Min", [](const char *n) { return ReduceType{std::in_place_type_t<af::ascir_op::Min>{}, n}; }},79 {"Min", [](const char *n) { return ReduceType{std::in_place_type_t<af::ascir_op::Min>{}, n}; }},
80 {"Prod", [](const char *n) { return ReduceType{std::in_place_type_t<af::ascir_op::Prod>{}, n}; }},80 {"Prod", [](const char *n) { return ReduceType{std::in_place_type_t<af::ascir_op::Prod>{}, n}; }},
81 {"Any", [](const char *n) { return ReduceType{std::in_place_type_t<af::ascir_op::Any>{}, n}; }},81 {"Any", [](const char *n) { return ReduceType{std::in_place_type_t<af::ascir_op::Any>{}, n}; }},
82- {"All", [](const char *n) { return ReduceType{std::in_place_type_t<af::ascir_op::All>{}, n}; }}};82+ {"All", [](const char *n) { return ReduceType{std::in_place_type_t<af::ascir_op::All>{}, n}; }},
83+ {"Softmax", [](const char *n) { return ReduceType{std::in_place_type_t<af::ascir_op::Softmax>{}, n}; }}};
83 84 
84bool IsNotPartitionReduce(const af::AscNodePtr &reduce_node, size_t threshold) {85bool IsNotPartitionReduce(const af::AscNodePtr &reduce_node, size_t threshold) {
85 std::queue<af::NodePtr> node_queue;86 std::queue<af::NodePtr> node_queue;
@@ -252,6 +253,13 @@ Status ReducePartitionCaseGenerator::Generate([[maybe_unused]] ascir::HintGraph
252}253}
253 254 
254bool ReducePartitionCaseGenerator::ShouldForceAllLoad(ascir::HintGraph &graph) {255bool ReducePartitionCaseGenerator::ShouldForceAllLoad(ascir::HintGraph &graph) {
256+ for (const auto &node : graph.GetAllNodes()) {
257+ if (node->GetType() == "Softmax") {
258+ GELOGI("Graph %s contains Softmax node %s, force AllLoad", graph.GetName().c_str(), node->GetName().c_str());
259+ return true;
260+ }
261+ }
262+ 
255 if (!IsGroupGraphLegal(graph)) {263 if (!IsGroupGraphLegal(graph)) {
256 return true;264 return true;
257 }265 }
@@ -18,9 +18,9 @@
18 18 
19namespace optimize {19namespace optimize {
20 20 
21-using ReduceType =21+using ReduceType = std::variant<af::ascir_op::Max, af::ascir_op::Sum, af::ascir_op::Min, af::ascir_op::Prod,
22- std::variant<af::ascir_op::Max, af::ascir_op::Sum, af::ascir_op::Min, af::ascir_op::Prod, af::ascir_op::Any,22+ af::ascir_op::Any, af::ascir_op::All, af::ascir_op::ArgMaxMultiRPhase1,
23- af::ascir_op::All, af::ascir_op::ArgMaxMultiRPhase1, af::ascir_op::ArgMaxMultiRPhase2>;23+ af::ascir_op::ArgMaxMultiRPhase2, af::ascir_op::Softmax>;
24 24 
25class ReducePartitionCaseGenerator : public FusionCaseGenerator {25class ReducePartitionCaseGenerator : public FusionCaseGenerator {
26 public:26 public:
@@ -36,6 +36,7 @@ set(ascendc_api_regbase_extend_src
36 square.h36 square.h
37 trunc_div.h37 trunc_div.h
38 remainder.h38 remainder.h
39+ softmax_af.h
39 bessel_j_utils.h40 bessel_j_utils.h
40 bessel_j0.h41 bessel_j0.h
41 bessel_j1.h42 bessel_j1.h
@@ -0,0 +1,12 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+#ifndef __ASCENDC_API_SOFTMAX_AF_H__
11+#define __ASCENDC_API_SOFTMAX_AF_H__
12+#endif
@@ -1145,6 +1145,15 @@ REG_ASC_IR(Conv2DOffsetBias)
1145 {"T2", TensorType{DT_FLOAT16, DT_FLOAT, DT_BF16, DT_HIFLOAT8}},1145 {"T2", TensorType{DT_FLOAT16, DT_FLOAT, DT_BF16, DT_HIFLOAT8}},
1146 {"T3", TensorType{DT_INT8}}}});1146 {"T3", TensorType{DT_INT8}}}});
1147 1147 
1148+REG_ASC_IR(Softmax)
1149+ .Input("x", "T")
1150+ .Output("y", "T")
1151+ .ComputeType(ComputeType::kComputeReduce)
1152+ .Impl(v2_soc_versions,
1153+ {af::ascir::AscIrImplCreator<af::ascir::SoftmaxAscIrAttImplV2>(),
1154+ af::ascir::AscIrImplCreator<af::ascir::SoftmaxAscIrCodegenImplV2>(),
1155+ {{"T", TensorType{DT_INT8, DT_UINT8, DT_INT16, DT_INT32, DT_BF16, DT_FLOAT16, DT_FLOAT, DT_INT64}}}});
1156+ 
1148REG_ASC_IR(Sin)1157REG_ASC_IR(Sin)
1149 .Input("x", "T")1158 .Input("x", "T")
1150 .Output("y", "T")1159 .Output("y", "T")
@@ -185,6 +185,7 @@ REG_ASC_IR_ATT_V2_CLASS_DEFINE(LaguerrePolynomialL);
185REG_ASC_IR_ATT_V2_CLASS_DEFINE(LegendrePolynomialP);185REG_ASC_IR_ATT_V2_CLASS_DEFINE(LegendrePolynomialP);
186REG_ASC_IR_ATT_V2_CLASS_DEFINE(AiryAi);186REG_ASC_IR_ATT_V2_CLASS_DEFINE(AiryAi);
187REG_ASC_IR_ATT_V2_CLASS_DEFINE(Erfinv);187REG_ASC_IR_ATT_V2_CLASS_DEFINE(Erfinv);
188+REG_ASC_IR_ATT_V2_CLASS_DEFINE(Softmax);
188} // namespace ascir189} // namespace ascir
189} // namespace af190} // namespace af
190 191 
@@ -4181,6 +4181,35 @@ class RemainderAscIrCodegenImplV2 : public AscIrCodegenV2 {
4181 }4181 }
4182};4182};
4183 4183 
4184+/*********************************************************************************/
4185+class SoftmaxAscIrCodegenImplV2 : public AscIrCodegenV2 {
4186+ public:
4187+ [[nodiscard]] std::vector<std::unique_ptr<TmpBufDesc>> CalcTmpBufSize(const AscNode &node) override {
4188+ return CalcSoftmaxTmpSizeV2(node);
4189+ }
4190+ [[nodiscard]] std::string GetApiCallName() const override {
4191+ return "SoftmaxApiCall";
4192+ }
4193+ [[nodiscard]] std::string GetApiName() const override {
4194+ return "SoftmaxARFullLoadExtend";
4195+ }
4196+ [[nodiscard]] std::vector<std::string> LoadApiHeaderFiles([[maybe_unused]] bool is_dynamic) const override {
4197+ return {"softmax_af_reg_base.h"};
4198+ }
4199+ [[nodiscard]] std::vector<std::string> IncludeApiHeaderFiles() const override {
4200+ return {
4201+ "basic_api/kernel_operator_scalar_intf.h",
4202+ };
4203+ }
4204+ [[nodiscard]] bool IsNodeValid(const AscNode &node) const override {
4205+ GE_ASSERT_TRUE(!IsNodeHasScalarInput(node), "Node %s[%s] not support scalar input", node.GetTypePtr(),
4206+ node.GetNamePtr());
4207+ GE_ASSERT_SUCCESS(ValidateShapeConsistencyWithSingleOutput(node), "Node %s[%s] check shape consistency failed",
4208+ node.GetTypePtr(), node.GetNamePtr());
4209+ return true;
4210+ }
4211+};
4212+ 
4184/*********************************************************************************/4213/*********************************************************************************/
4185class UnsupportedAscIrCodegenImplV2 : public AscIrCodegenV2 {4214class UnsupportedAscIrCodegenImplV2 : public AscIrCodegenV2 {
4186 public:4215 public:
@@ -55,6 +55,7 @@ std::vector<std::unique_ptr<TmpBufDesc>> CalcModifiedBesselK1TmpSizeV2(const Asc
55std::vector<std::unique_ptr<TmpBufDesc>> CalcScaledModifiedBesselK0TmpSizeV2(const AscNode &node);55std::vector<std::unique_ptr<TmpBufDesc>> CalcScaledModifiedBesselK0TmpSizeV2(const AscNode &node);
56std::vector<std::unique_ptr<TmpBufDesc>> CalcScaledModifiedBesselK1TmpSizeV2(const AscNode &node);56std::vector<std::unique_ptr<TmpBufDesc>> CalcScaledModifiedBesselK1TmpSizeV2(const AscNode &node);
57std::vector<std::unique_ptr<TmpBufDesc>> CalcIsInfTmpSize(const AscNode &node);57std::vector<std::unique_ptr<TmpBufDesc>> CalcIsInfTmpSize(const AscNode &node);
58+std::vector<std::unique_ptr<TmpBufDesc>> CalcSoftmaxTmpSizeV2(const AscNode &node);
58std::vector<std::unique_ptr<TmpBufDesc>> CalcMaskedFillTmpSize(const AscNode &node);59std::vector<std::unique_ptr<TmpBufDesc>> CalcMaskedFillTmpSize(const AscNode &node);
59} // namespace ascir60} // namespace ascir
60} // namespace af61} // namespace af
@@ -0,0 +1,62 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+#include "default_reg_func_v2.h"
11+#include "graph/symbolizer/symbolic_utils.h"
12+ 
13+namespace af {
14+namespace ascir {
15+ 
16+constexpr uint32_t kAlignWidth = 32U;
17+constexpr int32_t kSoftmaxExtraMul = 8;
18+ 
19+static AscGraphAttr *GetOrCreateSoftmaxGraphAttrsGroup(const ComputeGraphPtr &graph) {
20+ GE_CHECK_NOTNULL_EXEC(graph, return nullptr;);
21+ auto attr = graph->GetOrCreateAttrsGroup<AscGraphAttr>();
22+ GE_CHECK_NOTNULL_EXEC(attr, return nullptr;);
23+ return attr;
24+}
25+ 
26+std::vector<std::unique_ptr<TmpBufDesc>> CalcSoftmaxTmpSizeV2(const AscNode &node) {
27+ std::vector<std::unique_ptr<TmpBufDesc>> tmp_buf_desc;
28+ AscNodeInputs node_inputs = node.inputs;
29+ AscNodeOutputs node_outputs = node.outputs;
30+ if (node_outputs[0].attr.vectorized_strides.empty()) {
31+ return tmp_buf_desc;
32+ }
33+ 
34+ const auto dtype_size = ge::GetSizeByDataType(node_inputs[0].attr.dtype);
35+ GE_ASSERT_TRUE(dtype_size > 0, "Softmax node %s[%s] invalid dtype size.", node.GetTypePtr(), node.GetNamePtr());
36+ const uint32_t align_size = kAlignWidth / static_cast<uint32_t>(dtype_size);
37+ 
38+ auto attr = GetOrCreateSoftmaxGraphAttrsGroup(node.GetOwnerComputeGraph());
39+ 
40+ Expression a_exp = Symbol(1);
41+ Expression r_exp = Symbol(1);
42+ const size_t num_axes = node_outputs[0].attr.vectorized_strides.size();
43+ for (size_t i = 0; i < num_axes; i++) {
44+ uint64_t vectorized_axis_id = node_outputs[0].attr.vectorized_axis[i];
45+ Expression axis_size = attr->axis[vectorized_axis_id]->size;
46+ if (i == num_axes - 1) {
47+ r_exp = sym::Align(axis_size, align_size);
48+ } else {
49+ a_exp = sym::Mul(a_exp, axis_size);
50+ }
51+ }
52+ 
53+ Expression element_size = sym::Add(sym::Mul(a_exp, r_exp), sym::Mul(a_exp, Symbol(kSoftmaxExtraMul)));
54+ Expression tmp_size = sym::Mul(element_size, Symbol(dtype_size));
55+ GELOGD("Softmax node %s[%s] temp buffer size: %s", node.GetTypePtr(), node.GetNamePtr(), tmp_size.Str().get());
56+ TmpBufDesc desc = {tmp_size, -1};
57+ tmp_buf_desc.emplace_back(std::make_unique<TmpBufDesc>(desc));
58+ return tmp_buf_desc;
59+}
60+ 
61+} // namespace ascir
62+} // namespace af
@@ -128,6 +128,9 @@ Register::Register() {
128 };128 };
129 const std::string kAscendcRemainderRegBaseStr = {129 const std::string kAscendcRemainderRegBaseStr = {
130#include "remainder_reg_base.h"130#include "remainder_reg_base.h"
131+ };
132+ const std::string kAscendcSoftmaxAfRegBaseStr = {
133+#include "softmax_af_reg_base.h"
131 };134 };
132 const std::string kAscendcNegRegBaseStr = {135 const std::string kAscendcNegRegBaseStr = {
133#include "neg_reg_base.h"136#include "neg_reg_base.h"
@@ -272,6 +275,7 @@ Register::Register() {
272 {"fmod_reg_base.h", kAscendcFmodRegBaseStr},275 {"fmod_reg_base.h", kAscendcFmodRegBaseStr},
273 {"trunc_div_reg_base.h", kAscendcTruncDivRegBaseStr},276 {"trunc_div_reg_base.h", kAscendcTruncDivRegBaseStr},
274 {"remainder_reg_base.h", kAscendcRemainderRegBaseStr},277 {"remainder_reg_base.h", kAscendcRemainderRegBaseStr},
278+ {"softmax_af_reg_base.h", kAscendcSoftmaxAfRegBaseStr},
275 {"neg_reg_base.h", kAscendcNegRegBaseStr},279 {"neg_reg_base.h", kAscendcNegRegBaseStr},
276 {"square_reg_base.h", kAscendcSquareRegBaseStr},280 {"square_reg_base.h", kAscendcSquareRegBaseStr},
277 {"transpose_reg_base.h", kAscendcTransposeRegBaseStr},281 {"transpose_reg_base.h", kAscendcTransposeRegBaseStr},
@@ -0,0 +1,73 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+#include "softmax_api_call.h"
11+ 
12+#include <sstream>
13+#include "attr_utils.h"
14+#include "ascir_ops.h"
15+#include "common_utils.h"
16+#include "common/ge_common/debug/log.h"
17+#include "common/checker.h"
18+#include "graph/ascendc_ir/utils/asc_tensor_utils.h"
19+#include "api_call/utils/api_call_factory.h"
20+#include "codegen/expression_convert_struct.h"
21+ 
22+namespace codegen {
23+using namespace af::ops;
24+using namespace af::ascir_op;
25+using namespace ascgen_utils;
26+ 
27+Status SoftmaxApiCall::Generate(const TPipe &tpipe, const std::vector<ascir::AxisId> &current_axis,
28+ const std::vector<std::reference_wrapper<const Tensor>> &inputs,
29+ const std::vector<std::reference_wrapper<const Tensor>> &outputs,
30+ std::string &result) const {
31+ auto x = inputs[0].get();
32+ auto y = outputs[0].get();
33+ (void)RegisterBasicDumpParam(this->api_name_, inputs, outputs, CombinedExprFactory::SymbolVar(x.actual_size.Str()));
34+ std::string dtype_name;
35+ GE_CHK_STATUS_RET(Tensor::DtypeName(x.dtype, dtype_name), "Codegen(softmax) get data type:%d failed",
36+ static_cast<int32_t>(x.dtype));
37+ 
38+ int64_t life_time_axis_id = -1L;
39+ int64_t id = -1L;
40+ auto it = this->tmp_buf_id.find(life_time_axis_id);
41+ GE_ASSERT_TRUE(it != this->tmp_buf_id.end(), "SoftmaxApiCall cannot find tmp buffer id to use.");
42+ id = it->second;
43+ std::string tmp_buf_name = tpipe.tmp_buf.name + "_" + std::to_string(id);
44+ 
45+ std::stringstream a_actual;
46+ std::stringstream r_actual;
47+ a_actual << "uint32_t a_actual = 1";
48+ r_actual << "uint32_t r_actual = 1";
49+ const size_t num_axes = x.vectorized_axis.size();
50+ for (size_t i = 0; i < num_axes; ++i) {
51+ const auto axis = tpipe.tiler.GetAxis(x.vectorized_axis[i]);
52+ if (i == num_axes - 1) {
53+ r_actual << " * " << axis.actual_size;
54+ } else {
55+ a_actual << " * " << axis.actual_size;
56+ }
57+ }
58+ 
59+ std::stringstream ss;
60+ ss << "{" << std::endl;
61+ ss << a_actual.str() << ";" << std::endl;
62+ ss << r_actual.str() << ";" << std::endl;
63+ ss << this->api_name_ << "<" << dtype_name << ">(" << y << "[" << tpipe.tiler.TensorVectorizedOffset(current_axis, y)
64+ << "], " << x << "[" << tpipe.tiler.TensorVectorizedOffset(current_axis, x) << "], " << tmp_buf_name << ", "
65+ << "{a_actual, r_actual});" << std::endl;
66+ ss << "}" << std::endl;
67+ result = ss.str();
68+ return af::SUCCESS;
69+}
70+ 
71+static ApiCallRegister<SoftmaxApiCall> register_softmax_api_call("SoftmaxApiCall");
72+ 
73+} // namespace codegen
@@ -0,0 +1,25 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+#ifndef __AUTOFUSE_SOFTMAX_API_CALL_H__
11+#define __AUTOFUSE_SOFTMAX_API_CALL_H__
12+#include "codegen_kernel.h"
13+ 
14+namespace codegen {
15+class SoftmaxApiCall final : public ApiCall {
16+ public:
17+ using ApiCall::Generate;
18+ explicit SoftmaxApiCall(const std::string &api_name) : ApiCall(api_name) {}
19+ ~SoftmaxApiCall() final = default;
20+ Status Generate(const TPipe &tpipe, const std::vector<ascir::AxisId> &current_axis,
21+ const std::vector<std::reference_wrapper<const Tensor>> &inputs,
22+ const std::vector<std::reference_wrapper<const Tensor>> &outputs, std::string &result) const override;
23+};
24+} // namespace codegen
25+#endif // __AUTOFUSE_SOFTMAX_API_CALL_H__