已合并
feat: add SoftMax AscIR operator for V2 (A5/v35) #1222
WangYanMale创建于 7月7日
feat: add SoftMax AscIR operator for V2 (A5/v35) #1222
已合并
共 18 个文件变更+279-14
| @@ -177,6 +177,7 @@ inline const std::string kDuplicate = "Duplicate"; | |||
| 177 | inline const std::string kGatherMask = "GatherMask"; | 177 | inline const std::string kGatherMask = "GatherMask"; |
| 178 | inline const std::string kMaxs = "Maxs"; | 178 | inline const std::string kMaxs = "Maxs"; |
| 179 | inline const std::string kMax = "Max"; | 179 | inline const std::string kMax = "Max"; |
| 180 | +inline const std::string kSoftmax = "Softmax"; | ||
| 180 | inline const std::string kArgMax = "ArgMax"; | 181 | inline const std::string kArgMax = "ArgMax"; |
| 181 | inline const std::string kArgMaxMultiRPhase1 = "ArgMaxMultiRPhase1"; | 182 | inline const std::string kArgMaxMultiRPhase1 = "ArgMaxMultiRPhase1"; |
| 182 | inline const std::string kArgMaxMultiRPhase2 = "ArgMaxMultiRPhase2"; | 183 | inline const std::string kArgMaxMultiRPhase2 = "ArgMaxMultiRPhase2"; |
| @@ -237,13 +238,14 @@ inline const std::string kMul = "Mul"; | |||
| 237 | inline const std::string kNeg = "Neg"; | 238 | inline const std::string kNeg = "Neg"; |
| 238 | inline const std::string kReciprocal = "Reciprocal"; | 239 | inline const std::string kReciprocal = "Reciprocal"; |
| 239 | inline const std::string kRelu = "Relu"; | 240 | inline const std::string kRelu = "Relu"; |
| 240 | -inline const std::string kReduceAll = "ReduceAll"; // All | 241 | +inline const std::string kReduceAll = "ReduceAll"; // All |
| 241 | -inline const std::string kReduceAny = "ReduceAny"; // Any | 242 | +inline const std::string kReduceAny = "ReduceAny"; // Any |
| 242 | -inline const std::string kReduceMax = "ReduceMax"; // Max | 243 | +inline const std::string kReduceMax = "ReduceMax"; // Max |
| 243 | -inline const std::string kReduceMean = "ReduceMean"; // Mean | 244 | +inline const std::string kReduceSoftmax = "ReduceSoftmax"; // Softmax |
| 244 | -inline const std::string kReduceMin = "ReduceMin"; // Min | 245 | +inline const std::string kReduceMean = "ReduceMean"; // Mean |
| 245 | -inline const std::string kReduceSum = "ReduceSum"; // Sum | 246 | +inline const std::string kReduceMin = "ReduceMin"; // Min |
| 246 | -inline const std::string kReduceProd = "ReduceProd"; // Prod | 247 | +inline const std::string kReduceSum = "ReduceSum"; // Sum |
| 248 | +inline const std::string kReduceProd = "ReduceProd"; // Prod | ||
| 247 | inline const std::string kAll = "All"; | 249 | inline const std::string kAll = "All"; |
| 248 | inline const std::string kProd = "Prod"; | 250 | inline const std::string kProd = "Prod"; |
| 249 | inline const std::string kAny = "Any"; | 251 | inline 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 | + | ||
| 1343 | def ArgMaxMultiRPhase1( | 1356 | def 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) { | |||
| 533 | bool ScheduleUtils::IsLastAxisReduce(const ascir::ImplGraph &impl_graph) { | 533 | bool 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 | ||
| 84 | bool IsNotPartitionReduce(const af::AscNodePtr &reduce_node, size_t threshold) { | 85 | bool 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 | ||
| 254 | bool ReducePartitionCaseGenerator::ShouldForceAllLoad(ascir::HintGraph &graph) { | 255 | bool 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 | ||
| 19 | namespace optimize { | 19 | namespace 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 | ||
| 25 | class ReducePartitionCaseGenerator : public FusionCaseGenerator { | 25 | class ReducePartitionCaseGenerator : public FusionCaseGenerator { |
| 26 | public: | 26 | public: |
| @@ -36,6 +36,7 @@ set(ascendc_api_regbase_extend_src | |||
| 36 | square.h | 36 | square.h |
| 37 | trunc_div.h | 37 | trunc_div.h |
| 38 | remainder.h | 38 | remainder.h |
| 39 | + softmax_af.h | ||
| 39 | bessel_j_utils.h | 40 | bessel_j_utils.h |
| 40 | bessel_j0.h | 41 | bessel_j0.h |
| 41 | bessel_j1.h | 42 | 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 | + | ||
| 11 | + | ||
| 12 | + | ||
| @@ -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 | + | ||
| 1148 | REG_ASC_IR(Sin) | 1157 | REG_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); | |||
| 185 | REG_ASC_IR_ATT_V2_CLASS_DEFINE(LegendrePolynomialP); | 185 | REG_ASC_IR_ATT_V2_CLASS_DEFINE(LegendrePolynomialP); |
| 186 | REG_ASC_IR_ATT_V2_CLASS_DEFINE(AiryAi); | 186 | REG_ASC_IR_ATT_V2_CLASS_DEFINE(AiryAi); |
| 187 | REG_ASC_IR_ATT_V2_CLASS_DEFINE(Erfinv); | 187 | REG_ASC_IR_ATT_V2_CLASS_DEFINE(Erfinv); |
| 188 | +REG_ASC_IR_ATT_V2_CLASS_DEFINE(Softmax); | ||
| 188 | } // namespace ascir | 189 | } // namespace ascir |
| 189 | } // namespace af | 190 | } // 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 | /*********************************************************************************/ |
| 4185 | class UnsupportedAscIrCodegenImplV2 : public AscIrCodegenV2 { | 4214 | class UnsupportedAscIrCodegenImplV2 : public AscIrCodegenV2 { |
| 4186 | public: | 4215 | public: |
| @@ -55,6 +55,7 @@ std::vector<std::unique_ptr<TmpBufDesc>> CalcModifiedBesselK1TmpSizeV2(const Asc | |||
| 55 | std::vector<std::unique_ptr<TmpBufDesc>> CalcScaledModifiedBesselK0TmpSizeV2(const AscNode &node); | 55 | std::vector<std::unique_ptr<TmpBufDesc>> CalcScaledModifiedBesselK0TmpSizeV2(const AscNode &node); |
| 56 | std::vector<std::unique_ptr<TmpBufDesc>> CalcScaledModifiedBesselK1TmpSizeV2(const AscNode &node); | 56 | std::vector<std::unique_ptr<TmpBufDesc>> CalcScaledModifiedBesselK1TmpSizeV2(const AscNode &node); |
| 57 | std::vector<std::unique_ptr<TmpBufDesc>> CalcIsInfTmpSize(const AscNode &node); | 57 | std::vector<std::unique_ptr<TmpBufDesc>> CalcIsInfTmpSize(const AscNode &node); |
| 58 | +std::vector<std::unique_ptr<TmpBufDesc>> CalcSoftmaxTmpSizeV2(const AscNode &node); | ||
| 58 | std::vector<std::unique_ptr<TmpBufDesc>> CalcMaskedFillTmpSize(const AscNode &node); | 59 | std::vector<std::unique_ptr<TmpBufDesc>> CalcMaskedFillTmpSize(const AscNode &node); |
| 59 | } // namespace ascir | 60 | } // namespace ascir |
| 60 | } // namespace af | 61 | } // 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 | + | ||
| 11 | + | ||
| 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 | 130 | ||
| 131 | + }; | ||
| 132 | + const std::string kAscendcSoftmaxAfRegBaseStr = { | ||
| 133 | + | ||
| 131 | }; | 134 | }; |
| 132 | const std::string kAscendcNegRegBaseStr = { | 135 | const std::string kAscendcNegRegBaseStr = { |
| 133 | 136 | ||
| @@ -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 | + | ||
| 11 | + | ||
| 12 | + | ||
| 13 | + | ||
| 14 | + | ||
| 15 | + | ||
| 16 | + | ||
| 17 | + | ||
| 18 | + | ||
| 19 | + | ||
| 20 | + | ||
| 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> ¤t_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 | + | ||
| 11 | + | ||
| 12 | + | ||
| 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> ¤t_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 | + | ||