已合并
Revert AvgPoolUpdate、NPUClearFloatStatus算子950支持 #10779
吴成文创建于 14 天前
Revert AvgPoolUpdate、NPUClearFloatStatus算子950支持 #10779
已合并
共 40 个文件变更+46-3796
| @@ -1722,6 +1722,14 @@ currently supported. | |||
| 1722 | .OUTPUT(y, TensorType({DT_FLOAT, DT_DOUBLE, DT_FLOAT16, DT_BFLOAT16})) | 1722 | .OUTPUT(y, TensorType({DT_FLOAT, DT_DOUBLE, DT_FLOAT16, DT_BFLOAT16})) |
| 1723 | .OP_END_FACTORY_REG(SparseSegmentMeanGrad) | 1723 | .OP_END_FACTORY_REG(SparseSegmentMeanGrad) |
| 1724 | 1724 | ||
| 1725 | + | ||
| 1726 | + | ||
| 1727 | + REG_OP(NPUClearFloatStatus) | ||
| 1728 | + .INPUT(addr, TensorType({DT_FLOAT})) | ||
| 1729 | + .OUTPUT(data, TensorType({DT_FLOAT})) | ||
| 1730 | + .OP_END_FACTORY_REG(NPUClearFloatStatus) | ||
| 1731 | + | ||
| 1732 | + | ||
| 1725 | /** | 1733 | /** |
| 1726 | * @brief Gather slices from "params" according to "indices"."indices" must be | 1734 | * @brief Gather slices from "params" according to "indices"."indices" must be |
| 1727 | an integer tensor of any dimension(usually 0-D or 1-D). | 1735 | an integer tensor of any dimension(usually 0-D or 1-D). |
| @@ -2793,6 +2801,44 @@ currently supported. | |||
| 2793 | .REQUIRED_ATTR(out_idx, Type) | 2801 | .REQUIRED_ATTR(out_idx, Type) |
| 2794 | .OP_END_FACTORY_REG(UniqueWithCounts) | 2802 | .OP_END_FACTORY_REG(UniqueWithCounts) |
| 2795 | 2803 | ||
| 2804 | + | ||
| 2805 | + | ||
| 2806 | + /** | ||
| 2807 | + *@brief Average pooling update operator. | ||
| 2808 | + *@par Inputs: | ||
| 2809 | + *Two inputs, including: | ||
| 2810 | + * @li x1: A Tensor. Must be one of the following types: float16, float32. | ||
| 2811 | + * @li x2: A Tensor. Must be one of the following types: int4, int8, float16, float32. \n | ||
| 2812 | + | ||
| 2813 | + *@par Attributes: | ||
| 2814 | + * @li ksize: A required ListInt. The size of the sliding window for each dimension of the input tensor. | ||
| 2815 | + * @li strides: A required ListInt. The stride of the sliding window for each dimension of the input tensor. | ||
| 2816 | + * @li padding_mode: An optional String. Padding mode, defaults to "CALCULATED". | ||
| 2817 | + * @li pads: An optional ListInt. Padding sizes, defaults to {0, 0, 0, 0}. | ||
| 2818 | + * @li data_format: An optional String. Data format, defaults to "NHWC". | ||
| 2819 | + * @li ceil_mode: An optional Bool. Whether to use ceil mode, defaults to false. | ||
| 2820 | + * @li exclusive: An optional Bool. Whether to use exclusive mode, defaults to true. \n | ||
| 2821 | + | ||
| 2822 | + *@par Outputs: | ||
| 2823 | + *y: A Tensor. Must be one of the following types: float16, float32. | ||
| 2824 | + *Has the same type and shape as input x1. | ||
| 2825 | + *@par Third-party framework compatibility | ||
| 2826 | + *Compatible with the TensorFlow operator AvgPool. | ||
| 2827 | + */ | ||
| 2828 | + REG_OP(AvgPoolUpdate) | ||
| 2829 | + .INPUT(x1, TensorType({DT_FLOAT16, DT_FLOAT})) | ||
| 2830 | + .INPUT(x2, TensorType({DT_INT4, DT_INT8, DT_FLOAT16, DT_FLOAT})) | ||
| 2831 | + .OUTPUT(y, TensorType({DT_FLOAT16, DT_FLOAT})) | ||
| 2832 | + .REQUIRED_ATTR(ksize, ListInt) | ||
| 2833 | + .REQUIRED_ATTR(strides, ListInt) | ||
| 2834 | + .ATTR(padding_mode, String, "CALCULATED") | ||
| 2835 | + .ATTR(pads, ListInt, {0, 0, 0, 0}) | ||
| 2836 | + .ATTR(data_format, String, "NHWC") | ||
| 2837 | + .ATTR(ceil_mode, Bool, false) | ||
| 2838 | + .ATTR(exclusive, Bool, true) | ||
| 2839 | + .OP_END_FACTORY_REG(AvgPoolUpdate) | ||
| 2840 | + | ||
| 2841 | + | ||
| 2796 | /** | 2842 | /** |
| 2797 | *@brief Finds unique elements in a 1D tensor. \n | 2843 | *@brief Finds unique elements in a 1D tensor. \n |
| 2798 | 2844 | ||
| @@ -1,17 +0,0 @@ | |||
| 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 FILE 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 | -# Generated By CANNBot | ||
| 12 | - | ||
| 13 | - | ||
| 14 | -set(SUPPORT_COMPUTE_UNIT "ascend950") | ||
| 15 | -set(SUPPORT_TILING_DIR "arch35") | ||
| 16 | - | ||
| 17 | -add_modules_sources(HOSTNAME ${OPHOST_NAME} MODE PRIVATE DIR ${CMAKE_CURRENT_SOURCE_DIR} OPTYPE npu_clear_float_status ACLNNTYPE aclnn_exclude COMPUTE_UNIT ${SUPPORT_COMPUTE_UNIT} TILING_DIR ${SUPPORT_TILING_DIR} DISABLE_IN_OPP TRUE) | ||
| @@ -1,73 +0,0 @@ | |||
| 1 | -# NpuClearFloatStatus | ||
| 2 | - | ||
| 3 | -## 产品支持情况 | ||
| 4 | - | ||
| 5 | -| 产品 | 是否支持 | | ||
| 6 | -| :----------------------------------------------------------- | :------: | | ||
| 7 | -| <term>Ascend 950PR/Ascend 950DT</term> | √ | | ||
| 8 | -| <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term> | × | | ||
| 9 | -| <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term> | × | | ||
| 10 | -| <term>Atlas 200I/500 A2 推理产品</term> | √ | | ||
| 11 | -| <term>Atlas 推理系列产品</term> | √ | | ||
| 12 | -| <term>Atlas 训练系列产品</term> | √ | | ||
| 13 | - | ||
| 14 | -## 功能说明 | ||
| 15 | - | ||
| 16 | -- 算子功能:清除NPU每个AI core的浮点溢出状态寄存器,输出固定为8个float32零值。 | ||
| 17 | - | ||
| 18 | -- 计算公式: | ||
| 19 | - | ||
| 20 | -$$ | ||
| 21 | -data = zeros(8, dtype=float32) | ||
| 22 | -$$ | ||
| 23 | - | ||
| 24 | -## 参数说明 | ||
| 25 | - | ||
| 26 | -<table><thead> | ||
| 27 | - <tr> | ||
| 28 | - <th>参数名</th> | ||
| 29 | - <th>输入/输出/属性</th> | ||
| 30 | - <th>描述</th> | ||
| 31 | - <th>数据类型</th> | ||
| 32 | - <th>数据格式</th> | ||
| 33 | - </tr></thead> | ||
| 34 | -<tbody> | ||
| 35 | - <tr> | ||
| 36 | - <td>addr</td> | ||
| 37 | - <td>输入</td> | ||
| 38 | - <td>地址占位符,shape为(8,),数据内容不参与计算。</td> | ||
| 39 | - <td>FLOAT</td> | ||
| 40 | - <td>ND</td> | ||
| 41 | - </tr> | ||
| 42 | - <tr> | ||
| 43 | - <td>data</td> | ||
| 44 | - <td>输出</td> | ||
| 45 | - <td>固定输出8个float32零值,shape为(8,)。</td> | ||
| 46 | - <td>FLOAT</td> | ||
| 47 | - <td>ND</td> | ||
| 48 | - </tr> | ||
| 49 | -</tbody> | ||
| 50 | -</table> | ||
| 51 | - | ||
| 52 | -## 约束说明 | ||
| 53 | - | ||
| 54 | -- addr数据类型必须为float32。 | ||
| 55 | -- 输出data固定为8个float32零值,与输入数据内容无关。 | ||
| 56 | -- addr仅作为算子输入接口占位符,其数据内容不参与计算。 | ||
| 57 | - | ||
| 58 | -## 调用说明 | ||
| 59 | - | ||
| 60 | -<table><thead> | ||
| 61 | - <tr> | ||
| 62 | - <th>调用方式</th> | ||
| 63 | - <th>调用样例</th> | ||
| 64 | - <th>说明</th> | ||
| 65 | - </tr></thead> | ||
| 66 | -<tbody> | ||
| 67 | - <tr> | ||
| 68 | - <td>图模式调用</td> | ||
| 69 | - <td><a href="./examples/test_geir_npu_clear_float_status.cpp">test_geir_npu_clear_float_status</a></td> | ||
| 70 | - <td>参见<a href="../../docs/zh/invocation/quick_op_invocation.md">算子调用</a>完成算子编译和验证。</td> | ||
| 71 | - </tr> | ||
| 72 | -</tbody> | ||
| 73 | -</table> | ||
| @@ -1,375 +0,0 @@ | |||
| 1 | - | ||
| 2 | -/** | ||
| 3 | - * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 4 | - * This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 5 | - * CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 6 | - * Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 7 | - * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 8 | - * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 9 | - * See LICENSE in the root of the software repository for the full text of the License. | ||
| 10 | - */ | ||
| 11 | - | ||
| 12 | -/* Generated By CANNBot */ | ||
| 13 | - | ||
| 14 | - | ||
| 15 | - | ||
| 16 | - | ||
| 17 | - | ||
| 18 | - | ||
| 19 | - | ||
| 20 | - | ||
| 21 | - | ||
| 22 | - | ||
| 23 | - | ||
| 24 | - | ||
| 25 | - | ||
| 26 | - | ||
| 27 | - | ||
| 28 | - | ||
| 29 | - | ||
| 30 | - | ||
| 31 | - | ||
| 32 | - | ||
| 33 | - | ||
| 34 | - | ||
| 35 | - | ||
| 36 | -// 使用 constexpr 替代 #define,类型安全且支持调试器查看 | ||
| 37 | -constexpr int32_t kFailed = -1; | ||
| 38 | -constexpr int32_t kSuccess = 0; | ||
| 39 | - | ||
| 40 | -using namespace ge; | ||
| 41 | -using std::map; | ||
| 42 | -using std::string; | ||
| 43 | -using std::vector; | ||
| 44 | - | ||
| 45 | -enum RunMode { RUN_MODE_S = 0, RUN_MODE_D = 1 }; | ||
| 46 | - | ||
| 47 | -struct CaseResult { | ||
| 48 | - std::string case_name; | ||
| 49 | - bool build_ok; | ||
| 50 | - bool run_ok; | ||
| 51 | - bool output_exists; | ||
| 52 | - int output_count; | ||
| 53 | - std::string err_msg; | ||
| 54 | -}; | ||
| 55 | - | ||
| 56 | -// 以下三个宏定义必须原样复用,禁止重写、禁止改签名、禁止删行 | ||
| 57 | -// ADD_INPUT_MODE 内部封装了三个关键逻辑,重写会丢失: | ||
| 58 | -// 1. S/D 双 desc 分离:_desc_graph(D 模式含 -1,给图引擎做 infershape)和 _desc_real(具体值,给 FillTensorWithValue | ||
| 59 | -// 构造 Tensor) | ||
| 60 | -// 2. #inputIndex 字符串化:生成 "placeholder1" 等合法 op 名(用 + 会变成指针算术) | ||
| 61 | -// 3. update_output_desc_y:Data op 需要同时设置 input 和 output desc | ||
| 62 | -// 多输入依赖算子只需在宏调用处传入各输入的 shape,无需修改宏定义。 | ||
| 63 | - | ||
| 64 | - vector<int64_t> placeholder##inputIndex##_real_shape = inputShape; \ | ||
| 65 | - vector<int64_t> placeholder##inputIndex##_graph_shape = ((mode) == RUN_MODE_D) ? \ | ||
| 66 | - vector<int64_t>( \ | ||
| 67 | - placeholder##inputIndex##_real_shape.size(), -1) : \ | ||
| 68 | - placeholder##inputIndex##_real_shape; \ | ||
| 69 | - auto placeholder##inputIndex = op::Data("placeholder" #inputIndex).set_attr_index(0); \ | ||
| 70 | - TensorDesc placeholder##inputIndex##_desc_graph = TensorDesc(ge::Shape(placeholder##inputIndex##_graph_shape), \ | ||
| 71 | - FORMAT_ND, inputDtype); \ | ||
| 72 | - placeholder##inputIndex##_desc_graph.SetPlacement(ge::kPlacementHost); \ | ||
| 73 | - placeholder##inputIndex##_desc_graph.SetFormat(FORMAT_ND); \ | ||
| 74 | - TensorDesc placeholder##inputIndex##_desc_real = TensorDesc(ge::Shape(placeholder##inputIndex##_real_shape), \ | ||
| 75 | - FORMAT_ND, inputDtype); \ | ||
| 76 | - placeholder##inputIndex##_desc_real.SetPlacement(ge::kPlacementHost); \ | ||
| 77 | - placeholder##inputIndex##_desc_real.SetFormat(FORMAT_ND); \ | ||
| 78 | - placeholder##inputIndex##_desc_real.SetRealDimCnt(placeholder##inputIndex##_real_shape.size()); \ | ||
| 79 | - Tensor tensor_placeholder##inputIndex; \ | ||
| 80 | - ret = FillTensorWithValue(placeholder##inputIndex##_real_shape, tensor_placeholder##inputIndex, \ | ||
| 81 | - placeholder##inputIndex##_desc_real, inputDtype, 2); \ | ||
| 82 | - if (ret != kSuccess) { \ | ||
| 83 | - printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \ | ||
| 84 | - return kFailed; \ | ||
| 85 | - } \ | ||
| 86 | - placeholder##inputIndex.update_input_desc_x(placeholder##inputIndex##_desc_graph); \ | ||
| 87 | - placeholder##inputIndex.update_output_desc_y(placeholder##inputIndex##_desc_graph); \ | ||
| 88 | - input.push_back(tensor_placeholder##inputIndex); \ | ||
| 89 | - graph.AddOp(placeholder##inputIndex); \ | ||
| 90 | - add1.set_input_##inputName(placeholder##inputIndex); \ | ||
| 91 | - inputs.push_back(placeholder##inputIndex) | ||
| 92 | - | ||
| 93 | - | ||
| 94 | - vector<int64_t> placeholder##inputIndex##_shape = inputShape; \ | ||
| 95 | - auto placeholder##inputIndex = op::Const("placeholder" #inputIndex); \ | ||
| 96 | - TensorDesc placeholder##inputIndex##_desc = TensorDesc(ge::Shape(placeholder##inputIndex##_shape), FORMAT_ND, \ | ||
| 97 | - inputDtype); \ | ||
| 98 | - placeholder##inputIndex##_desc.SetPlacement(ge::kPlacementHost); \ | ||
| 99 | - placeholder##inputIndex##_desc.SetFormat(FORMAT_ND); \ | ||
| 100 | - Tensor tensor_placeholder##inputIndex; \ | ||
| 101 | - ret = FillTensorWithValue(placeholder##inputIndex##_shape, tensor_placeholder##inputIndex, \ | ||
| 102 | - placeholder##inputIndex##_desc, inputDtype, 2); \ | ||
| 103 | - if (ret != kSuccess) { \ | ||
| 104 | - printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \ | ||
| 105 | - return kFailed; \ | ||
| 106 | - } \ | ||
| 107 | - placeholder##inputIndex.SetAttr("value", tensor_placeholder##inputIndex); \ | ||
| 108 | - placeholder##inputIndex.update_output_desc_y(placeholder##inputIndex##_desc); \ | ||
| 109 | - graph.AddOp(placeholder##inputIndex); \ | ||
| 110 | - add1.set_input_##inputName(placeholder##inputIndex); \ | ||
| 111 | - add1.update_input_desc_##inputName(placeholder##inputIndex##_desc); \ | ||
| 112 | - inputs.push_back(placeholder##inputIndex) | ||
| 113 | - | ||
| 114 | -// ADD_OUTPUT_MODE 内部封装了 S/D 模式 graph_shape 自动生成(D 模式含 -1),禁止重写。 | ||
| 115 | - | ||
| 116 | - vector<int64_t> output##outputIndex##_graph_shape = ((mode) == RUN_MODE_D) ? \ | ||
| 117 | - vector<int64_t>(outputShape.size(), -1) : \ | ||
| 118 | - outputShape; \ | ||
| 119 | - TensorDesc outputName##outputIndex##_desc = TensorDesc(ge::Shape(output##outputIndex##_graph_shape), FORMAT_ND, \ | ||
| 120 | - outputDtype); \ | ||
| 121 | - add1.update_output_desc_##outputName(outputName##outputIndex##_desc) | ||
| 122 | - | ||
| 123 | -string GetTime() | ||
| 124 | -{ | ||
| 125 | - time_t timep; | ||
| 126 | - time(&timep); | ||
| 127 | - char tmp[64]; | ||
| 128 | - strftime(tmp, sizeof(tmp), "%Y-%m-%d %H:%M:%S,000", localtime(&timep)); | ||
| 129 | - return tmp; | ||
| 130 | -} | ||
| 131 | - | ||
| 132 | -// 原 GenOnesData 重命名为 FillTensorWithValue(原名误导,实际生成 value 填充而非 ones) | ||
| 133 | -// 根据 dtype 选择正确的 vector 类型,避免 int32_t vector 构造 float32 数据的类型不安全 | ||
| 134 | -// Tensor(const TensorDesc&, const uint8_t*, size_t) 深拷贝数据 | ||
| 135 | -// (构造后 Tensor 内部 ptr != 传入 ptr,源数据释放后值仍保留) | ||
| 136 | -// 因此 pData 离开作用域析构后 input_tensor 仍有效,无需调用者持有 pData | ||
| 137 | -int32_t FillTensorWithValue(vector<int64_t> shapes, Tensor& input_tensor, TensorDesc& input_tensor_desc, | ||
| 138 | - DataType data_type, int value) | ||
| 139 | -{ | ||
| 140 | - input_tensor_desc.SetRealDimCnt(shapes.size()); | ||
| 141 | - size_t size = 1; | ||
| 142 | - for (uint32_t i = 0; i < shapes.size(); i++) { | ||
| 143 | - size *= shapes[i]; | ||
| 144 | - } | ||
| 145 | - int dtypeSize = ge::GetSizeByDataType(data_type); | ||
| 146 | - if (dtypeSize <= 0) { | ||
| 147 | - printf("%s - ERROR - unsupported data type: %d\n", GetTime().c_str(), data_type); | ||
| 148 | - return kFailed; | ||
| 149 | - } | ||
| 150 | - uint32_t data_len = size * static_cast<uint32_t>(dtypeSize); | ||
| 151 | - // float32 dtype 用 vector<float> 构造,保证类型安全 | ||
| 152 | - // 其他 dtype 仍用 vector<int32_t>(int 值填充,按字节 reinterpret) | ||
| 153 | - if (data_type == ge::DT_FLOAT) { | ||
| 154 | - auto pData = std::make_shared<std::vector<float>>(size, static_cast<float>(value)); | ||
| 155 | - input_tensor = Tensor(input_tensor_desc, reinterpret_cast<uint8_t*>(pData->data()), data_len); | ||
| 156 | - } else { | ||
| 157 | - auto pData = std::make_shared<std::vector<int32_t>>(size, value); | ||
| 158 | - input_tensor = Tensor(input_tensor_desc, reinterpret_cast<uint8_t*>(pData->data()), data_len); | ||
| 159 | - } | ||
| 160 | - return kSuccess; | ||
| 161 | -} | ||
| 162 | - | ||
| 163 | -int CreateOppInGraph(RunMode mode, DataType inDtype, const std::vector<int64_t>& xShape, std::vector<ge::Tensor>& input, | ||
| 164 | - std::vector<Operator>& inputs, std::vector<Operator>& outputs, Graph& graph) | ||
| 165 | -{ | ||
| 166 | - Status ret = kSuccess; | ||
| 167 | - // 自定义代码:添加单算子定义到图中 | ||
| 168 | - auto add1 = op::NPUClearFloatStatus("add1"); | ||
| 169 | - ADD_INPUT_MODE(1, addr, inDtype, xShape, mode); | ||
| 170 | - | ||
| 171 | - // 输出 data shape 固定为 (8,),不依赖输入 addr 的 shape | ||
| 172 | - std::vector<int64_t> outShape = {8}; | ||
| 173 | - ADD_OUTPUT_MODE(1, data, inDtype, outShape, mode); | ||
| 174 | - | ||
| 175 | - // NPUClearFloatStatus 无属性 | ||
| 176 | - | ||
| 177 | - outputs.push_back(add1); | ||
| 178 | - // 添加完毕 | ||
| 179 | - return kSuccess; | ||
| 180 | -} | ||
| 181 | - | ||
| 182 | -// 从 RunOneCase 拆分 — 构建图并添加到 session | ||
| 183 | -static bool BuildAndAddGraph(ge::Session* session, uint32_t graph_id, RunMode mode, DataType dtype, | ||
| 184 | - const std::vector<int64_t>& shape, std::vector<ge::Tensor>& input, std::string& err_msg) | ||
| 185 | -{ | ||
| 186 | - std::string graph_name = "tc_ge_irrun_test_" + std::to_string(graph_id); | ||
| 187 | - Graph graph(graph_name.c_str()); | ||
| 188 | - std::vector<Operator> inputs{}; | ||
| 189 | - std::vector<Operator> outputs{}; | ||
| 190 | - | ||
| 191 | - Status ret = CreateOppInGraph(mode, dtype, shape, input, inputs, outputs, graph); | ||
| 192 | - if (ret != kSuccess) { | ||
| 193 | - err_msg = "CreateOppInGraph failed"; | ||
| 194 | - return false; | ||
| 195 | - } | ||
| 196 | - if (!inputs.empty() && !outputs.empty()) { | ||
| 197 | - graph.SetInputs(inputs).SetOutputs(outputs); | ||
| 198 | - } | ||
| 199 | - std::map<AscendString, AscendString> graph_options = {}; | ||
| 200 | - ret = session->AddGraph(graph_id, graph, graph_options); | ||
| 201 | - if (ret != kSuccess) { | ||
| 202 | - err_msg = "AddGraph failed, ret=" + std::to_string(ret); | ||
| 203 | - return false; | ||
| 204 | - } | ||
| 205 | - return true; | ||
| 206 | -} | ||
| 207 | - | ||
| 208 | -// 从 RunOneCase 拆分 — 执行图并收集结果 | ||
| 209 | -static CaseResult ExecuteGraphAndCollect(ge::Session* session, uint32_t graph_id, std::vector<ge::Tensor>& input, | ||
| 210 | - const std::string& case_name, CaseResult r) | ||
| 211 | -{ | ||
| 212 | - std::vector<ge::Tensor> output; | ||
| 213 | - Status ret = session->RunGraph(graph_id, input, output); | ||
| 214 | - session->RemoveGraph(graph_id); | ||
| 215 | - if (ret != kSuccess) { | ||
| 216 | - r.err_msg = "RunGraph failed, ret=" + std::to_string(ret); | ||
| 217 | - return r; | ||
| 218 | - } | ||
| 219 | - r.run_ok = true; | ||
| 220 | - r.output_count = output.size(); | ||
| 221 | - r.output_exists = (output.size() > 0); | ||
| 222 | - | ||
| 223 | - for (size_t i = 0; i < output.size(); i++) { | ||
| 224 | - int64_t shape_size = output[i].GetTensorDesc().GetShape().GetShapeSize(); | ||
| 225 | - printf(" [%s] output[%zu] dtype=%d shape_size=%lld\n", case_name.c_str(), i, | ||
| 226 | - output[i].GetTensorDesc().GetDataType(), (long long)shape_size); | ||
| 227 | - } | ||
| 228 | - return r; | ||
| 229 | -} | ||
| 230 | - | ||
| 231 | -CaseResult RunOneCase(ge::Session* session, uint32_t graph_id, RunMode mode, DataType dtype, | ||
| 232 | - const std::vector<int64_t>& shape, const std::string& case_name) | ||
| 233 | -{ | ||
| 234 | - CaseResult r; | ||
| 235 | - r.case_name = case_name; | ||
| 236 | - r.build_ok = false; | ||
| 237 | - r.run_ok = false; | ||
| 238 | - r.output_exists = false; | ||
| 239 | - r.output_count = 0; | ||
| 240 | - r.err_msg = ""; | ||
| 241 | - | ||
| 242 | - // 构建阶段 | ||
| 243 | - std::vector<ge::Tensor> input; | ||
| 244 | - if (!BuildAndAddGraph(session, graph_id, mode, dtype, shape, input, r.err_msg)) { | ||
| 245 | - return r; | ||
| 246 | - } | ||
| 247 | - r.build_ok = true; | ||
| 248 | - // 执行阶段 | ||
| 249 | - return ExecuteGraphAndCollect(session, graph_id, input, case_name, std::move(r)); | ||
| 250 | -} | ||
| 251 | - | ||
| 252 | -void PrintReport(const std::vector<CaseResult>& results) | ||
| 253 | -{ | ||
| 254 | - printf("\n"); | ||
| 255 | - printf("====================================================================================================\n"); | ||
| 256 | - printf("| %-22s | %-8s | %-9s | %-12s | %-7s | %-20s\n", "Case", "Build", "RunGraph", "OutputExists", "OutCnt", | ||
| 257 | - "ErrMsg"); | ||
| 258 | - printf("----------------------------------------------------------------------------------------------------\n"); | ||
| 259 | - int pass_cnt = 0; | ||
| 260 | - int total = results.size(); | ||
| 261 | - for (const auto& r : results) { | ||
| 262 | - bool pass = r.build_ok && r.run_ok && r.output_exists; | ||
| 263 | - if (pass) | ||
| 264 | - pass_cnt++; | ||
| 265 | - printf("| %-22s | %-8s | %-9s | %-12s | %-7d | %-20s\n", r.case_name.c_str(), r.build_ok ? "OK" : "FAIL", | ||
| 266 | - r.run_ok ? "OK" : "FAIL", r.output_exists ? "OK" : "FAIL", r.output_count, | ||
| 267 | - r.err_msg.empty() ? "-" : r.err_msg.c_str()); | ||
| 268 | - } | ||
| 269 | - printf("====================================================================================================\n"); | ||
| 270 | - printf("Summary: %d/%d passed\n", pass_cnt, total); | ||
| 271 | -} | ||
| 272 | - | ||
| 273 | -// 从 main 拆分 — 初始化图引擎 | ||
| 274 | -static int32_t InitGE() | ||
| 275 | -{ | ||
| 276 | - printf("%s - INFO - [XIR]: Start to initialize ge using ge global options\n", GetTime().c_str()); | ||
| 277 | - std::map<AscendString, AscendString> global_options = { | ||
| 278 | - {"ge.exec.deviceId", "0"}, {"ge.graphRunMode", "0"}, {"ge.exec.precision_mode", "must_keep_origin_dtype"}}; | ||
| 279 | - Status ret = ge::GEInitialize(global_options); | ||
| 280 | - if (ret != kSuccess) { | ||
| 281 | - printf("%s - INFO - [XIR]: Initialize ge using ge global options failed\n", GetTime().c_str()); | ||
| 282 | - return kFailed; | ||
| 283 | - } | ||
| 284 | - printf("%s - INFO - [XIR]: Initialize ge using ge global options success\n", GetTime().c_str()); | ||
| 285 | - return kSuccess; | ||
| 286 | -} | ||
| 287 | - | ||
| 288 | -// 从 main 拆分 — 创建 session | ||
| 289 | -static ge::Session* CreateSession() | ||
| 290 | -{ | ||
| 291 | - std::map<AscendString, AscendString> build_options = {}; | ||
| 292 | - ge::Session* session = new Session(build_options); | ||
| 293 | - if (session == nullptr) { | ||
| 294 | - printf("%s - ERROR - [XIR]: create session failed\n", GetTime().c_str()); | ||
| 295 | - ge::GEFinalize(); | ||
| 296 | - return nullptr; | ||
| 297 | - } | ||
| 298 | - return session; | ||
| 299 | -} | ||
| 300 | - | ||
| 301 | -// 从 main 拆分 — 运行所有测试用例 | ||
| 302 | -static void RunAllCases(ge::Session* session, std::vector<CaseResult>& results) | ||
| 303 | -{ | ||
| 304 | - // dtype 矩阵(取自 reg_op.dtype_set: DT_FLOAT) | ||
| 305 | - struct DtypeEntry { | ||
| 306 | - DataType dt; | ||
| 307 | - std::string name; | ||
| 308 | - }; | ||
| 309 | - std::vector<DtypeEntry> dtype_list = {{DT_FLOAT, "FP32"}}; | ||
| 310 | - | ||
| 311 | - // shape 场景矩阵:输入 addr 和输出 data 均为固定 shape (8,) | ||
| 312 | - // addr 是硬件状态寄存器地址占位符,标准5类场景对本算子均不适用,仅测试固定 (8,) 场景 | ||
| 313 | - struct ShapeEntry { | ||
| 314 | - std::vector<int64_t> shape; | ||
| 315 | - std::string name; | ||
| 316 | - }; | ||
| 317 | - std::vector<ShapeEntry> shape_list = {{{8}, "fixed_8"}}; | ||
| 318 | - | ||
| 319 | - uint32_t graph_id = 0; | ||
| 320 | - for (const auto& d : dtype_list) { | ||
| 321 | - for (const auto& s : shape_list) { | ||
| 322 | - for (auto mode : {RUN_MODE_S, RUN_MODE_D}) { | ||
| 323 | - std::string mode_name = (mode == RUN_MODE_S) ? "S" : "D"; | ||
| 324 | - std::string case_name = d.name + "_" + s.name + "_" + mode_name; | ||
| 325 | - printf("\n%s - INFO - [XIR]: ===== %s =====\n", GetTime().c_str(), case_name.c_str()); | ||
| 326 | - CaseResult r = RunOneCase(session, graph_id, mode, d.dt, s.shape, case_name); | ||
| 327 | - results.push_back(r); | ||
| 328 | - graph_id++; | ||
| 329 | - } | ||
| 330 | - } | ||
| 331 | - } | ||
| 332 | -} | ||
| 333 | - | ||
| 334 | -// 从 main 拆分 — 清理资源并返回最终状态 | ||
| 335 | -static int32_t Cleanup(ge::Session* session, bool all_pass) | ||
| 336 | -{ | ||
| 337 | - delete session; | ||
| 338 | - printf("%s - INFO - [XIR]: Start to finalize ir graph session\n", GetTime().c_str()); | ||
| 339 | - Status ret = ge::GEFinalize(); | ||
| 340 | - if (ret != kSuccess) { | ||
| 341 | - printf("%s - INFO - [XIR]: Finalize ir graph session failed\n", GetTime().c_str()); | ||
| 342 | - return kFailed; | ||
| 343 | - } | ||
| 344 | - printf("%s - INFO - [XIR]: Finalize ir graph session success\n", GetTime().c_str()); | ||
| 345 | - return all_pass ? kSuccess : kFailed; | ||
| 346 | -} | ||
| 347 | - | ||
| 348 | -int main(int argc, char* argv[]) | ||
| 349 | -{ | ||
| 350 | - // main 拆分为 InitGE / CreateSession / RunAllCases / Cleanup | ||
| 351 | - if (InitGE() != kSuccess) { | ||
| 352 | - return kFailed; | ||
| 353 | - } | ||
| 354 | - ge::Session* session = CreateSession(); | ||
| 355 | - if (session == nullptr) { | ||
| 356 | - return kFailed; | ||
| 357 | - } | ||
| 358 | - | ||
| 359 | - std::vector<CaseResult> results; | ||
| 360 | - RunAllCases(session, results); | ||
| 361 | - PrintReport(results); | ||
| 362 | - | ||
| 363 | - bool all_pass = true; | ||
| 364 | - for (const auto& r : results) { | ||
| 365 | - if (!r.build_ok || !r.run_ok || !r.output_exists) { | ||
| 366 | - all_pass = false; | ||
| 367 | - } | ||
| 368 | - } | ||
| 369 | - if (all_pass) { | ||
| 370 | - printf("\n%s - INFO - [XIR]: ALL CASES PASSED\n", GetTime().c_str()); | ||
| 371 | - } else { | ||
| 372 | - printf("\n%s - ERROR - [XIR]: SOME CASES FAILED, see report above\n", GetTime().c_str()); | ||
| 373 | - } | ||
| 374 | - return Cleanup(session, all_pass); | ||
| 375 | -} | ||
| @@ -1,9 +0,0 @@ | |||
| 1 | -# Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 2 | -# This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 3 | -# CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 4 | -# Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 5 | -# THIS FILE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 6 | -# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 7 | -# See LICENSE in the root of the software repository for the full text of the License. | ||
| 8 | - | ||
| 9 | -add_graph_plugin_sources() | ||
| @@ -1,39 +0,0 @@ | |||
| 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 | -/* Generated By CANNBot */ | ||
| 11 | - | ||
| 12 | -/*! | ||
| 13 | - * \file npu_clear_float_status_graph_infer.cpp | ||
| 14 | - * \brief npu_clear_float_status operator graph infer resource | ||
| 15 | - */ | ||
| 16 | - | ||
| 17 | - | ||
| 18 | - | ||
| 19 | - | ||
| 20 | - | ||
| 21 | -namespace ops { | ||
| 22 | -using namespace ge; | ||
| 23 | - | ||
| 24 | -static ge::graphStatus InferDataTypeNPUClearFloatStatus(gert::InferDataTypeContext* context) | ||
| 25 | -{ | ||
| 26 | - OP_LOGD(context->GetNodeName(), "Begin to do InferDataTypeNPUClearFloatStatus"); | ||
| 27 | - | ||
| 28 | - // data dtype 固定为 float32,与输入 addr 的 dtype 无关 | ||
| 29 | - // (addr 是地址占位符,非数据 Tensor) | ||
| 30 | - // 原实现跟随输入 dtype,若输入非 float32 会导致图编译阶段下游 dtype 推导异常 | ||
| 31 | - context->SetOutputDataType(NpuCfs::DATA_IDX, ge::DT_FLOAT); | ||
| 32 | - | ||
| 33 | - OP_LOGD(context->GetNodeName(), "End to do InferDataTypeNPUClearFloatStatus"); | ||
| 34 | - return GRAPH_SUCCESS; | ||
| 35 | -} | ||
| 36 | - | ||
| 37 | -IMPL_OP(NPUClearFloatStatus).InferDataType(InferDataTypeNPUClearFloatStatus); | ||
| 38 | - | ||
| 39 | -}; // namespace ops | ||
| @@ -1,43 +0,0 @@ | |||
| 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 | -/* Generated By CANNBot */ | ||
| 11 | - | ||
| 12 | -/*! | ||
| 13 | - * \file npu_clear_float_status_proto.h | ||
| 14 | - * \brief | ||
| 15 | - */ | ||
| 16 | - | ||
| 17 | - | ||
| 18 | - | ||
| 19 | - | ||
| 20 | - | ||
| 21 | - | ||
| 22 | -namespace ge { | ||
| 23 | - | ||
| 24 | -/** | ||
| 25 | - *@brief Clears the float overflow status on NPU. | ||
| 26 | - *@par Inputs: | ||
| 27 | - *One input, including: | ||
| 28 | - * @li addr: A ND Tensor of type float32. Shape (8,). \n | ||
| 29 | - | ||
| 30 | - *@par Outputs: | ||
| 31 | - *data: A ND Tensor of type float32. Shape (8,). | ||
| 32 | - */ | ||
| 33 | - | ||
| 34 | - | ||
| 35 | -REG_OP(NPUClearFloatStatus) | ||
| 36 | - .INPUT(addr, TensorType({DT_FLOAT})) | ||
| 37 | - .OUTPUT(data, TensorType({DT_FLOAT})) | ||
| 38 | - .OP_END_FACTORY_REG(NPUClearFloatStatus) | ||
| 39 | - | ||
| 40 | - | ||
| 41 | -} // namespace ge | ||
| 42 | - | ||
| 43 | - | ||
| @@ -1,184 +0,0 @@ | |||
| 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 | -/* Generated By CANNBot */ | ||
| 12 | - | ||
| 13 | - | ||
| 14 | - | ||
| 15 | - | ||
| 16 | - | ||
| 17 | - | ||
| 18 | - | ||
| 19 | - | ||
| 20 | - | ||
| 21 | - | ||
| 22 | -namespace optiling { | ||
| 23 | - | ||
| 24 | -// UB 分配策略 — TBuf float16[38400] = 75KB | ||
| 25 | -// DCACHE_SIZE = 128KB (默认), STATIC_UB_ESTIMATE = 0 (无 __ubuf__ 静态数组) | ||
| 26 | -// 动态 UB 池 = ubSize - 128KB, 需 >= 75KB (TBuf 分配) | ||
| 27 | -constexpr uint32_t DCACHE_SIZE = 128 * 1024; | ||
| 28 | -constexpr uint32_t STATIC_UB_ESTIMATE = 0; | ||
| 29 | -// TBuf 所需容量 = VEC_DUP_SIZE * sizeof(half) = 38400 * 2 = 75KB | ||
| 30 | -// VEC_DUP_SIZE 定义于 op_kernel/arch35/npu_clear_float_status_simt.h | ||
| 31 | -// 无法直接引用 simt.h 的 VEC_DUP_SIZE(kernel 头文件不应引入 host 代码), | ||
| 32 | -// 若 simt.h 的 VEC_DUP_SIZE 变化,此处需同步修改,static_assert 会校验一致性 | ||
| 33 | -constexpr uint32_t TBUF_REQUIRED_SIZE = 38400 * 2; | ||
| 34 | -static_assert(TBUF_REQUIRED_SIZE == 38400 * 2, "TBUF_REQUIRED_SIZE must match VEC_DUP_SIZE * sizeof(half)"); | ||
| 35 | - | ||
| 36 | -struct NPUClearFloatStatusCompileInfo {}; | ||
| 37 | - | ||
| 38 | -static ge::graphStatus GetPlatformInfo(gert::TilingContext* context, uint64_t& ubSize, int64_t& coreNum) | ||
| 39 | -{ | ||
| 40 | - fe::PlatFormInfos* platformInfoPtr = context->GetPlatformInfo(); | ||
| 41 | - OP_CHECK_NULL_WITH_CONTEXT(context, platformInfoPtr); | ||
| 42 | - auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfoPtr); | ||
| 43 | - coreNum = ascendcPlatform.GetCoreNumAiv(); | ||
| 44 | - OP_CHECK_IF(coreNum == 0, OP_LOGE(context, "coreNum is 0"), return ge::GRAPH_FAILED); | ||
| 45 | - ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSize); | ||
| 46 | - OP_CHECK_IF(ubSize == 0, OP_LOGE(context, "ubSize is 0"), return ge::GRAPH_FAILED); | ||
| 47 | - return ge::GRAPH_SUCCESS; | ||
| 48 | -} | ||
| 49 | - | ||
| 50 | -static ge::graphStatus ValidateDtype(const gert::TilingContext* context) | ||
| 51 | -{ | ||
| 52 | - auto addrDesc = context->GetInputDesc(NpuCfs::ADDR_IDX); | ||
| 53 | - OP_CHECK_NULL_WITH_CONTEXT(context, addrDesc); | ||
| 54 | - OP_CHECK_IF(addrDesc->GetDataType() != ge::DT_FLOAT, | ||
| 55 | - OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(context->GetNodeName(), "addr", | ||
| 56 | - ge::TypeUtils::DataTypeToSerialString(addrDesc->GetDataType()), | ||
| 57 | - "The dtype of addr must be FLOAT"), | ||
| 58 | - return ge::GRAPH_FAILED); | ||
| 59 | - auto dataDesc = context->GetOutputDesc(NpuCfs::DATA_IDX); | ||
| 60 | - OP_CHECK_NULL_WITH_CONTEXT(context, dataDesc); | ||
| 61 | - OP_CHECK_IF(dataDesc->GetDataType() != ge::DT_FLOAT, | ||
| 62 | - OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(context->GetNodeName(), "data", | ||
| 63 | - ge::TypeUtils::DataTypeToSerialString(dataDesc->GetDataType()), | ||
| 64 | - "The dtype of data must be FLOAT"), | ||
| 65 | - return ge::GRAPH_FAILED); | ||
| 66 | - return ge::GRAPH_SUCCESS; | ||
| 67 | -} | ||
| 68 | - | ||
| 69 | -static ge::graphStatus ValidateInputs(gert::TilingContext* context) | ||
| 70 | -{ | ||
| 71 | - OP_CHECK_IF(ValidateDtype(context) != ge::GRAPH_SUCCESS, OP_LOGE(context, "ValidateDtype failed"), | ||
| 72 | - return ge::GRAPH_FAILED); | ||
| 73 | - return ge::GRAPH_SUCCESS; | ||
| 74 | -} | ||
| 75 | - | ||
| 76 | -static void DumpTilingData(gert::TilingContext* context, const NPUClearFloatStatusTilingData* tiling) | ||
| 77 | -{ | ||
| 78 | - OP_LOGD(context, "NPUClearFloatStatusTilingData: needCoreNum=%d", tiling->needCoreNum); | ||
| 79 | -} | ||
| 80 | - | ||
| 81 | -// 本算子无需用户 workspace,仅需系统 workspace | ||
| 82 | -static ge::graphStatus GetWorkspaceSize(gert::TilingContext* context) | ||
| 83 | -{ | ||
| 84 | - fe::PlatFormInfos* platformInfoPtr = context->GetPlatformInfo(); | ||
| 85 | - OP_CHECK_NULL_WITH_CONTEXT(context, platformInfoPtr); | ||
| 86 | - auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfoPtr); | ||
| 87 | - uint64_t sysWorkspaceSize = ascendcPlatform.GetLibApiWorkSpaceSize(); | ||
| 88 | - size_t* currentWorkspace = context->GetWorkspaceSizes(1); | ||
| 89 | - OP_CHECK_NULL_WITH_CONTEXT(context, currentWorkspace); | ||
| 90 | - currentWorkspace[0] = static_cast<size_t>(sysWorkspaceSize); | ||
| 91 | - return ge::GRAPH_SUCCESS; | ||
| 92 | -} | ||
| 93 | - | ||
| 94 | -// 核数切分 — needCoreNum = coreNum(全核启动) | ||
| 95 | -static ge::graphStatus ComputeTiling(gert::TilingContext* context, NPUClearFloatStatusTilingData* tiling, | ||
| 96 | - int64_t coreNum) | ||
| 97 | -{ | ||
| 98 | - // 防御性检查 coreNum 不超过 INT32_MAX | ||
| 99 | - // 物理核数实际 < 1000,此检查为纯防御性,避免 int64→int32 强转溢出导致未定义行为 | ||
| 100 | - OP_CHECK_IF( | ||
| 101 | - coreNum > INT32_MAX, | ||
| 102 | - OP_LOGE(context, "coreNum %lld exceeds INT32_MAX, cannot cast to int32_t", static_cast<long long>(coreNum)), | ||
| 103 | - return ge::GRAPH_FAILED); | ||
| 104 | - tiling->needCoreNum = static_cast<int32_t>(coreNum); | ||
| 105 | - return ge::GRAPH_SUCCESS; | ||
| 106 | -} | ||
| 107 | - | ||
| 108 | -// 无需 TilingKey 区分场景(单 dtype 单场景,DTYPE_ 宏实例化) | ||
| 109 | -static uint64_t GetTilingKey() | ||
| 110 | -{ | ||
| 111 | - return GET_TPL_TILING_KEY(static_cast<uint64_t>(NPU_CLEAR_FLOAT_STATUS_SCH_MODE_DEFAULT)); | ||
| 112 | -} | ||
| 113 | - | ||
| 114 | -// UB 分配策略 — 校验动态 UB 池容量并设置 LocalMemorySize | ||
| 115 | -static ge::graphStatus ApplyUbConfig(gert::TilingContext* context, uint64_t ubSize) | ||
| 116 | -{ | ||
| 117 | - OP_CHECK_IF( | ||
| 118 | - (ubSize <= DCACHE_SIZE + STATIC_UB_ESTIMATE), | ||
| 119 | - OP_LOGE(context, "ubSize %llu <= DCACHE_SIZE + STATIC_UB_ESTIMATE", static_cast<unsigned long long>(ubSize)), | ||
| 120 | - return ge::GRAPH_FAILED); | ||
| 121 | - OP_CHECK_IF((ubSize - DCACHE_SIZE - STATIC_UB_ESTIMATE < TBUF_REQUIRED_SIZE), | ||
| 122 | - OP_LOGE(context, "dynamic UB pool %llu < TBuf required %u", | ||
| 123 | - static_cast<unsigned long long>(ubSize - DCACHE_SIZE - STATIC_UB_ESTIMATE), TBUF_REQUIRED_SIZE), | ||
| 124 | - return ge::GRAPH_FAILED); | ||
| 125 | - auto res = context->SetLocalMemorySize(static_cast<uint32_t>(ubSize - DCACHE_SIZE - STATIC_UB_ESTIMATE)); | ||
| 126 | - OP_CHECK_IF((res != ge::GRAPH_SUCCESS), | ||
| 127 | - OP_LOGE(context, "SetLocalMemorySize failed, ubSize=%llu", static_cast<unsigned long long>(ubSize)), | ||
| 128 | - return ge::GRAPH_FAILED); | ||
| 129 | - return ge::GRAPH_SUCCESS; | ||
| 130 | -} | ||
| 131 | - | ||
| 132 | -static ge::graphStatus NPUClearFloatStatusTilingFunc(gert::TilingContext* context) | ||
| 133 | -{ | ||
| 134 | - // 1. validate inputs | ||
| 135 | - OP_CHECK_IF(ValidateInputs(context) != ge::GRAPH_SUCCESS, OP_LOGE(context, "ValidateInputs failed"), | ||
| 136 | - return ge::GRAPH_FAILED); | ||
| 137 | - | ||
| 138 | - // 2. get platform info | ||
| 139 | - uint64_t ubSize; | ||
| 140 | - int64_t coreNum; | ||
| 141 | - OP_CHECK_IF(GetPlatformInfo(context, ubSize, coreNum) != ge::GRAPH_SUCCESS, | ||
| 142 | - OP_LOGE(context, "NPUClearFloatStatus: GetPlatformInfo failed, see previous log for details"), | ||
| 143 | - return ge::GRAPH_FAILED); | ||
| 144 | - | ||
| 145 | - // 3. compute tiling (BlockDim = coreNum) | ||
| 146 | - NPUClearFloatStatusTilingData* tiling = context->GetTilingData<NPUClearFloatStatusTilingData>(); | ||
| 147 | - OP_CHECK_NULL_WITH_CONTEXT(context, tiling); | ||
| 148 | - OP_CHECK_IF( | ||
| 149 | - memset_s(tiling, sizeof(NPUClearFloatStatusTilingData), 0, sizeof(NPUClearFloatStatusTilingData)) != EOK, | ||
| 150 | - OP_LOGE(context, "set tiling data error"), return ge::GRAPH_FAILED); | ||
| 151 | - | ||
| 152 | - OP_CHECK_IF(ComputeTiling(context, tiling, coreNum) != ge::GRAPH_SUCCESS, | ||
| 153 | - OP_LOGE(context, "NPUClearFloatStatus: ComputeTiling failed, see previous log for details"), | ||
| 154 | - return ge::GRAPH_FAILED); | ||
| 155 | - | ||
| 156 | - DumpTilingData(context, tiling); | ||
| 157 | - | ||
| 158 | - // 4. get workspace size (仅系统 workspace) | ||
| 159 | - OP_CHECK_IF(GetWorkspaceSize(context) != ge::GRAPH_SUCCESS, | ||
| 160 | - OP_LOGE(context, "NPUClearFloatStatus: GetWorkspaceSize failed, see previous log for details"), | ||
| 161 | - return ge::GRAPH_FAILED); | ||
| 162 | - | ||
| 163 | - // 5. set block dim (BlockDim = coreNum, 所有 AI Core 必须执行) | ||
| 164 | - context->SetBlockDim(tiling->needCoreNum); | ||
| 165 | - | ||
| 166 | - // 6. set local memory size | ||
| 167 | - OP_CHECK_IF(ApplyUbConfig(context, ubSize) != ge::GRAPH_SUCCESS, OP_LOGE(context, "ApplyUbConfig failed"), | ||
| 168 | - return ge::GRAPH_FAILED); | ||
| 169 | - | ||
| 170 | - // 7. set tiling key | ||
| 171 | - context->SetTilingKey(GetTilingKey()); | ||
| 172 | - return ge::GRAPH_SUCCESS; | ||
| 173 | -} | ||
| 174 | - | ||
| 175 | -static ge::graphStatus TilingParseForNPUClearFloatStatus([[maybe_unused]] gert::TilingParseContext* context) | ||
| 176 | -{ | ||
| 177 | - return ge::GRAPH_SUCCESS; | ||
| 178 | -} | ||
| 179 | - | ||
| 180 | -IMPL_OP_OPTILING(NPUClearFloatStatus) | ||
| 181 | - .Tiling(NPUClearFloatStatusTilingFunc) | ||
| 182 | - .TilingParse<NPUClearFloatStatusCompileInfo>(TilingParseForNPUClearFloatStatus); | ||
| 183 | - | ||
| 184 | -} // namespace optiling | ||
| @@ -1,31 +0,0 @@ | |||
| 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 | -/* Generated By CANNBot */ | ||
| 12 | - | ||
| 13 | - | ||
| 14 | - | ||
| 15 | - | ||
| 16 | - | ||
| 17 | - | ||
| 18 | - | ||
| 19 | -namespace NpuCfs { | ||
| 20 | -// 输出固定为 8 个 float32 零值 | ||
| 21 | -constexpr int64_t OUTPUT_DIM = 8; | ||
| 22 | - | ||
| 23 | -// 输出 data 期望维度数 (1D) | ||
| 24 | -constexpr size_t EXPECTED_DIM_NUM = 1; | ||
| 25 | - | ||
| 26 | -// input/output 索引常量 | ||
| 27 | -constexpr int32_t ADDR_IDX = 0; // input addr 索引 | ||
| 28 | -constexpr int32_t DATA_IDX = 0; // output data 索引 | ||
| 29 | -} // namespace NpuCfs | ||
| 30 | - | ||
| 31 | - | ||
| @@ -1,50 +0,0 @@ | |||
| 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 | -/* Generated By CANNBot */ | ||
| 12 | - | ||
| 13 | - | ||
| 14 | - | ||
| 15 | -namespace ops { | ||
| 16 | -class NPUClearFloatStatus : public OpDef { | ||
| 17 | -public: | ||
| 18 | - explicit NPUClearFloatStatus(const char* name) : OpDef(name) | ||
| 19 | - { | ||
| 20 | - // addr (8,) float32 ND — 地址占位符 | ||
| 21 | - this->Input("addr") | ||
| 22 | - .ParamType(REQUIRED) | ||
| 23 | - .DataType({ge::DT_FLOAT}) | ||
| 24 | - .Format({ge::FORMAT_ND}) | ||
| 25 | - .UnknownShapeFormat({ge::FORMAT_ND}) | ||
| 26 | - .AutoContiguous(); | ||
| 27 | - | ||
| 28 | - // data (8,) float32 ND — 固定输出 8 个 float32 零值 | ||
| 29 | - this->Output("data") | ||
| 30 | - .ParamType(REQUIRED) | ||
| 31 | - .DataType({ge::DT_FLOAT}) | ||
| 32 | - .Format({ge::FORMAT_ND}) | ||
| 33 | - .UnknownShapeFormat({ge::FORMAT_ND}) | ||
| 34 | - .AutoContiguous(); | ||
| 35 | - | ||
| 36 | - // 无属性 | ||
| 37 | - | ||
| 38 | - OpAICoreConfig aicoreConfig; | ||
| 39 | - aicoreConfig.DynamicCompileStaticFlag(true) | ||
| 40 | - .DynamicFormatFlag(false) | ||
| 41 | - .DynamicRankSupportFlag(true) | ||
| 42 | - .DynamicShapeSupportFlag(true) | ||
| 43 | - .NeedCheckSupportFlag(false) | ||
| 44 | - .PrecisionReduceFlag(true) | ||
| 45 | - .ExtendCfgInfo("opFile.value", "npu_clear_float_status"); | ||
| 46 | - this->AICore().AddConfig("ascend950", aicoreConfig); | ||
| 47 | - } | ||
| 48 | -}; | ||
| 49 | -OP_ADD(NPUClearFloatStatus); | ||
| 50 | -} // namespace ops | ||
| @@ -1,87 +0,0 @@ | |||
| 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 | -/* Generated By CANNBot */ | ||
| 11 | - | ||
| 12 | -/*! | ||
| 13 | - * \file npu_clear_float_status_infershape.cpp | ||
| 14 | - * \brief Shape/Format inference for npu_clear_float_status | ||
| 15 | - * | ||
| 16 | - * InferDataType 已迁移至 op_graph/npu_clear_float_status_graph_infer.cpp (IMPL_OP 注册)。 | ||
| 17 | - * 本文件仅保留 InferShape + InferFormat (IMPL_OP_INFERSHAPE 注册)。 | ||
| 18 | - * | ||
| 19 | - * InferShape: 输出 shape 固定为 (8,),不依赖输入 addr 的 shape | ||
| 20 | - * InferFormat: 输出 format 固定为 ND | ||
| 21 | - * infershape 侧校验: addr dtype == float32 | ||
| 22 | - */ | ||
| 23 | - | ||
| 24 | - | ||
| 25 | - | ||
| 26 | - | ||
| 27 | - | ||
| 28 | - | ||
| 29 | - | ||
| 30 | -using namespace ge; | ||
| 31 | - | ||
| 32 | -namespace ops { | ||
| 33 | - | ||
| 34 | -// infershape 侧仅校验 input addr,output data 校验由 tiling 阶段负责 | ||
| 35 | -static ge::graphStatus InferShapeValidateDtype(const gert::InferShapeContext* context) | ||
| 36 | -{ | ||
| 37 | - auto addrDesc = context->GetInputDesc(NpuCfs::ADDR_IDX); | ||
| 38 | - OP_CHECK_NULL_WITH_CONTEXT(context, addrDesc); | ||
| 39 | - OP_CHECK_IF(addrDesc->GetDataType() != ge::DT_FLOAT, | ||
| 40 | - OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(context->GetNodeName(), "addr", | ||
| 41 | - ge::TypeUtils::DataTypeToSerialString(addrDesc->GetDataType()), | ||
| 42 | - "The dtype of addr must be FLOAT"), | ||
| 43 | - return ge::GRAPH_FAILED); | ||
| 44 | - return ge::GRAPH_SUCCESS; | ||
| 45 | -} | ||
| 46 | - | ||
| 47 | -static ge::graphStatus InferShapeValidateInputs(gert::InferShapeContext* context) | ||
| 48 | -{ | ||
| 49 | - OP_CHECK_IF(InferShapeValidateDtype(context) != ge::GRAPH_SUCCESS, | ||
| 50 | - OP_LOGE(context, "InferShapeValidateDtype failed"), return ge::GRAPH_FAILED); | ||
| 51 | - return ge::GRAPH_SUCCESS; | ||
| 52 | -} | ||
| 53 | - | ||
| 54 | -// 输出 shape 固定为 (8,) | ||
| 55 | -static ge::graphStatus InferShapeNPUClearFloatStatus(gert::InferShapeContext* context) | ||
| 56 | -{ | ||
| 57 | - OP_LOGD(context->GetNodeName(), "Begin to do InferShapeNPUClearFloatStatus"); | ||
| 58 | - | ||
| 59 | - OP_CHECK_IF(InferShapeValidateInputs(context) != ge::GRAPH_SUCCESS, | ||
| 60 | - OP_LOGE(context, "InferShapeValidateInputs failed"), return ge::GRAPH_FAILED); | ||
| 61 | - | ||
| 62 | - gert::Shape* dataShape = context->GetOutputShape(NpuCfs::DATA_IDX); | ||
| 63 | - OP_CHECK_NULL_WITH_CONTEXT(context, dataShape); | ||
| 64 | - dataShape->SetDimNum(NpuCfs::EXPECTED_DIM_NUM); | ||
| 65 | - dataShape->SetDim(0, NpuCfs::OUTPUT_DIM); | ||
| 66 | - | ||
| 67 | - OP_LOGD(context->GetNodeName(), "End to do InferShapeNPUClearFloatStatus"); | ||
| 68 | - return GRAPH_SUCCESS; | ||
| 69 | -} | ||
| 70 | - | ||
| 71 | -// InferFormat — 输出 format 固定为 ND | ||
| 72 | -static ge::graphStatus InferFormatNPUClearFloatStatus(gert::InferFormatContext* context) | ||
| 73 | -{ | ||
| 74 | - OP_LOGD(context->GetNodeName(), "Begin to do InferFormatNPUClearFloatStatus"); | ||
| 75 | - auto* outputFormat = context->GetOutputFormat(NpuCfs::DATA_IDX); | ||
| 76 | - OP_CHECK_NULL_WITH_CONTEXT(context, outputFormat); | ||
| 77 | - outputFormat->SetOriginFormat(ge::FORMAT_ND); | ||
| 78 | - outputFormat->SetStorageFormat(ge::FORMAT_ND); | ||
| 79 | - OP_LOGD(context->GetNodeName(), "End to do InferFormatNPUClearFloatStatus"); | ||
| 80 | - return GRAPH_SUCCESS; | ||
| 81 | -} | ||
| 82 | - | ||
| 83 | -// InferDataType 已迁移至 op_graph/npu_clear_float_status_graph_infer.cpp | ||
| 84 | -IMPL_OP_INFERSHAPE(NPUClearFloatStatus) | ||
| 85 | - .InferShape(InferShapeNPUClearFloatStatus) | ||
| 86 | - .InferFormat(InferFormatNPUClearFloatStatus); | ||
| 87 | -} // namespace ops | ||
| @@ -1,83 +0,0 @@ | |||
| 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 | -/* Generated By CANNBot */ | ||
| 12 | - | ||
| 13 | - | ||
| 14 | - | ||
| 15 | - | ||
| 16 | - | ||
| 17 | - | ||
| 18 | - | ||
| 19 | - | ||
| 20 | - | ||
| 21 | - | ||
| 22 | - | ||
| 23 | -namespace NsNpuClearFloatStatus { | ||
| 24 | - | ||
| 25 | -using namespace AscendC; | ||
| 26 | - | ||
| 27 | -constexpr uint32_t THREAD_NUM = 128; | ||
| 28 | -constexpr int32_t OUTPUT_SIZE = 8; // 输出元素数(固定 8 个 float32) | ||
| 29 | -constexpr int32_t VEC_DUP_SIZE = 38400; // vector_dup Tensor 元素数(float16[38400]) | ||
| 30 | - | ||
| 31 | -// SIMT VF: 写 totalElements 个零值到输出 GM (Grid-Stride 循环) | ||
| 32 | -// 由于 totalElements=8 远小于 blockDim*gridDim,实际仅 core 0 前 8 线程 | ||
| 33 | -// 执行写入,其余线程在循环条件判断后立即退出(Grid-Stride 自然结果) | ||
| 34 | -template <typename T> | ||
| 35 | -__simt_vf__ __aicore__ __launch_bounds__(THREAD_NUM) inline void OpNPUClearFloatStatusSimt(int32_t totalElements, | ||
| 36 | - __gm__ T* output) | ||
| 37 | -{ | ||
| 38 | - // 防御性检查,避免负数经 static_cast<uint32_t> 转为巨大无符号数导致循环越界 | ||
| 39 | - if (totalElements <= 0) { | ||
| 40 | - return; | ||
| 41 | - } | ||
| 42 | - // Grid-Stride 循环:写 totalElements 个零值到输出 GM | ||
| 43 | - for (uint32_t idx = blockIdx.x * blockDim.x + threadIdx.x; idx < static_cast<uint32_t>(totalElements); | ||
| 44 | - idx += blockDim.x * gridDim.x) { | ||
| 45 | - output[idx] = static_cast<T>(0); | ||
| 46 | - } | ||
| 47 | -} | ||
| 48 | - | ||
| 49 | -// Process: 标量作用域执行 vector_dup + 启动 SIMT VF | ||
| 50 | -template <typename T> | ||
| 51 | -__aicore__ inline void Process(GM_ADDR addr, GM_ADDR data, const NPUClearFloatStatusTilingData* tilingData) | ||
| 52 | -{ | ||
| 53 | - (void)addr; // 输入 addr 仅占位符,不读取数据 | ||
| 54 | - (void)tilingData; // TilingData 仅 needCoreNum,Process 不直接使用 | ||
| 55 | - | ||
| 56 | - // 辅助 vector_dup 操作(SIMD 向量触发) | ||
| 57 | - // 触发向量计算单元 6 次,依次填充值 3~8,每次覆盖整个 Tensor | ||
| 58 | - // V1 TIK vector_dup 参数签名与当前 AscendC Duplicate API 不同, | ||
| 59 | - // 无法精确对应参数含义,仅保留当前实现以与 V1 行为一致 | ||
| 60 | - TPipe pipe; | ||
| 61 | - TBuf<QuePosition::VECCALC> dataUbInputBuf; | ||
| 62 | - if (!pipe.InitBuffer(dataUbInputBuf, VEC_DUP_SIZE * sizeof(half))) { | ||
| 63 | - return; | ||
| 64 | - } | ||
| 65 | - LocalTensor<half> dataUbInput = dataUbInputBuf.Get<half>(); | ||
| 66 | - | ||
| 67 | - // V1 TIK vector_dup(tensor, value, count) 连续调用 6 次,value 依次取 3~8。 | ||
| 68 | - // 填充值本身无实际意义,目的是触发向量计算单元执行以清除溢出状态标志 | ||
| 69 | - // (ascend950 不支持 set_overflow_status 接口)。 | ||
| 70 | - constexpr int32_t DUP_START_VAL = 3; | ||
| 71 | - constexpr int32_t DUP_COUNT = 6; | ||
| 72 | - for (int32_t i = 0; i < DUP_COUNT; ++i) { | ||
| 73 | - Duplicate(dataUbInput, static_cast<half>(DUP_START_VAL + i), VEC_DUP_SIZE); | ||
| 74 | - } | ||
| 75 | - | ||
| 76 | - // SIMT VF 写 8 个 float32 零值到输出 GM | ||
| 77 | - // SIMT VF 写 GM 后框架自动保证 cache 一致性,无需显式 DataCacheCleanAndInvalid | ||
| 78 | - __gm__ T* outputGm = reinterpret_cast<__gm__ T*>(data); | ||
| 79 | - asc_vf_call<OpNPUClearFloatStatusSimt<T>>(dim3(THREAD_NUM), OUTPUT_SIZE, outputGm); | ||
| 80 | -} | ||
| 81 | - | ||
| 82 | -} // namespace NsNpuClearFloatStatus | ||
| 83 | - | ||
| @@ -1,22 +0,0 @@ | |||
| 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 | -/* Generated By CANNBot */ | ||
| 12 | - | ||
| 13 | - | ||
| 14 | - | ||
| 15 | - | ||
| 16 | - | ||
| 17 | - | ||
| 18 | -struct NPUClearFloatStatusTilingData { | ||
| 19 | - int32_t needCoreNum = 0; // 需要启动的核数(= 物理核数) | ||
| 20 | -}; | ||
| 21 | - | ||
| 22 | - | ||
| @@ -1,26 +0,0 @@ | |||
| 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 | -/* Generated By CANNBot */ | ||
| 12 | - | ||
| 13 | - | ||
| 14 | - | ||
| 15 | - | ||
| 16 | - | ||
| 17 | - | ||
| 18 | - | ||
| 19 | - | ||
| 20 | -ASCENDC_TPL_ARGS_DECL(NPUClearFloatStatus, | ||
| 21 | - ASCENDC_TPL_UINT_DECL(schMode, 1, ASCENDC_TPL_UI_LIST, NPU_CLEAR_FLOAT_STATUS_SCH_MODE_DEFAULT)); | ||
| 22 | - | ||
| 23 | -ASCENDC_TPL_SEL(ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(schMode, ASCENDC_TPL_UI_LIST, | ||
| 24 | - NPU_CLEAR_FLOAT_STATUS_SCH_MODE_DEFAULT))); | ||
| 25 | - | ||
| 26 | - | ||
| @@ -1,27 +0,0 @@ | |||
| 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 | -/* Generated By CANNBot */ | ||
| 12 | - | ||
| 13 | - | ||
| 14 | - | ||
| 15 | -enum class NPUClearFloatStatusTilingKey : uint32_t { | ||
| 16 | - TILING_KEY_DEFAULT = 0, | ||
| 17 | -}; | ||
| 18 | - | ||
| 19 | -template <uint32_t schMode> | ||
| 20 | -__global__ __aicore__ void npu_clear_float_status(GM_ADDR addr, GM_ADDR data, GM_ADDR workspace, GM_ADDR tiling) | ||
| 21 | -{ | ||
| 22 | - REGISTER_TILING_DEFAULT(NPUClearFloatStatusTilingData); | ||
| 23 | - GET_TILING_DATA_WITH_STRUCT(NPUClearFloatStatusTilingData, tilingData, tiling); | ||
| 24 | - if constexpr (schMode == static_cast<uint32_t>(NPUClearFloatStatusTilingKey::TILING_KEY_DEFAULT)) { | ||
| 25 | - NsNpuClearFloatStatus::Process<DTYPE_ADDR>(addr, data, &tilingData); | ||
| 26 | - } | ||
| 27 | -} | ||
| @@ -1,16 +0,0 @@ | |||
| 1 | -# Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 2 | -# This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 3 | -# CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 4 | -# Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 5 | -# THIS FILE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 6 | -# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 7 | -# See LICENSE in the root of the software repository for the full text of the License. | ||
| 8 | -#/ | ||
| 9 | - | ||
| 10 | -file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*) | ||
| 11 | -message(STATUS "=== Debug: CURRENT_DIRS =${CURRENT_DIRS} ") | ||
| 12 | -foreach(SUB_DIR ${CURRENT_DIRS}) | ||
| 13 | - if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt") | ||
| 14 | - add_subdirectory(${SUB_DIR}) | ||
| 15 | - endif() | ||
| 16 | -endforeach() | ||
| @@ -1,76 +0,0 @@ | |||
| 1 | -#!/usr/bin/env python3 | ||
| 2 | -# -*- coding: utf-8 -*- | ||
| 3 | -# Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 4 | -# This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 5 | -# CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 6 | -# Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 7 | -# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 8 | -# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 9 | -# See LICENSE in the root of the software repository for the full text of the License. | ||
| 10 | -""" | ||
| 11 | -Golden TestSpec for npu_clear_float_status operator. | ||
| 12 | - | ||
| 13 | -NPUClearFloatStatus is an Ascend NPU hardware status management operator. | ||
| 14 | -It clears the float overflow status register on each AI core. | ||
| 15 | -Output is always 8 float32 zeros, independent of input data. | ||
| 16 | -""" | ||
| 17 | - | ||
| 18 | -import tensorflow as tf | ||
| 19 | - | ||
| 20 | -__spec__ = { | ||
| 21 | - "npu_clear_float_status": "NpuClearFloatStatusTestSpec", | ||
| 22 | -} | ||
| 23 | - | ||
| 24 | - | ||
| 25 | -class NpuClearFloatStatusTestSpec: | ||
| 26 | - """ | ||
| 27 | - TestSpec shared by Kernel and GEIR for npu_clear_float_status. | ||
| 28 | - | ||
| 29 | - Input: | ||
| 30 | - addr: numpy.ndarray (Kernel) / tf.Tensor (third_party) | ||
| 31 | - float32, shape (8,) - address placeholder, data not used | ||
| 32 | - | ||
| 33 | - Output: | ||
| 34 | - data: float32, shape (8,) - always 8 zeros | ||
| 35 | - | ||
| 36 | - Precision: | ||
| 37 | - binary_equal - output is bit-exact zeros (non-computational operator) | ||
| 38 | - """ | ||
| 39 | - | ||
| 40 | - | ||
| 41 | - def golden(addr, **kwargs): | ||
| 42 | - """ | ||
| 43 | - Kernel/GEIR Golden function. | ||
| 44 | - | ||
| 45 | - Parameters: | ||
| 46 | - addr: numpy.ndarray, float32, shape (8,) | ||
| 47 | - Address placeholder tensor. Data content does not affect output. | ||
| 48 | - | ||
| 49 | - Returns: | ||
| 50 | - list: [numpy.ndarray of 8 float32 zeros] | ||
| 51 | - """ | ||
| 52 | - result = tf.zeros([8], dtype=tf.float32) | ||
| 53 | - return [result.numpy()] | ||
| 54 | - | ||
| 55 | - class ThirdPartyImpl: | ||
| 56 | - """ | ||
| 57 | - Provider third-party implementation for GEIR cross-check. | ||
| 58 | - | ||
| 59 | - Receives original dtype tf.Tensor, independently computes output. | ||
| 60 | - Since output is always zeros, the result is deterministic | ||
| 61 | - and bit-exact regardless of input data. | ||
| 62 | - """ | ||
| 63 | - | ||
| 64 | - def __init__(self, **kwargs): | ||
| 65 | - pass | ||
| 66 | - | ||
| 67 | - def __call__(self, addr, **kwargs): | ||
| 68 | - return [tf.zeros([8], dtype=tf.float32)] | ||
| 69 | - | ||
| 70 | - # Maps framework key to its third-party implementation class | ||
| 71 | - third_party = {"tf": ThirdPartyImpl} | ||
| 72 | - | ||
| 73 | - # Bit-exact float32 output: use binary_equal (output is always zeros) | ||
| 74 | - tolerance = { | ||
| 75 | - "float32": {"standard": "binary_equal"}, | ||
| 76 | - } | ||
| @@ -1,16 +0,0 @@ | |||
| 1 | -# Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 2 | -# This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 3 | -# CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 4 | -# Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 5 | -# THIS FILE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 6 | -# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 7 | -# See LICENSE in the root of the software repository for the full text of the License. | ||
| 8 | -#/ | ||
| 9 | - | ||
| 10 | -file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*) | ||
| 11 | -message(STATUS "=== Debug: CURRENT_DIRS =${CURRENT_DIRS} ") | ||
| 12 | -foreach(SUB_DIR ${CURRENT_DIRS}) | ||
| 13 | - if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt") | ||
| 14 | - add_subdirectory(${SUB_DIR}) | ||
| 15 | - endif() | ||
| 16 | -endforeach() | ||
| @@ -1,14 +0,0 @@ | |||
| 1 | -# Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 2 | -# This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 3 | -# CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 4 | -# Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 5 | -# THIS FILE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 6 | -# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 7 | -# See LICENSE in the root of the software repository for the full text of the License. | ||
| 8 | -#/ | ||
| 9 | - | ||
| 10 | -file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*) | ||
| 11 | -if(UT_TEST_ALL OR OP_HOST_UT) | ||
| 12 | - add_modules_ut_sources(HOSTNAME ${OP_TILING_MODULE_NAME} MODE PRIVATE DIR ${CMAKE_CURRENT_SOURCE_DIR}) | ||
| 13 | - add_modules_ut_sources(HOSTNAME ${OP_INFERSHAPE_MODULE_NAME} MODE PRIVATE DIR ${CMAKE_CURRENT_SOURCE_DIR}) | ||
| 14 | -endif() | ||
| @@ -1,95 +0,0 @@ | |||
| 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 | -/* Generated By CANNBot */ | ||
| 12 | - | ||
| 13 | - | ||
| 14 | - | ||
| 15 | - | ||
| 16 | - | ||
| 17 | - | ||
| 18 | - | ||
| 19 | - | ||
| 20 | - | ||
| 21 | - | ||
| 22 | - | ||
| 23 | - | ||
| 24 | - | ||
| 25 | - | ||
| 26 | - | ||
| 27 | -using namespace ut_util; | ||
| 28 | -using namespace ge; | ||
| 29 | - | ||
| 30 | -class NPUClearFloatStatusTiling : public testing::Test { | ||
| 31 | -protected: | ||
| 32 | - static void SetUpTestCase() { std::cout << "NPUClearFloatStatusTiling SetUp" << std::endl; } | ||
| 33 | - | ||
| 34 | - static void TearDownTestCase() { std::cout << "NPUClearFloatStatusTiling TearDown" << std::endl; } | ||
| 35 | -}; | ||
| 36 | - | ||
| 37 | -TEST_F(NPUClearFloatStatusTiling, npu_clear_float_status_tiling_test1) | ||
| 38 | -{ | ||
| 39 | - std::map<std::string, std::string> socInfos; | ||
| 40 | - std::map<std::string, std::string> aicoreSpec; | ||
| 41 | - std::map<std::string, std::string> intrinsics; | ||
| 42 | - std::string compileInfoStr = R"({"hardware_info":{"UB_SIZE":262144,"CORE_NUM":64}})"; | ||
| 43 | - GetPlatFormInfos(compileInfoStr.c_str(), socInfos, aicoreSpec, intrinsics); | ||
| 44 | - | ||
| 45 | - fe::PlatFormInfos platformInfo; | ||
| 46 | - platformInfo.Init(); | ||
| 47 | - | ||
| 48 | - struct NPUClearFloatStatusCompileInfo { | ||
| 49 | - } compileInfo; | ||
| 50 | - | ||
| 51 | - auto tilingData = gert::TilingData::CreateCap(4096); | ||
| 52 | - auto workspaceHolder = gert::ContinuousVector::Create<size_t>(4096); | ||
| 53 | - auto workspace = reinterpret_cast<gert::ContinuousVector*>(workspaceHolder.get()); | ||
| 54 | - | ||
| 55 | - gert::StorageShape addrShape = {{8}, {8}}; | ||
| 56 | - gert::StorageShape outShape = {{8}, {8}}; | ||
| 57 | - | ||
| 58 | - auto holder = gert::TilingContextFaker() | ||
| 59 | - .SetOpType("NPUClearFloatStatus") | ||
| 60 | - .NodeIoNum(1, 1) | ||
| 61 | - .IrInstanceNum(std::vector<uint32_t>{1}, std::vector<uint32_t>{1}) | ||
| 62 | - .InputShapes({&addrShape}) | ||
| 63 | - .OutputShapes({&outShape}) | ||
| 64 | - .CompileInfo(&compileInfo) | ||
| 65 | - .PlatformInfo(reinterpret_cast<char*>(&platformInfo)) | ||
| 66 | - .NodeInputTd(0, ge::DT_FLOAT, ge::FORMAT_ND, ge::FORMAT_ND) | ||
| 67 | - .NodeOutputTd(0, ge::DT_FLOAT, ge::FORMAT_ND, ge::FORMAT_ND) | ||
| 68 | - .TilingData(tilingData.get()) | ||
| 69 | - .Workspace(workspace) | ||
| 70 | - .Build(); | ||
| 71 | - | ||
| 72 | - gert::TilingContext* tilingContext = holder.GetContext<gert::TilingContext>(); | ||
| 73 | - ASSERT_NE(tilingContext, nullptr); | ||
| 74 | - ASSERT_NE(tilingContext->GetPlatformInfo(), nullptr); | ||
| 75 | - tilingContext->GetPlatformInfo()->SetPlatformRes("SoCInfo", socInfos); | ||
| 76 | - tilingContext->GetPlatformInfo()->SetPlatformRes("AICoreSpec", aicoreSpec); | ||
| 77 | - tilingContext->GetPlatformInfo()->SetCoreNumByCoreType("AICore"); | ||
| 78 | - tilingContext->GetPlatformInfo()->SetPlatformRes("AICoreintrinsicDtypeMap", intrinsics); | ||
| 79 | - std::map<std::string, std::string> socVersionInfos = {{"Short_SoC_version", "Ascend950"}, {"NpuArch", "3510"}}; | ||
| 80 | - tilingContext->GetPlatformInfo()->SetPlatformRes("version", socVersionInfos); | ||
| 81 | - | ||
| 82 | - std::string opType("NPUClearFloatStatus"); | ||
| 83 | - ASSERT_NE(gert::OpImplRegistry::GetInstance().GetOpImpl(opType.c_str()), nullptr); | ||
| 84 | - auto tilingFunc = gert::OpImplRegistry::GetInstance().GetOpImpl(opType.c_str())->tiling; | ||
| 85 | - ASSERT_NE(tilingFunc, nullptr); | ||
| 86 | - | ||
| 87 | - EXPECT_EQ(tilingFunc(tilingContext), ge::GRAPH_SUCCESS); | ||
| 88 | - EXPECT_EQ(tilingContext->GetTilingKey(), 0U); | ||
| 89 | - auto* tilingDataPtr = tilingContext->GetTilingData<NPUClearFloatStatusTilingData>(); | ||
| 90 | - ASSERT_NE(tilingDataPtr, nullptr); | ||
| 91 | - EXPECT_EQ(tilingDataPtr->needCoreNum, 64); | ||
| 92 | - EXPECT_EQ(tilingContext->GetBlockDim(), 64U); | ||
| 93 | - ASSERT_NE(tilingContext->GetWorkspaceSizes(1), nullptr); | ||
| 94 | - EXPECT_EQ(tilingContext->GetWorkspaceSizes(1)[0], 16777216U); | ||
| 95 | -} | ||
| @@ -1,47 +0,0 @@ | |||
| 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 | -/* Generated By CANNBot */ | ||
| 12 | - | ||
| 13 | - | ||
| 14 | - | ||
| 15 | - | ||
| 16 | - | ||
| 17 | - | ||
| 18 | - | ||
| 19 | - | ||
| 20 | -class NPUClearFloatStatusInfershape : public testing::Test { | ||
| 21 | -protected: | ||
| 22 | - static void SetUpTestCase() { std::cout << "NPUClearFloatStatusInfershape SetUp" << std::endl; } | ||
| 23 | - | ||
| 24 | - static void TearDownTestCase() { std::cout << "NPUClearFloatStatusInfershape TearDown" << std::endl; } | ||
| 25 | -}; | ||
| 26 | - | ||
| 27 | -TEST_F(NPUClearFloatStatusInfershape, npu_clear_float_status_infershape_test1) | ||
| 28 | -{ | ||
| 29 | - std::string opType("NPUClearFloatStatus"); | ||
| 30 | - ASSERT_NE(gert::OpImplRegistry::GetInstance().GetOpImpl(opType.c_str()), nullptr); | ||
| 31 | - auto inferShapeFunc = gert::OpImplRegistry::GetInstance().GetOpImpl(opType.c_str())->infer_shape; | ||
| 32 | - ASSERT_NE(inferShapeFunc, nullptr); | ||
| 33 | - | ||
| 34 | - gert::StorageShape addrShape = {{8}, {8}}; | ||
| 35 | - gert::StorageShape outputShape = {}; | ||
| 36 | - | ||
| 37 | - auto holder = gert::InferShapeContextFaker() | ||
| 38 | - .NodeIoNum(1, 1) | ||
| 39 | - .IrInstanceNum({1}, {1}) | ||
| 40 | - .InputShapes({&addrShape}) | ||
| 41 | - .OutputShapes({&outputShape}) | ||
| 42 | - .NodeInputTd(0, ge::DT_FLOAT, ge::FORMAT_ND, ge::FORMAT_ND) | ||
| 43 | - .NodeOutputTd(0, ge::DT_FLOAT, ge::FORMAT_ND, ge::FORMAT_ND) | ||
| 44 | - .Build(); | ||
| 45 | - | ||
| 46 | - EXPECT_EQ(inferShapeFunc(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_SUCCESS); | ||
| 47 | -} | ||
| @@ -907,17 +907,6 @@ | |||
| 907 | <td>AI Core</td> | 907 | <td>AI Core</td> |
| 908 | <td>分配8个float32零值的Tensor,用于NPU溢出状态检测。</td> | 908 | <td>分配8个float32零值的Tensor,用于NPU溢出状态检测。</td> |
| 909 | </tr> | 909 | </tr> |
| 910 | - <tr> | ||
| 911 | - <td>control</td> | ||
| 912 | - <td><a href="../../control/npu_clear_float_status/README.md">npu_clear_float_status</a></td> | ||
| 913 | - <td>√</td> | ||
| 914 | - <td>×</td> | ||
| 915 | - <td>×</td> | ||
| 916 | - <td>√</td> | ||
| 917 | - <td>√</td> | ||
| 918 | - <td>AI Core</td> | ||
| 919 | - <td>清除NPU浮点溢出状态寄存器,输出固定8个float32零值。</td> | ||
| 920 | - </tr> | ||
| 921 | <tr> | 910 | <tr> |
| 922 | <td>control</td> | 911 | <td>control</td> |
| 923 | <td><a href="../../control/update_tensor_desc/README.md">update_tensor_desc</a></td> | 912 | <td><a href="../../control/update_tensor_desc/README.md">update_tensor_desc</a></td> |
| @@ -4818,16 +4807,6 @@ | |||
| 4818 | <td>AI Core</td> | 4807 | <td>AI Core</td> |
| 4819 | <td>对于输入信号的输入通道,提供3维最大池化(Max pooling)操作,输出池化后的值out和索引indices。</td> | 4808 | <td>对于输入信号的输入通道,提供3维最大池化(Max pooling)操作,输出池化后的值out和索引indices。</td> |
| 4820 | </tr> | 4809 | </tr> |
| 4821 | - <tr> | ||
| 4822 | - <td>pooling</td> | ||
| 4823 | - <td><a href="../../pooling/avg_pool_update/README.md">avg_pool_update</a></td> | ||
| 4824 | - <td>✓</td> | ||
| 4825 | - <td>✓</td> | ||
| 4826 | - <td>✗</td> | ||
| 4827 | - <td>✓</td> | ||
| 4828 | - <td>AI Core</td> | ||
| 4829 | - <td>计算平均池化的更新值,将求和池化结果除以池化窗口实际覆盖的有效元素个数得到平均值。</td> | ||
| 4830 | - </tr> | ||
| 4831 | <tr> | 4810 | <tr> |
| 4832 | <td>pooling</td> | 4811 | <td>pooling</td> |
| 4833 | <td><a href="../../pooling/psroi_pooling_v2/README.md">psroi_pooling_v2</a></td> | 4812 | <td><a href="../../pooling/psroi_pooling_v2/README.md">psroi_pooling_v2</a></td> |
| @@ -1,16 +0,0 @@ | |||
| 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 | -set(SUPPORT_COMPUTE_UNIT "ascend950") | ||
| 12 | - | ||
| 13 | -set(SUPPORT_TILING_DIR "arch35") | ||
| 14 | - | ||
| 15 | -add_modules_sources(HOSTNAME ${OPHOST_NAME} MODE PRIVATE DIR ${CMAKE_CURRENT_SOURCE_DIR} OPTYPE avg_pool_update | ||
| 16 | - ACLNNTYPE aclnn_exclude COMPUTE_UNIT ${SUPPORT_COMPUTE_UNIT} TILING_DIR ${SUPPORT_TILING_DIR} DISABLE_IN_OPP TRUE) | ||
| @@ -1,157 +0,0 @@ | |||
| 1 | -# AvgPoolUpdate | ||
| 2 | - | ||
| 3 | -## 产品支持情况 | ||
| 4 | - | ||
| 5 | -| 产品 | 是否支持 | | ||
| 6 | -| :----------------------------------------------------------- | :------: | | ||
| 7 | -| <term>Ascend 950PR/Ascend 950DT</term> | √ | | ||
| 8 | -| <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term> | √ | | ||
| 9 | -| <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term> | √ | | ||
| 10 | -| <term>Atlas 200I/500 A2 推理产品</term> | √ | | ||
| 11 | -| <term>Atlas 推理系列产品</term> | √ | | ||
| 12 | -| <term>Atlas 训练系列产品</term> | √ | | ||
| 13 | - | ||
| 14 | -## 功能说明 | ||
| 15 | - | ||
| 16 | -- 算子功能:计算平均池化的更新值。将求和池化结果除以池化窗口实际覆盖的有效元素个数,得到平均值。x1为求和池化输出,x2为原始输入feature map(仅用于获取输入空间尺寸,不参与数值计算)。 | ||
| 17 | - | ||
| 18 | -- 计算公式: | ||
| 19 | - | ||
| 20 | -$$ | ||
| 21 | -y = x1 \oslash \text{mean\_matrix} | ||
| 22 | -$$ | ||
| 23 | - | ||
| 24 | -其中mean_matrix为池化除数矩阵: | ||
| 25 | - | ||
| 26 | -$$ | ||
| 27 | -\text{mean\_matrix}[h, w] = \text{mean\_h} \times \text{mean\_w} | ||
| 28 | -$$ | ||
| 29 | - | ||
| 30 | -$$ | ||
| 31 | -\text{mean\_h} = \max(\min(\min(h \cdot s_h - p_t + k_h,\; (H_{out}-1-h) \cdot s_h - p_b + k_h),\; \min(k_h,\; H_{in})),\; 1) | ||
| 32 | -$$ | ||
| 33 | - | ||
| 34 | -$$ | ||
| 35 | -\text{mean\_w} = \max(\min(\min(w \cdot s_w - p_l + k_w,\; (W_{out}-1-w) \cdot s_w - p_r + k_w),\; \min(k_w,\; W_{in})),\; 1) | ||
| 36 | -$$ | ||
| 37 | - | ||
| 38 | -式中,$s_h$、$s_w$为strides的H/W分量,$p_t$、$p_b$、$p_l$、$p_r$为pads的上/下/左/右分量,$k_h$、$k_w$为ksize的H/W分量,$H_{in}$、$W_{in}$为x2的空间输入尺寸,$H_{out}$、$W_{out}$为x1的空间输出尺寸。mean_matrix计算使用int64中间变量,除法前cast为x1的dtype。 | ||
| 39 | - | ||
| 40 | -## 参数说明 | ||
| 41 | - | ||
| 42 | -<table style="undefined;table-layout: fixed; width: 980px"><colgroup> | ||
| 43 | - <col style="width: 100px"> | ||
| 44 | - <col style="width: 150px"> | ||
| 45 | - <col style="width: 280px"> | ||
| 46 | - <col style="width: 330px"> | ||
| 47 | - <col style="width: 120px"> | ||
| 48 | - </colgroup> | ||
| 49 | - <thead> | ||
| 50 | - <tr> | ||
| 51 | - <th>参数名</th> | ||
| 52 | - <th>输入/输出/属性</th> | ||
| 53 | - <th>描述</th> | ||
| 54 | - <th>数据类型</th> | ||
| 55 | - <th>数据格式</th> | ||
| 56 | - </tr></thead> | ||
| 57 | - <tbody> | ||
| 58 | - <tr> | ||
| 59 | - <td>x1</td> | ||
| 60 | - <td>输入</td> | ||
| 61 | - <td>求和池化输出结果,公式中的x1。shape为(N, C, H_out, W_out)(NCHW)或(N, H_out, W_out, C)(NHWC)。</td> | ||
| 62 | - <td>FLOAT16、FLOAT</td> | ||
| 63 | - <td>ND</td> | ||
| 64 | - </tr> | ||
| 65 | - <tr> | ||
| 66 | - <td>x2</td> | ||
| 67 | - <td>输入</td> | ||
| 68 | - <td>原始输入feature map,公式中的x2。仅用于获取输入空间尺寸H_in/W_in,不参与数值计算。shape为(N, C, H_in, W_in)(NCHW)或(N, H_in, W_in, C)(NHWC)。</td> | ||
| 69 | - <td>INT4、INT8、FLOAT16、FLOAT</td> | ||
| 70 | - <td>ND</td> | ||
| 71 | - </tr> | ||
| 72 | - <tr> | ||
| 73 | - <td>y</td> | ||
| 74 | - <td>输出</td> | ||
| 75 | - <td>平均池化更新结果,公式中的y。shape与x1相同。</td> | ||
| 76 | - <td>FLOAT16、FLOAT</td> | ||
| 77 | - <td>ND</td> | ||
| 78 | - </tr> | ||
| 79 | - <tr> | ||
| 80 | - <td>ksize</td> | ||
| 81 | - <td>属性</td> | ||
| 82 | - <td>池化窗口大小,长度为4的列表,各维度顺序与data_format一致。</td> | ||
| 83 | - <td>ListInt</td> | ||
| 84 | - <td>-</td> | ||
| 85 | - </tr> | ||
| 86 | - <tr> | ||
| 87 | - <td>strides</td> | ||
| 88 | - <td>属性</td> | ||
| 89 | - <td>池化步长,长度为4的列表,各维度顺序与data_format一致。</td> | ||
| 90 | - <td>ListInt</td> | ||
| 91 | - <td>-</td> | ||
| 92 | - </tr> | ||
| 93 | - <tr> | ||
| 94 | - <td>padding_mode</td> | ||
| 95 | - <td>属性</td> | ||
| 96 | - <td>填充模式,取值范围:CALCULATED、VALID、SAME。默认值为CALCULATED。</td> | ||
| 97 | - <td>String</td> | ||
| 98 | - <td>-</td> | ||
| 99 | - </tr> | ||
| 100 | - <tr> | ||
| 101 | - <td>pads</td> | ||
| 102 | - <td>属性</td> | ||
| 103 | - <td>填充值,长度为4的列表(top, bottom, left, right)。仅在padding_mode为CALCULATED时生效。默认值为{0, 0, 0, 0}。</td> | ||
| 104 | - <td>ListInt</td> | ||
| 105 | - <td>-</td> | ||
| 106 | - </tr> | ||
| 107 | - <tr> | ||
| 108 | - <td>data_format</td> | ||
| 109 | - <td>属性</td> | ||
| 110 | - <td>数据格式,取值范围:NCHW、NHWC。默认值为NHWC。</td> | ||
| 111 | - <td>String</td> | ||
| 112 | - <td>-</td> | ||
| 113 | - </tr> | ||
| 114 | - <tr> | ||
| 115 | - <td>ceil_mode</td> | ||
| 116 | - <td>属性</td> | ||
| 117 | - <td>是否使用ceil模式计算输出尺寸。true为ceil,false为floor。默认值为false。</td> | ||
| 118 | - <td>Bool</td> | ||
| 119 | - <td>-</td> | ||
| 120 | - </tr> | ||
| 121 | - <tr> | ||
| 122 | - <td>exclusive</td> | ||
| 123 | - <td>属性</td> | ||
| 124 | - <td>是否排除padding区域计入窗口。true为排除padding(使用实际覆盖元素个数),false为使用常量窗口大小。默认值为true。</td> | ||
| 125 | - <td>Bool</td> | ||
| 126 | - <td>-</td> | ||
| 127 | - </tr> | ||
| 128 | - </tbody></table> | ||
| 129 | - | ||
| 130 | -## 约束说明 | ||
| 131 | - | ||
| 132 | -- **exclusive约束**:exclusive必须为true。当exclusive为false时算子无意义(池化因子为常量,无需更新),会报错退出。 | ||
| 133 | -- **padding_mode与ceil_mode约束**:当padding_mode为VALID且ceil_mode为false时,算子无意义(池化因子为常量),会报错退出。 | ||
| 134 | -- **输入维度约束**:x1和x2必须为4D张量,支持NCHW和NHWC两种data_format。不支持其他维度数。 | ||
| 135 | -- **ksize约束**:ksize为长度4的列表,各分量必须为正整数。H/W分量(k_h, k_w)的取值由data_format决定位置。 | ||
| 136 | -- **strides约束**:strides为长度4的列表,各分量必须为正整数。H/W分量(s_h, s_w)的取值由data_format决定位置。 | ||
| 137 | -- **pads约束**:pads为长度4的列表(top, bottom, left, right),各分量必须为非负整数。仅在padding_mode为CALCULATED时生效;padding_mode为VALID时pads被忽略(置为0);padding_mode为SAME时pads由框架自动计算。 | ||
| 138 | -- **data_format约束**:仅支持NCHW和NHWC。 | ||
| 139 | -- **padding_mode约束**:仅支持CALCULATED、VALID、SAME三种取值。 | ||
| 140 | -- **dtype约束**:x1和y的dtype必须相同(FLOAT16或FLOAT)。x2的dtype独立,支持INT4、INT8、FLOAT16、FLOAT(x2仅用于获取输入空间尺寸,不参与数值计算)。 | ||
| 141 | - | ||
| 142 | -## 调用说明 | ||
| 143 | - | ||
| 144 | -<table><thead> | ||
| 145 | - <tr> | ||
| 146 | - <th>调用方式</th> | ||
| 147 | - <th>调用样例</th> | ||
| 148 | - <th>说明</th> | ||
| 149 | - </tr></thead> | ||
| 150 | -<tbody> | ||
| 151 | - <tr> | ||
| 152 | - <td>图模式调用</td> | ||
| 153 | - <td><a href="examples/arch35/test_geir_avg_pool_update.cpp">test_geir_avg_pool_update</a></td> | ||
| 154 | - <td>参见<a href="../../docs/zh/invocation/quick_op_invocation.md">算子调用</a>完成算子编译和验证。</td> | ||
| 155 | - </tr> | ||
| 156 | -</tbody> | ||
| 157 | -</table> | ||
| @@ -1,476 +0,0 @@ | |||
| 1 | - | ||
| 2 | -/** | ||
| 3 | - * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 4 | - * This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 5 | - * CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 6 | - * Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 7 | - * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 8 | - * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 9 | - * See LICENSE in the root of the software repository for the full text of the License. | ||
| 10 | - */ | ||
| 11 | - | ||
| 12 | -/* Generated By CANNBot */ | ||
| 13 | - | ||
| 14 | - | ||
| 15 | - | ||
| 16 | - | ||
| 17 | - | ||
| 18 | - | ||
| 19 | - | ||
| 20 | - | ||
| 21 | - | ||
| 22 | - | ||
| 23 | - | ||
| 24 | - | ||
| 25 | - | ||
| 26 | - | ||
| 27 | - | ||
| 28 | - | ||
| 29 | - | ||
| 30 | - | ||
| 31 | - | ||
| 32 | - | ||
| 33 | -constexpr int32_t kFailed = -1; | ||
| 34 | -constexpr int32_t kSuccess = 0; | ||
| 35 | - | ||
| 36 | -using namespace ge; | ||
| 37 | -using std::map; | ||
| 38 | -using std::string; | ||
| 39 | -using std::vector; | ||
| 40 | - | ||
| 41 | -enum RunMode { RUN_MODE_S = 0, RUN_MODE_D = 1 }; | ||
| 42 | - | ||
| 43 | -struct CaseResult { | ||
| 44 | - std::string case_name; | ||
| 45 | - bool build_ok; | ||
| 46 | - bool run_ok; | ||
| 47 | - bool output_exists; | ||
| 48 | - int output_count; | ||
| 49 | - std::string err_msg; | ||
| 50 | -}; | ||
| 51 | - | ||
| 52 | -struct ShapeConfig { | ||
| 53 | - std::vector<int64_t> x1Shape; | ||
| 54 | - std::vector<int64_t> x2Shape; | ||
| 55 | - std::vector<int64_t> ksize; | ||
| 56 | - std::vector<int64_t> strides; | ||
| 57 | - std::string padding_mode; | ||
| 58 | - std::vector<int64_t> pads; | ||
| 59 | - std::string data_format; | ||
| 60 | - bool ceil_mode; | ||
| 61 | - bool exclusive; | ||
| 62 | - std::string name; | ||
| 63 | - // x2 独立 dtype,独立于 x1 | ||
| 64 | - // DT_UNDEFINED 表示跟随 x1 的 loop dtype;指定具体 dtype 时 x2 使用该 dtype | ||
| 65 | - DataType x2Dtype; | ||
| 66 | -}; | ||
| 67 | - | ||
| 68 | - | ||
| 69 | - vector<int64_t> placeholder##inputIndex##_real_shape = inputShape; \ | ||
| 70 | - vector<int64_t> placeholder##inputIndex##_graph_shape = ((mode) == RUN_MODE_D) ? \ | ||
| 71 | - vector<int64_t>( \ | ||
| 72 | - placeholder##inputIndex##_real_shape.size(), -1) : \ | ||
| 73 | - placeholder##inputIndex##_real_shape; \ | ||
| 74 | - auto placeholder##inputIndex = op::Data("placeholder" + inputIndex).set_attr_index(0); \ | ||
| 75 | - TensorDesc placeholder##inputIndex##_desc_graph = TensorDesc(ge::Shape(placeholder##inputIndex##_graph_shape), \ | ||
| 76 | - FORMAT_ND, intputDtype); \ | ||
| 77 | - placeholder##inputIndex##_desc_graph.SetPlacement(ge::kPlacementHost); \ | ||
| 78 | - placeholder##inputIndex##_desc_graph.SetFormat(FORMAT_ND); \ | ||
| 79 | - TensorDesc placeholder##inputIndex##_desc_real = TensorDesc(ge::Shape(placeholder##inputIndex##_real_shape), \ | ||
| 80 | - FORMAT_ND, intputDtype); \ | ||
| 81 | - placeholder##inputIndex##_desc_real.SetPlacement(ge::kPlacementHost); \ | ||
| 82 | - placeholder##inputIndex##_desc_real.SetFormat(FORMAT_ND); \ | ||
| 83 | - placeholder##inputIndex##_desc_real.SetRealDimCnt(placeholder##inputIndex##_real_shape.size()); \ | ||
| 84 | - Tensor tensor_placeholder##inputIndex; \ | ||
| 85 | - ret = GenUniformData(placeholder##inputIndex##_real_shape, tensor_placeholder##inputIndex, \ | ||
| 86 | - placeholder##inputIndex##_desc_real, intputDtype, 2); \ | ||
| 87 | - if (ret != kSuccess) { \ | ||
| 88 | - printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \ | ||
| 89 | - return kFailed; \ | ||
| 90 | - } \ | ||
| 91 | - placeholder##inputIndex.update_input_desc_x(placeholder##inputIndex##_desc_graph); \ | ||
| 92 | - placeholder##inputIndex.update_output_desc_y(placeholder##inputIndex##_desc_graph); \ | ||
| 93 | - input.push_back(tensor_placeholder##inputIndex); \ | ||
| 94 | - graph.AddOp(placeholder##inputIndex); \ | ||
| 95 | - add1.set_input_##inputName(placeholder##inputIndex); \ | ||
| 96 | - inputs.push_back(placeholder##inputIndex); | ||
| 97 | - | ||
| 98 | - | ||
| 99 | - vector<int64_t> output##outputIndex##_graph_shape = ((mode) == RUN_MODE_D) ? \ | ||
| 100 | - vector<int64_t>(outputShape.size(), -1) : \ | ||
| 101 | - outputShape; \ | ||
| 102 | - TensorDesc outputName##outputIndex##_desc = TensorDesc(ge::Shape(output##outputIndex##_graph_shape), FORMAT_ND, \ | ||
| 103 | - outputDtype); \ | ||
| 104 | - add1.update_output_desc_##outputName(outputName##outputIndex##_desc); | ||
| 105 | - | ||
| 106 | -string GetTime() | ||
| 107 | -{ | ||
| 108 | - time_t timep; | ||
| 109 | - time(&timep); | ||
| 110 | - char tmp[64]; | ||
| 111 | - struct tm tmResult; | ||
| 112 | - localtime_r(&timep, &tmResult); | ||
| 113 | - strftime(tmp, sizeof(tmp), "%Y-%m-%d %H:%M:%S,000", &tmResult); | ||
| 114 | - return tmp; | ||
| 115 | -} | ||
| 116 | - | ||
| 117 | -uint32_t GetDataTypeSize(DataType dt) | ||
| 118 | -{ | ||
| 119 | - if (dt == ge::DT_FLOAT) { | ||
| 120 | - return 4; | ||
| 121 | - } else if (dt == ge::DT_FLOAT16) { | ||
| 122 | - return 2; | ||
| 123 | - } else if (dt == ge::DT_BF16) { | ||
| 124 | - return 2; | ||
| 125 | - } else if (dt == ge::DT_INT16) { | ||
| 126 | - return 2; | ||
| 127 | - } else if (dt == ge::DT_UINT16) { | ||
| 128 | - return 2; | ||
| 129 | - } else if (dt == ge::DT_INT32) { | ||
| 130 | - return 4; | ||
| 131 | - } else if (dt == ge::DT_UINT32) { | ||
| 132 | - return 4; | ||
| 133 | - } else if (dt == ge::DT_INT64) { | ||
| 134 | - return 8; | ||
| 135 | - } else if (dt == ge::DT_UINT64) { | ||
| 136 | - return 8; | ||
| 137 | - } else if (dt == ge::DT_INT8) { | ||
| 138 | - return 1; | ||
| 139 | - } else if (dt == ge::DT_INT4) { | ||
| 140 | - return 1; | ||
| 141 | - } | ||
| 142 | - return 1; | ||
| 143 | -} | ||
| 144 | - | ||
| 145 | -// FP16 位模式构造:float→FP16 转换(纯位运算,不依赖 __fp16) | ||
| 146 | -static uint16_t FloatToHalfBits(float fval) | ||
| 147 | -{ | ||
| 148 | - uint32_t bits; | ||
| 149 | - memcpy(&bits, &fval, sizeof(bits)); | ||
| 150 | - uint32_t sign = (bits >> 16) & 0x8000; | ||
| 151 | - int32_t exponent = ((bits >> 23) & 0xff) - 127 + 15; | ||
| 152 | - uint32_t mantissa = (bits >> 13) & 0x3ff; | ||
| 153 | - if (exponent <= 0) { | ||
| 154 | - return sign; // 零或非规格化数简化为零 | ||
| 155 | - } | ||
| 156 | - if (exponent >= 0x1f) { | ||
| 157 | - return sign | 0x7c00; // 无穷大 | ||
| 158 | - } | ||
| 159 | - return static_cast<uint16_t>(sign | (exponent << 10) | mantissa); | ||
| 160 | -} | ||
| 161 | - | ||
| 162 | -int32_t GenUniformData(vector<int64_t> shapes, Tensor& input_tensor, TensorDesc& input_tensor_desc, DataType data_type, | ||
| 163 | - int value) | ||
| 164 | -{ | ||
| 165 | - input_tensor_desc.SetRealDimCnt(shapes.size()); | ||
| 166 | - size_t size = 1; | ||
| 167 | - for (uint32_t i = 0; i < shapes.size(); i++) { | ||
| 168 | - size *= shapes[i]; | ||
| 169 | - } | ||
| 170 | - size_t data_len = size * GetDataTypeSize(data_type); | ||
| 171 | - uint8_t* pData = new (std::nothrow) uint8_t[data_len]; | ||
| 172 | - if (pData == nullptr) { | ||
| 173 | - return kFailed; | ||
| 174 | - } | ||
| 175 | - for (size_t i = 0; i < size; ++i) { | ||
| 176 | - size_t offset = i * GetDataTypeSize(data_type); | ||
| 177 | - if (data_type == ge::DT_FLOAT) { | ||
| 178 | - float fval = static_cast<float>(value); | ||
| 179 | - memcpy(pData + offset, &fval, GetDataTypeSize(data_type)); | ||
| 180 | - } else if (data_type == ge::DT_FLOAT16) { | ||
| 181 | - // FP16 位模式构造:float->FP16 转换,避免直接截取 float 地址导致零值 | ||
| 182 | - float fval = static_cast<float>(value); | ||
| 183 | - uint16_t fp16Bits = FloatToHalfBits(fval); | ||
| 184 | - memcpy(pData + offset, &fp16Bits, sizeof(uint16_t)); | ||
| 185 | - } else { | ||
| 186 | - memset(pData + offset, value, GetDataTypeSize(data_type)); | ||
| 187 | - } | ||
| 188 | - } | ||
| 189 | - input_tensor = Tensor(input_tensor_desc, pData, data_len); | ||
| 190 | - // Tensor 构造为深拷贝,delete 源数据安全 | ||
| 191 | - delete[] pData; | ||
| 192 | - return kSuccess; | ||
| 193 | -} | ||
| 194 | - | ||
| 195 | -int CreateOppInGraph(RunMode mode, DataType inDtype, const ShapeConfig& cfg, std::vector<ge::Tensor>& input, | ||
| 196 | - std::vector<Operator>& inputs, std::vector<Operator>& outputs, Graph& graph) | ||
| 197 | -{ | ||
| 198 | - Status ret = kSuccess; | ||
| 199 | - // 自定义代码:添加单算子定义到图中 | ||
| 200 | - auto add1 = op::AvgPoolUpdate("add1"); | ||
| 201 | - | ||
| 202 | - // 属性设置 | ||
| 203 | - add1.set_attr_ksize(cfg.ksize); | ||
| 204 | - add1.set_attr_strides(cfg.strides); | ||
| 205 | - add1.SetAttr("padding_mode", cfg.padding_mode.c_str()); | ||
| 206 | - add1.set_attr_pads(cfg.pads); | ||
| 207 | - add1.SetAttr("data_format", cfg.data_format.c_str()); | ||
| 208 | - add1.set_attr_ceil_mode(cfg.ceil_mode); | ||
| 209 | - add1.set_attr_exclusive(cfg.exclusive); | ||
| 210 | - | ||
| 211 | - // x1 dtype 始终使用 loop dtype(inDtype) | ||
| 212 | - ADD_INPUT_MODE(1, x1, inDtype, cfg.x1Shape, mode); | ||
| 213 | - // x2 dtype 独立于 x1 | ||
| 214 | - // cfg.x2Dtype == DT_UNDEFINED 时跟随 inDtype,否则使用指定 dtype | ||
| 215 | - DataType x2Actual = (cfg.x2Dtype == ge::DT_UNDEFINED) ? inDtype : cfg.x2Dtype; | ||
| 216 | - ADD_INPUT_MODE(2, x2, x2Actual, cfg.x2Shape, mode); | ||
| 217 | - | ||
| 218 | - // y dtype 与 x1 一致 | ||
| 219 | - ADD_OUTPUT_MODE(1, y, inDtype, cfg.x1Shape, mode); | ||
| 220 | - | ||
| 221 | - outputs.push_back(add1); | ||
| 222 | - // 添加完毕 | ||
| 223 | - return kSuccess; | ||
| 224 | -} | ||
| 225 | - | ||
| 226 | -CaseResult RunOneCase(ge::Session* session, uint32_t graph_id, RunMode mode, DataType dtype, const ShapeConfig& cfg, | ||
| 227 | - const std::string& case_name) | ||
| 228 | -{ | ||
| 229 | - CaseResult r; | ||
| 230 | - r.case_name = case_name; | ||
| 231 | - r.build_ok = false; | ||
| 232 | - r.run_ok = false; | ||
| 233 | - r.output_exists = false; | ||
| 234 | - r.output_count = 0; | ||
| 235 | - r.err_msg = ""; | ||
| 236 | - | ||
| 237 | - std::string graph_name = "tc_ge_irrun_test_" + std::to_string(graph_id); | ||
| 238 | - Graph graph(graph_name.c_str()); | ||
| 239 | - std::vector<ge::Tensor> input; | ||
| 240 | - std::vector<Operator> inputs{}; | ||
| 241 | - std::vector<Operator> outputs{}; | ||
| 242 | - | ||
| 243 | - Status ret = CreateOppInGraph(mode, dtype, cfg, input, inputs, outputs, graph); | ||
| 244 | - if (ret != kSuccess) { | ||
| 245 | - r.err_msg = "CreateOppInGraph failed"; | ||
| 246 | - return r; | ||
| 247 | - } | ||
| 248 | - if (!inputs.empty() && !outputs.empty()) { | ||
| 249 | - graph.SetInputs(inputs).SetOutputs(outputs); | ||
| 250 | - } | ||
| 251 | - | ||
| 252 | - std::map<AscendString, AscendString> graph_options = {}; | ||
| 253 | - ret = session->AddGraph(graph_id, graph, graph_options); | ||
| 254 | - if (ret != kSuccess) { | ||
| 255 | - r.err_msg = "AddGraph failed, ret=" + std::to_string(ret); | ||
| 256 | - return r; | ||
| 257 | - } | ||
| 258 | - r.build_ok = true; | ||
| 259 | - | ||
| 260 | - std::vector<ge::Tensor> output; | ||
| 261 | - ret = session->RunGraph(graph_id, input, output); | ||
| 262 | - session->RemoveGraph(graph_id); | ||
| 263 | - if (ret != kSuccess) { | ||
| 264 | - r.err_msg = "RunGraph failed, ret=" + std::to_string(ret); | ||
| 265 | - return r; | ||
| 266 | - } | ||
| 267 | - r.run_ok = true; | ||
| 268 | - r.output_count = output.size(); | ||
| 269 | - r.output_exists = (output.size() > 0); | ||
| 270 | - | ||
| 271 | - for (size_t i = 0; i < output.size(); i++) { | ||
| 272 | - int64_t shape_size = output[i].GetTensorDesc().GetShape().GetShapeSize(); | ||
| 273 | - printf(" [%s] output[%zu] dtype=%d shape_size=%lld\n", case_name.c_str(), i, | ||
| 274 | - output[i].GetTensorDesc().GetDataType(), (long long)shape_size); | ||
| 275 | - } | ||
| 276 | - | ||
| 277 | - return r; | ||
| 278 | -} | ||
| 279 | - | ||
| 280 | -void PrintReport(const std::vector<CaseResult>& results) | ||
| 281 | -{ | ||
| 282 | - printf("\n"); | ||
| 283 | - printf("====================================================================================================\n"); | ||
| 284 | - printf("| %-22s | %-8s | %-9s | %-12s | %-7s | %-20s\n", "Case", "Build", "RunGraph", "OutputExists", "OutCnt", | ||
| 285 | - "ErrMsg"); | ||
| 286 | - printf("----------------------------------------------------------------------------------------------------\n"); | ||
| 287 | - int pass_cnt = 0; | ||
| 288 | - int total = results.size(); | ||
| 289 | - for (const auto& r : results) { | ||
| 290 | - bool pass = r.build_ok && r.run_ok && r.output_exists; | ||
| 291 | - if (pass) | ||
| 292 | - pass_cnt++; | ||
| 293 | - printf("| %-22s | %-8s | %-9s | %-12s | %-7d | %-20s\n", r.case_name.c_str(), r.build_ok ? "OK" : "FAIL", | ||
| 294 | - r.run_ok ? "OK" : "FAIL", r.output_exists ? "OK" : "FAIL", r.output_count, | ||
| 295 | - r.err_msg.empty() ? "-" : r.err_msg.c_str()); | ||
| 296 | - } | ||
| 297 | - printf("====================================================================================================\n"); | ||
| 298 | - printf("Summary: %d/%d passed\n", pass_cnt, total); | ||
| 299 | -} | ||
| 300 | - | ||
| 301 | -int main(int argc, char* argv[]) | ||
| 302 | -{ | ||
| 303 | - printf("%s - INFO - [XIR]: Start to initialize ge using ge global options\n", GetTime().c_str()); | ||
| 304 | - // 设置全局选项 | ||
| 305 | - std::map<AscendString, AscendString> global_options = { | ||
| 306 | - {"ge.exec.deviceId", "0"}, {"ge.graphRunMode", "0"}, {"ge.exec.precision_mode", "must_keep_origin_dtype"}}; | ||
| 307 | - // 初始化图引擎 | ||
| 308 | - Status ret = ge::GEInitialize(global_options); | ||
| 309 | - if (ret != kSuccess) { | ||
| 310 | - printf("%s - INFO - [XIR]: Initialize ge using ge global options failed\n", GetTime().c_str()); | ||
| 311 | - return kFailed; | ||
| 312 | - } | ||
| 313 | - printf("%s - INFO - [XIR]: Initialize ge using ge global options success\n", GetTime().c_str()); | ||
| 314 | - | ||
| 315 | - // dtype 矩阵(取自 reg_op.dtype_set: DT_FLOAT16, DT_FLOAT) | ||
| 316 | - struct DtypeEntry { | ||
| 317 | - DataType dt; | ||
| 318 | - std::string name; | ||
| 319 | - }; | ||
| 320 | - std::vector<DtypeEntry> dtype_list = { | ||
| 321 | - {DT_FLOAT, "FP32"}, | ||
| 322 | - {DT_FLOAT16, "FP16"}, | ||
| 323 | - }; | ||
| 324 | - | ||
| 325 | - // shape 场景矩阵(AvgPoolUpdate 为4D池化算子,需要NCHW/NHWC格式输入) | ||
| 326 | - // 跳过 scalar({1})、1d({64})、8d({2,2,2,2,2,2,2,2}) 场景:池化算子要求4D输入,这些场景不合法 | ||
| 327 | - // x1Shape 为池化输出shape(out_h,out_w),x2Shape 为原始输入feature map shape(input_h,input_w) | ||
| 328 | - // out_h = floor((input_h + pad_t + pad_b - k_h) / stride_h) + 1 | ||
| 329 | - // out_w = floor((input_w + pad_l + pad_r - k_w) / stride_w) + 1 | ||
| 330 | - // x2Dtype: ge::DT_UNDEFINED 表示跟随 x1 loop dtype;指定具体 dtype 时 x2 使用该 dtype | ||
| 331 | - std::vector<ShapeConfig> shape_list = { | ||
| 332 | - // 常规4D NCHW: out=(8-2)/2+1=4 | ||
| 333 | - {{1, 3, 4, 4}, | ||
| 334 | - {1, 3, 8, 8}, | ||
| 335 | - {1, 1, 2, 2}, | ||
| 336 | - {1, 1, 2, 2}, | ||
| 337 | - "CALCULATED", | ||
| 338 | - {0, 0, 0, 0}, | ||
| 339 | - "NCHW", | ||
| 340 | - false, | ||
| 341 | - true, | ||
| 342 | - "regular_nchw", | ||
| 343 | - ge::DT_UNDEFINED}, | ||
| 344 | - // 最小有效4D: out=(2-2)/2+1=1 | ||
| 345 | - {{1, 1, 1, 1}, | ||
| 346 | - {1, 1, 2, 2}, | ||
| 347 | - {1, 1, 2, 2}, | ||
| 348 | - {1, 1, 2, 2}, | ||
| 349 | - "CALCULATED", | ||
| 350 | - {0, 0, 0, 0}, | ||
| 351 | - "NCHW", | ||
| 352 | - false, | ||
| 353 | - true, | ||
| 354 | - "minimal", | ||
| 355 | - ge::DT_UNDEFINED}, | ||
| 356 | - // 空batch边界: N=0 | ||
| 357 | - {{0, 3, 4, 4}, | ||
| 358 | - {0, 3, 8, 8}, | ||
| 359 | - {1, 1, 2, 2}, | ||
| 360 | - {1, 1, 2, 2}, | ||
| 361 | - "CALCULATED", | ||
| 362 | - {0, 0, 0, 0}, | ||
| 363 | - "NCHW", | ||
| 364 | - false, | ||
| 365 | - true, | ||
| 366 | - "empty_batch", | ||
| 367 | - ge::DT_UNDEFINED}, | ||
| 368 | - // NHWC格式: out=(8-2)/2+1=4 | ||
| 369 | - {{1, 4, 4, 3}, | ||
| 370 | - {1, 8, 8, 3}, | ||
| 371 | - {1, 2, 2, 1}, | ||
| 372 | - {1, 2, 2, 1}, | ||
| 373 | - "CALCULATED", | ||
| 374 | - {0, 0, 0, 0}, | ||
| 375 | - "NHWC", | ||
| 376 | - false, | ||
| 377 | - true, | ||
| 378 | - "nhwc", | ||
| 379 | - ge::DT_UNDEFINED}, | ||
| 380 | - // 带padding: out=(10+1+1-3)/2+1=5 | ||
| 381 | - {{1, 3, 5, 5}, | ||
| 382 | - {1, 3, 10, 10}, | ||
| 383 | - {1, 1, 3, 3}, | ||
| 384 | - {1, 1, 2, 2}, | ||
| 385 | - "CALCULATED", | ||
| 386 | - {1, 1, 1, 1}, | ||
| 387 | - "NCHW", | ||
| 388 | - false, | ||
| 389 | - true, | ||
| 390 | - "with_pads", | ||
| 391 | - ge::DT_UNDEFINED}, | ||
| 392 | - {{1, 3, 4, 4}, | ||
| 393 | - {1, 3, 8, 8}, | ||
| 394 | - {1, 1, 2, 2}, | ||
| 395 | - {1, 1, 2, 2}, | ||
| 396 | - "CALCULATED", | ||
| 397 | - {0, 0, 0, 0}, | ||
| 398 | - "NCHW", | ||
| 399 | - false, | ||
| 400 | - true, | ||
| 401 | - "x2_int8", | ||
| 402 | - ge::DT_INT8}, | ||
| 403 | - {{1, 3, 4, 4}, | ||
| 404 | - {1, 3, 8, 8}, | ||
| 405 | - {1, 1, 2, 2}, | ||
| 406 | - {1, 1, 2, 2}, | ||
| 407 | - "VALID", | ||
| 408 | - {0, 0, 0, 0}, | ||
| 409 | - "NCHW", | ||
| 410 | - true, | ||
| 411 | - true, | ||
| 412 | - "valid_ceil", | ||
| 413 | - ge::DT_UNDEFINED}, | ||
| 414 | - {{1, 3, 4, 4}, | ||
| 415 | - {1, 3, 8, 8}, | ||
| 416 | - {1, 1, 2, 2}, | ||
| 417 | - {1, 1, 2, 2}, | ||
| 418 | - "SAME", | ||
| 419 | - {0, 0, 0, 0}, | ||
| 420 | - "NCHW", | ||
| 421 | - false, | ||
| 422 | - true, | ||
| 423 | - "same_mode", | ||
| 424 | - ge::DT_UNDEFINED}, | ||
| 425 | - }; | ||
| 426 | - | ||
| 427 | - // 单 session 复用 | ||
| 428 | - std::map<AscendString, AscendString> build_options = {}; | ||
| 429 | - ge::Session* session = new (std::nothrow) Session(build_options); | ||
| 430 | - if (session == nullptr) { | ||
| 431 | - printf("%s - ERROR - [XIR]: create session failed\n", GetTime().c_str()); | ||
| 432 | - ge::GEFinalize(); | ||
| 433 | - return kFailed; | ||
| 434 | - } | ||
| 435 | - | ||
| 436 | - std::vector<CaseResult> results; | ||
| 437 | - uint32_t graph_id = 0; | ||
| 438 | - | ||
| 439 | - // N_dtype × N_shape × 2 mode 全矩阵 | ||
| 440 | - for (const auto& d : dtype_list) { | ||
| 441 | - for (const auto& s : shape_list) { | ||
| 442 | - for (auto mode : {RUN_MODE_S, RUN_MODE_D}) { | ||
| 443 | - std::string mode_name = (mode == RUN_MODE_S) ? "S" : "D"; | ||
| 444 | - std::string case_name = d.name + "_" + s.name + "_" + mode_name; | ||
| 445 | - printf("\n%s - INFO - [XIR]: ===== %s =====\n", GetTime().c_str(), case_name.c_str()); | ||
| 446 | - CaseResult r = RunOneCase(session, graph_id, mode, d.dt, s, case_name); | ||
| 447 | - results.push_back(r); | ||
| 448 | - graph_id++; | ||
| 449 | - } | ||
| 450 | - } | ||
| 451 | - } | ||
| 452 | - | ||
| 453 | - PrintReport(results); | ||
| 454 | - | ||
| 455 | - bool all_pass = true; | ||
| 456 | - for (const auto& r : results) { | ||
| 457 | - if (!r.build_ok || !r.run_ok || !r.output_exists) { | ||
| 458 | - all_pass = false; | ||
| 459 | - } | ||
| 460 | - } | ||
| 461 | - if (all_pass) { | ||
| 462 | - printf("\n%s - INFO - [XIR]: ALL CASES PASSED\n", GetTime().c_str()); | ||
| 463 | - } else { | ||
| 464 | - printf("\n%s - ERROR - [XIR]: SOME CASES FAILED, see report above\n", GetTime().c_str()); | ||
| 465 | - } | ||
| 466 | - | ||
| 467 | - delete session; | ||
| 468 | - printf("%s - INFO - [XIR]: Start to finalize ir graph session\n", GetTime().c_str()); | ||
| 469 | - ret = ge::GEFinalize(); | ||
| 470 | - if (ret != kSuccess) { | ||
| 471 | - printf("%s - INFO - [XIR]: Finalize ir graph session failed\n", GetTime().c_str()); | ||
| 472 | - return kFailed; | ||
| 473 | - } | ||
| 474 | - printf("%s - INFO - [XIR]: Finalize ir graph session success\n", GetTime().c_str()); | ||
| 475 | - return all_pass ? kSuccess : kFailed; | ||
| 476 | -} | ||
| @@ -1,40 +0,0 @@ | |||
| 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 | -/* Generated By CANNBot */ | ||
| 12 | - | ||
| 13 | -/*! | ||
| 14 | - * \file avg_pool_update_graph_infer.cpp | ||
| 15 | - * \brief avg_pool_update operator graph infer resource | ||
| 16 | - */ | ||
| 17 | - | ||
| 18 | - | ||
| 19 | - | ||
| 20 | - | ||
| 21 | -namespace ops { | ||
| 22 | -using namespace ge; | ||
| 23 | - | ||
| 24 | -static constexpr int64_t IDX_0 = 0; | ||
| 25 | - | ||
| 26 | -static ge::graphStatus InferDataTypeAvgPoolUpdate(gert::InferDataTypeContext* context) | ||
| 27 | -{ | ||
| 28 | - OP_LOGD(context->GetNodeName(), "Begin to do InferDataTypeAvgPoolUpdate"); | ||
| 29 | - | ||
| 30 | - // 输出 y 的 dtype 与第一个输入 x1 相同 | ||
| 31 | - ge::DataType yDtype = context->GetInputDataType(IDX_0); | ||
| 32 | - context->SetOutputDataType(IDX_0, yDtype); | ||
| 33 | - | ||
| 34 | - OP_LOGD(context->GetNodeName(), "End to do InferDataTypeAvgPoolUpdate"); | ||
| 35 | - return GRAPH_SUCCESS; | ||
| 36 | -} | ||
| 37 | - | ||
| 38 | -IMPL_OP(AvgPoolUpdate).InferDataType(InferDataTypeAvgPoolUpdate); | ||
| 39 | - | ||
| 40 | -} // namespace ops | ||
| @@ -1,62 +0,0 @@ | |||
| 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 | -/* Generated By CANNBot */ | ||
| 12 | - | ||
| 13 | -/*! | ||
| 14 | - * \file avg_pool_update_proto.h | ||
| 15 | - * \brief Prototype definition (REG_OP) for avg_pool_update operator | ||
| 16 | - */ | ||
| 17 | - | ||
| 18 | - | ||
| 19 | - | ||
| 20 | - | ||
| 21 | - | ||
| 22 | - | ||
| 23 | -namespace ge { | ||
| 24 | - | ||
| 25 | -/** | ||
| 26 | - *@brief Average pooling update operator. | ||
| 27 | - *@par Inputs: | ||
| 28 | - *Two inputs, including: | ||
| 29 | - * @li x1: A Tensor. Must be one of the following types: float16, float32. | ||
| 30 | - * @li x2: A Tensor. Must be one of the following types: int4, int8, float16, float32. \n | ||
| 31 | - | ||
| 32 | - *@par Attributes: | ||
| 33 | - * @li ksize: A required ListInt. The size of the sliding window for each dimension of the input tensor. | ||
| 34 | - * @li strides: A required ListInt. The stride of the sliding window for each dimension of the input tensor. | ||
| 35 | - * @li padding_mode: An optional String. Padding mode, defaults to "CALCULATED". | ||
| 36 | - * @li pads: An optional ListInt. Padding sizes, defaults to {0, 0, 0, 0}. | ||
| 37 | - * @li data_format: An optional String. Data format, defaults to "NHWC". | ||
| 38 | - * @li ceil_mode: An optional Bool. Whether to use ceil mode, defaults to false. | ||
| 39 | - * @li exclusive: An optional Bool. Whether to use exclusive mode, defaults to true. \n | ||
| 40 | - | ||
| 41 | - *@par Outputs: | ||
| 42 | - *y: A Tensor. Must be one of the following types: float16, float32. | ||
| 43 | - *Has the same type and shape as input x1. | ||
| 44 | - *@par Third-party framework compatibility | ||
| 45 | - *Compatible with the TensorFlow operator AvgPool. | ||
| 46 | - */ | ||
| 47 | -REG_OP(AvgPoolUpdate) | ||
| 48 | - .INPUT(x1, TensorType({DT_FLOAT16, DT_FLOAT})) | ||
| 49 | - .INPUT(x2, TensorType({DT_INT4, DT_INT8, DT_FLOAT16, DT_FLOAT})) | ||
| 50 | - .OUTPUT(y, TensorType({DT_FLOAT16, DT_FLOAT})) | ||
| 51 | - .REQUIRED_ATTR(ksize, ListInt) | ||
| 52 | - .REQUIRED_ATTR(strides, ListInt) | ||
| 53 | - .ATTR(padding_mode, String, "CALCULATED") | ||
| 54 | - .ATTR(pads, ListInt, {0, 0, 0, 0}) | ||
| 55 | - .ATTR(data_format, String, "NHWC") | ||
| 56 | - .ATTR(ceil_mode, Bool, false) | ||
| 57 | - .ATTR(exclusive, Bool, true) | ||
| 58 | - .OP_END_FACTORY_REG(AvgPoolUpdate) | ||
| 59 | - | ||
| 60 | -} // namespace ge | ||
| 61 | - | ||
| 62 | - | ||
| @@ -1,549 +0,0 @@ | |||
| 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 | - * \file avg_pool_update_tiling.cpp | ||
| 13 | - * \brief Tiling implementation for avg_pool_update operator | ||
| 14 | - */ | ||
| 15 | - | ||
| 16 | - | ||
| 17 | - | ||
| 18 | - | ||
| 19 | - | ||
| 20 | - | ||
| 21 | - | ||
| 22 | - | ||
| 23 | - | ||
| 24 | - | ||
| 25 | - | ||
| 26 | - | ||
| 27 | -namespace optiling { | ||
| 28 | - | ||
| 29 | -constexpr size_t WS_ARRAY_SIZE = 512; | ||
| 30 | -constexpr int64_t PER_CORE_MIN = 1024; | ||
| 31 | -constexpr uint32_t DCACHE_SIZE = 32 * 1024; | ||
| 32 | -constexpr uint32_t STATIC_UB_ESTIMATE = 160; // __ubuf__ uint64_t[20] = 160B(32B 对齐) | ||
| 33 | - | ||
| 34 | -struct AvgPoolUpdateCompileInfo {}; | ||
| 35 | - | ||
| 36 | -static constexpr size_t X1_IDX = 0; | ||
| 37 | -static constexpr size_t X2_IDX = 1; | ||
| 38 | -static constexpr size_t NDIM = 4; | ||
| 39 | - | ||
| 40 | -static constexpr size_t KSIZE_IDX = 0; | ||
| 41 | -static constexpr size_t STRIDES_IDX = 1; | ||
| 42 | -static constexpr size_t PADDING_MODE_IDX = 2; | ||
| 43 | -static constexpr size_t PADS_IDX = 3; | ||
| 44 | -static constexpr size_t DATA_FORMAT_IDX = 4; | ||
| 45 | -static constexpr size_t CEIL_MODE_IDX = 5; | ||
| 46 | -static constexpr size_t EXCLUSIVE_IDX = 6; | ||
| 47 | - | ||
| 48 | -// NCHW: N=0, C=1, H=2, W=3 | ||
| 49 | -static constexpr size_t NCHW_C_DIM_IDX = 1; | ||
| 50 | -static constexpr size_t NCHW_H_DIM_IDX = 2; | ||
| 51 | -static constexpr size_t NCHW_W_DIM_IDX = 3; | ||
| 52 | -// NHWC: N=0, H=1, W=2, C=3 | ||
| 53 | -static constexpr size_t NHWC_H_DIM_IDX = 1; | ||
| 54 | -static constexpr size_t NHWC_W_DIM_IDX = 2; | ||
| 55 | -static constexpr size_t NHWC_C_DIM_IDX = 3; | ||
| 56 | - | ||
| 57 | -// CALCULATED mode pads 数组顺序: [top, bottom, left, right] | ||
| 58 | -static constexpr size_t PAD_TOP_IDX = 0; | ||
| 59 | -static constexpr size_t PAD_BOTTOM_IDX = 1; | ||
| 60 | -static constexpr size_t PAD_LEFT_IDX = 2; | ||
| 61 | -static constexpr size_t PAD_RIGHT_IDX = 3; | ||
| 62 | - | ||
| 63 | -// 统一 data_format 解析,避免 TilingFunc 与 ParseAttrs 两处重复 strcmp | ||
| 64 | -struct DataFormatLayout { | ||
| 65 | - size_t hIdx, wIdx, cIdx; | ||
| 66 | - bool isNhwc; | ||
| 67 | -}; | ||
| 68 | - | ||
| 69 | -static DataFormatLayout ParseDataFormat(const char* dataFormat) | ||
| 70 | -{ | ||
| 71 | - if (strcmp(dataFormat, "NCHW") == 0) { | ||
| 72 | - return {NCHW_H_DIM_IDX, NCHW_W_DIM_IDX, NCHW_C_DIM_IDX, false}; | ||
| 73 | - } | ||
| 74 | - return {NHWC_H_DIM_IDX, NHWC_W_DIM_IDX, NHWC_C_DIM_IDX, true}; | ||
| 75 | -} | ||
| 76 | - | ||
| 77 | -// 统一创建一次 PlatformAscendC,避免重复创建 | ||
| 78 | -static ge::graphStatus GetPlatformInfo(gert::TilingContext* context, uint64_t& ubSize, int64_t& coreNum, | ||
| 79 | - uint64_t& sysWorkspaceSize) | ||
| 80 | -{ | ||
| 81 | - fe::PlatFormInfos* platformInfoPtr = context->GetPlatformInfo(); | ||
| 82 | - OP_CHECK_NULL_WITH_CONTEXT(context, platformInfoPtr); | ||
| 83 | - auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfoPtr); | ||
| 84 | - coreNum = ascendcPlatform.GetCoreNumAiv(); | ||
| 85 | - OP_CHECK_IF(coreNum == 0, OP_LOGE(context, "coreNum is 0"), return ge::GRAPH_FAILED); | ||
| 86 | - ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSize); | ||
| 87 | - OP_CHECK_IF(ubSize == 0, OP_LOGE(context, "ubSize is 0"), return ge::GRAPH_FAILED); | ||
| 88 | - sysWorkspaceSize = ascendcPlatform.GetLibApiWorkSpaceSize(); | ||
| 89 | - return ge::GRAPH_SUCCESS; | ||
| 90 | -} | ||
| 91 | - | ||
| 92 | -// x1/y: FP16/FP32;x2: INT4/INT8/FP16/FP32(独立校验,不要求与 x1 一致) | ||
| 93 | -// 算子 dtype 在编译期由 DTYPE_X1 宏决定,tiling 阶段只校验不返回 dtype | ||
| 94 | -static ge::graphStatus ValidateDtype(gert::TilingContext* context) | ||
| 95 | -{ | ||
| 96 | - auto x1Desc = context->GetInputDesc(X1_IDX); | ||
| 97 | - OP_CHECK_NULL_WITH_CONTEXT(context, x1Desc); | ||
| 98 | - ge::DataType x1Dtype = x1Desc->GetDataType(); | ||
| 99 | - OP_CHECK_IF( | ||
| 100 | - x1Dtype != ge::DT_FLOAT && x1Dtype != ge::DT_FLOAT16, | ||
| 101 | - OP_LOGE_FOR_INVALID_DTYPE(context->GetNodeName(), "x1", Ops::Base::ToString(x1Dtype).c_str(), "float/float16"), | ||
| 102 | - return ge::GRAPH_FAILED); | ||
| 103 | - | ||
| 104 | - // x2 dtype ∈ {INT4, INT8, FP16, FP32}(独立于 x1) | ||
| 105 | - auto x2Desc = context->GetInputDesc(X2_IDX); | ||
| 106 | - OP_CHECK_NULL_WITH_CONTEXT(context, x2Desc); | ||
| 107 | - ge::DataType x2Dtype = x2Desc->GetDataType(); | ||
| 108 | - OP_CHECK_IF( | ||
| 109 | - x2Dtype != ge::DT_INT4 && x2Dtype != ge::DT_INT8 && x2Dtype != ge::DT_FLOAT16 && x2Dtype != ge::DT_FLOAT, | ||
| 110 | - OP_LOGE_FOR_INVALID_DTYPE(context->GetNodeName(), "x2", Ops::Base::ToString(x2Dtype).c_str(), | ||
| 111 | - "int4/int8/float16/float"), | ||
| 112 | - return ge::GRAPH_FAILED); | ||
| 113 | - return ge::GRAPH_SUCCESS; | ||
| 114 | -} | ||
| 115 | - | ||
| 116 | -static ge::graphStatus ValidateShape(gert::TilingContext* context) | ||
| 117 | -{ | ||
| 118 | - auto x1Shape = context->GetInputShape(X1_IDX); | ||
| 119 | - OP_CHECK_NULL_WITH_CONTEXT(context, x1Shape); | ||
| 120 | - auto x1Storage = x1Shape->GetStorageShape(); | ||
| 121 | - OP_CHECK_IF(x1Storage.GetDimNum() != NDIM, | ||
| 122 | - OP_LOGE_FOR_INVALID_SHAPEDIM(context->GetNodeName(), "x1", | ||
| 123 | - (std::to_string(x1Storage.GetDimNum()) + "D").c_str(), "4D"), | ||
| 124 | - return ge::GRAPH_FAILED); | ||
| 125 | - | ||
| 126 | - auto x2Shape = context->GetInputShape(X2_IDX); | ||
| 127 | - OP_CHECK_NULL_WITH_CONTEXT(context, x2Shape); | ||
| 128 | - auto x2Storage = x2Shape->GetStorageShape(); | ||
| 129 | - OP_CHECK_IF(x2Storage.GetDimNum() != NDIM, | ||
| 130 | - OP_LOGE_FOR_INVALID_SHAPEDIM(context->GetNodeName(), "x2", | ||
| 131 | - (std::to_string(x2Storage.GetDimNum()) + "D").c_str(), "4D"), | ||
| 132 | - return ge::GRAPH_FAILED); | ||
| 133 | - return ge::GRAPH_SUCCESS; | ||
| 134 | -} | ||
| 135 | - | ||
| 136 | -static ge::graphStatus ValidateInputs(gert::TilingContext* context) | ||
| 137 | -{ | ||
| 138 | - OP_CHECK_IF(ValidateDtype(context) != ge::GRAPH_SUCCESS, OP_LOGE(context, "ValidateDtype failed"), | ||
| 139 | - return ge::GRAPH_FAILED); | ||
| 140 | - OP_CHECK_IF(ValidateShape(context) != ge::GRAPH_SUCCESS, OP_LOGE(context, "ValidateShape failed"), | ||
| 141 | - return ge::GRAPH_FAILED); | ||
| 142 | - return ge::GRAPH_SUCCESS; | ||
| 143 | -} | ||
| 144 | - | ||
| 145 | -// 计算 trailing padding: pad = max((outDim-1)*stride + kDim - inputDim - padBefore, 0) | ||
| 146 | -// 复用于 VALID(ceilMode=True, padBefore=0) 与 CALCULATED(padBefore=padT/padL) 分支 | ||
| 147 | -static ge::graphStatus ComputeCorrectPad(gert::TilingContext* context, int64_t outDim, int32_t stride, int32_t kDim, | ||
| 148 | - int64_t inputDim, int64_t padBefore, int64_t& result) | ||
| 149 | -{ | ||
| 150 | - // (outDim-1)*stride 溢出保护 | ||
| 151 | - int64_t product = 0; | ||
| 152 | - OP_CHECK_IF(__builtin_mul_overflow(outDim - 1, static_cast<int64_t>(stride), &product), | ||
| 153 | - OP_LOGE(context, "(outDim-1)*stride overflow: outDim=%ld stride=%d", outDim, stride), | ||
| 154 | - return ge::GRAPH_FAILED); | ||
| 155 | - // 逐项检查加/减法溢出 | ||
| 156 | - int64_t step1 = 0; | ||
| 157 | - OP_CHECK_IF(__builtin_add_overflow(product, static_cast<int64_t>(kDim), &step1), | ||
| 158 | - OP_LOGE(context, "product + kDim overflow: product=%ld kDim=%d", product, kDim), | ||
| 159 | - return ge::GRAPH_FAILED); | ||
| 160 | - int64_t step2 = 0; | ||
| 161 | - OP_CHECK_IF(__builtin_sub_overflow(step1, inputDim, &step2), | ||
| 162 | - OP_LOGE(context, "step1 - inputDim overflow: step1=%ld inputDim=%ld", step1, inputDim), | ||
| 163 | - return ge::GRAPH_FAILED); | ||
| 164 | - int64_t pad = 0; | ||
| 165 | - OP_CHECK_IF(__builtin_sub_overflow(step2, padBefore, &pad), | ||
| 166 | - OP_LOGE(context, "step2 - padBefore overflow: step2=%ld padBefore=%ld", step2, padBefore), | ||
| 167 | - return ge::GRAPH_FAILED); | ||
| 168 | - result = pad > 0 ? pad : 0; | ||
| 169 | - return ge::GRAPH_SUCCESS; | ||
| 170 | -} | ||
| 171 | - | ||
| 172 | -// 根据 padding_mode 计算 padT/padB/padL/padR(VALID/SAME/CALCULATED 三分支) | ||
| 173 | -// 前置条件:tiling->outH/outW/inputH/inputW/kH/kW/strideH/strideW 已赋值 | ||
| 174 | -static ge::graphStatus ComputePads(gert::TilingContext* context, const char* paddingMode, bool ceilMode, | ||
| 175 | - const int64_t* padsArr, AvgPoolUpdateTilingData* tiling) | ||
| 176 | -{ | ||
| 177 | - if (strcmp(paddingMode, "VALID") == 0) { | ||
| 178 | - // VALID 分支:ceil_mode=True 时补充 bottom/right padding | ||
| 179 | - tiling->padT = 0; | ||
| 180 | - tiling->padL = 0; | ||
| 181 | - if (ceilMode) { | ||
| 182 | - // VALID+ceilMode: padBefore=0,复用 ComputeCorrectPad | ||
| 183 | - OP_CHECK_IF(ComputeCorrectPad(context, tiling->outH, tiling->strideH, tiling->kH, tiling->inputH, 0, | ||
| 184 | - tiling->padB) != ge::GRAPH_SUCCESS, | ||
| 185 | - OP_LOGE(context, "ComputeCorrectPad padB failed"), return ge::GRAPH_FAILED); | ||
| 186 | - OP_CHECK_IF(ComputeCorrectPad(context, tiling->outW, tiling->strideW, tiling->kW, tiling->inputW, 0, | ||
| 187 | - tiling->padR) != ge::GRAPH_SUCCESS, | ||
| 188 | - OP_LOGE(context, "ComputeCorrectPad padR failed"), return ge::GRAPH_FAILED); | ||
| 189 | - } else { | ||
| 190 | - tiling->padB = 0; | ||
| 191 | - tiling->padR = 0; | ||
| 192 | - } | ||
| 193 | - } else if (strcmp(paddingMode, "SAME") == 0) { | ||
| 194 | - // SAME 分支:(outDim-1)*stride 乘法溢出保护 | ||
| 195 | - int64_t totalPadH = 0; | ||
| 196 | - OP_CHECK_IF( | ||
| 197 | - __builtin_mul_overflow(tiling->outH - 1, static_cast<int64_t>(tiling->strideH), &totalPadH), | ||
| 198 | - OP_LOGE(context, "SAME (outH-1)*strideH overflow: outH=%ld strideH=%d", tiling->outH, tiling->strideH), | ||
| 199 | - return ge::GRAPH_FAILED); | ||
| 200 | - int64_t stepH1 = 0; | ||
| 201 | - OP_CHECK_IF(__builtin_add_overflow(totalPadH, static_cast<int64_t>(tiling->kH), &stepH1), | ||
| 202 | - OP_LOGE(context, "SAME totalPadH + kH overflow: totalPadH=%ld kH=%d", totalPadH, tiling->kH), | ||
| 203 | - return ge::GRAPH_FAILED); | ||
| 204 | - OP_CHECK_IF(__builtin_sub_overflow(stepH1, tiling->inputH, &totalPadH), | ||
| 205 | - OP_LOGE(context, "SAME stepH1 - inputH overflow: stepH1=%ld inputH=%ld", stepH1, tiling->inputH), | ||
| 206 | - return ge::GRAPH_FAILED); | ||
| 207 | - if (totalPadH < 0) { | ||
| 208 | - totalPadH = 0; | ||
| 209 | - } | ||
| 210 | - int64_t totalPadW = 0; | ||
| 211 | - OP_CHECK_IF( | ||
| 212 | - __builtin_mul_overflow(tiling->outW - 1, static_cast<int64_t>(tiling->strideW), &totalPadW), | ||
| 213 | - OP_LOGE(context, "SAME (outW-1)*strideW overflow: outW=%ld strideW=%d", tiling->outW, tiling->strideW), | ||
| 214 | - return ge::GRAPH_FAILED); | ||
| 215 | - int64_t stepW1 = 0; | ||
| 216 | - OP_CHECK_IF(__builtin_add_overflow(totalPadW, static_cast<int64_t>(tiling->kW), &stepW1), | ||
| 217 | - OP_LOGE(context, "SAME totalPadW + kW overflow: totalPadW=%ld kW=%d", totalPadW, tiling->kW), | ||
| 218 | - return ge::GRAPH_FAILED); | ||
| 219 | - OP_CHECK_IF(__builtin_sub_overflow(stepW1, tiling->inputW, &totalPadW), | ||
| 220 | - OP_LOGE(context, "SAME stepW1 - inputW overflow: stepW1=%ld inputW=%ld", stepW1, tiling->inputW), | ||
| 221 | - return ge::GRAPH_FAILED); | ||
| 222 | - if (totalPadW < 0) { | ||
| 223 | - totalPadW = 0; | ||
| 224 | - } | ||
| 225 | - // SAME padding 均分:total/2 给 top/left,余数给 bottom/right | ||
| 226 | - tiling->padT = totalPadH / 2; | ||
| 227 | - tiling->padB = totalPadH - tiling->padT; | ||
| 228 | - tiling->padL = totalPadW / 2; | ||
| 229 | - tiling->padR = totalPadW - tiling->padL; | ||
| 230 | - } else { // CALCULATED | ||
| 231 | - // padT/padL 直接赋值,仅校验非负 | ||
| 232 | - OP_CHECK_IF(padsArr[PAD_TOP_IDX] < 0, | ||
| 233 | - OP_LOGE_FOR_INVALID_VALUE(context->GetNodeName(), "pads", | ||
| 234 | - std::to_string(padsArr[PAD_TOP_IDX]).c_str(), ">= 0"), | ||
| 235 | - return ge::GRAPH_FAILED); | ||
| 236 | - tiling->padT = padsArr[PAD_TOP_IDX]; | ||
| 237 | - OP_CHECK_IF(padsArr[PAD_LEFT_IDX] < 0, | ||
| 238 | - OP_LOGE_FOR_INVALID_VALUE(context->GetNodeName(), "pads", | ||
| 239 | - std::to_string(padsArr[PAD_LEFT_IDX]).c_str(), ">= 0"), | ||
| 240 | - return ge::GRAPH_FAILED); | ||
| 241 | - tiling->padL = padsArr[PAD_LEFT_IDX]; | ||
| 242 | - // padB/padR 根据 output 尺寸反推(信任框架 Infershape 的输出尺寸而非用户传入的 pad_bottom/pad_right) | ||
| 243 | - // 复用 ComputeCorrectPad,padBefore=padT/padL | ||
| 244 | - OP_CHECK_IF(ComputeCorrectPad(context, tiling->outH, tiling->strideH, tiling->kH, tiling->inputH, tiling->padT, | ||
| 245 | - tiling->padB) != ge::GRAPH_SUCCESS, | ||
| 246 | - OP_LOGE(context, "ComputeCorrectPad padB failed"), return ge::GRAPH_FAILED); | ||
| 247 | - OP_CHECK_IF(ComputeCorrectPad(context, tiling->outW, tiling->strideW, tiling->kW, tiling->inputW, tiling->padL, | ||
| 248 | - tiling->padR) != ge::GRAPH_SUCCESS, | ||
| 249 | - OP_LOGE(context, "ComputeCorrectPad padR failed"), return ge::GRAPH_FAILED); | ||
| 250 | - } | ||
| 251 | - | ||
| 252 | - OP_LOGD(context, "ComputePads: kH=%d kW=%d strideH=%d strideW=%d padT=%ld padB=%ld padL=%ld padR=%ld", tiling->kH, | ||
| 253 | - tiling->kW, tiling->strideH, tiling->strideW, tiling->padT, tiling->padB, tiling->padL, tiling->padR); | ||
| 254 | - return ge::GRAPH_SUCCESS; | ||
| 255 | -} | ||
| 256 | - | ||
| 257 | -// 承载属性指针/标量值,使获取与校验赋值逻辑分离 | ||
| 258 | -struct AvgPoolUpdateAttrs { | ||
| 259 | - const int64_t* ksizeArr; | ||
| 260 | - const int64_t* stridesArr; | ||
| 261 | - const char* paddingMode; | ||
| 262 | - const int64_t* padsArr; | ||
| 263 | - const char* dataFormat; | ||
| 264 | - bool ceilMode; | ||
| 265 | - bool exclusive; | ||
| 266 | -}; | ||
| 267 | - | ||
| 268 | -// 集中获取 ksize/strides/padding_mode/pads/ceil_mode/exclusive 属性指针,包含 size 校验 | ||
| 269 | -static ge::graphStatus GetAttrPointers(gert::TilingContext* context, AvgPoolUpdateAttrs* attrOut) | ||
| 270 | -{ | ||
| 271 | - auto attrs = context->GetAttrs(); | ||
| 272 | - OP_CHECK_NULL_WITH_CONTEXT(context, attrs); | ||
| 273 | - | ||
| 274 | - const auto* ksize = attrs->GetAttrPointer<gert::ContinuousVector>(KSIZE_IDX); | ||
| 275 | - OP_CHECK_NULL_WITH_CONTEXT(context, ksize); | ||
| 276 | - OP_CHECK_IF( | ||
| 277 | - ksize->GetSize() < NDIM, | ||
| 278 | - OP_LOGE_FOR_INVALID_VALUE(context->GetNodeName(), "ksize", std::to_string(ksize->GetSize()).c_str(), ">= 4"), | ||
| 279 | - return ge::GRAPH_FAILED); | ||
| 280 | - attrOut->ksizeArr = static_cast<const int64_t*>(ksize->GetData()); | ||
| 281 | - | ||
| 282 | - const auto* strides = attrs->GetAttrPointer<gert::ContinuousVector>(STRIDES_IDX); | ||
| 283 | - OP_CHECK_NULL_WITH_CONTEXT(context, strides); | ||
| 284 | - OP_CHECK_IF(strides->GetSize() < NDIM, | ||
| 285 | - OP_LOGE_FOR_INVALID_VALUE(context->GetNodeName(), "strides", std::to_string(strides->GetSize()).c_str(), | ||
| 286 | - ">= 4"), | ||
| 287 | - return ge::GRAPH_FAILED); | ||
| 288 | - attrOut->stridesArr = static_cast<const int64_t*>(strides->GetData()); | ||
| 289 | - | ||
| 290 | - // padding_mode | ||
| 291 | - attrOut->paddingMode = attrs->GetAttrPointer<char>(PADDING_MODE_IDX); | ||
| 292 | - OP_CHECK_NULL_WITH_CONTEXT(context, attrOut->paddingMode); | ||
| 293 | - | ||
| 294 | - const auto* pads = attrs->GetAttrPointer<gert::ContinuousVector>(PADS_IDX); | ||
| 295 | - OP_CHECK_NULL_WITH_CONTEXT(context, pads); | ||
| 296 | - OP_CHECK_IF( | ||
| 297 | - pads->GetSize() < NDIM, | ||
| 298 | - OP_LOGE_FOR_INVALID_VALUE(context->GetNodeName(), "pads", std::to_string(pads->GetSize()).c_str(), ">= 4"), | ||
| 299 | - return ge::GRAPH_FAILED); | ||
| 300 | - attrOut->padsArr = static_cast<const int64_t*>(pads->GetData()); | ||
| 301 | - | ||
| 302 | - // data_format | ||
| 303 | - attrOut->dataFormat = attrs->GetAttrPointer<char>(DATA_FORMAT_IDX); | ||
| 304 | - OP_CHECK_NULL_WITH_CONTEXT(context, attrOut->dataFormat); | ||
| 305 | - | ||
| 306 | - // ceil_mode | ||
| 307 | - const bool* ceilModePtr = attrs->GetAttrPointer<bool>(CEIL_MODE_IDX); | ||
| 308 | - OP_CHECK_NULL_WITH_CONTEXT(context, ceilModePtr); | ||
| 309 | - attrOut->ceilMode = *ceilModePtr; | ||
| 310 | - | ||
| 311 | - // exclusive | ||
| 312 | - const bool* exclusivePtr = attrs->GetAttrPointer<bool>(EXCLUSIVE_IDX); | ||
| 313 | - OP_CHECK_NULL_WITH_CONTEXT(context, exclusivePtr); | ||
| 314 | - attrOut->exclusive = *exclusivePtr; | ||
| 315 | - | ||
| 316 | - return ge::GRAPH_SUCCESS; | ||
| 317 | -} | ||
| 318 | - | ||
| 319 | -// 仅保留校验和赋值逻辑,属性指针获取已提取至 GetAttrPointers | ||
| 320 | -static ge::graphStatus ParseAttrs(gert::TilingContext* context, AvgPoolUpdateTilingData* tiling, | ||
| 321 | - const DataFormatLayout& layout, const AvgPoolUpdateAttrs& attr) | ||
| 322 | -{ | ||
| 323 | - // 前置校验:exclusive=false 或 VALID+ceil_mode=false 时池化因子为常数,无需此算子 | ||
| 324 | - OP_CHECK_IF( | ||
| 325 | - !attr.exclusive, | ||
| 326 | - OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(context->GetNodeName(), "exclusive", "false", | ||
| 327 | - "AvgPoolUpdate op is not required when pooling factor is a constant"), | ||
| 328 | - return ge::GRAPH_FAILED); | ||
| 329 | - OP_CHECK_IF( | ||
| 330 | - strcmp(attr.paddingMode, "VALID") == 0 && !attr.ceilMode, | ||
| 331 | - OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(context->GetNodeName(), "ceil_mode", "VALID+ceil_mode=false", | ||
| 332 | - "AvgPoolUpdate op is not required when pooling factor is a constant"), | ||
| 333 | - return ge::GRAPH_FAILED); | ||
| 334 | - | ||
| 335 | - // 校验 padding_mode 合法值(与 TBE check_padding 对齐,非 CALCULATED/VALID/SAME 报错) | ||
| 336 | - OP_CHECK_IF( | ||
| 337 | - strcmp(attr.paddingMode, "CALCULATED") != 0 && strcmp(attr.paddingMode, "VALID") != 0 && | ||
| 338 | - strcmp(attr.paddingMode, "SAME") != 0, | ||
| 339 | - OP_LOGE_FOR_INVALID_VALUE(context->GetNodeName(), "padding_mode", attr.paddingMode, "CALCULATED/VALID/SAME"), | ||
| 340 | - return ge::GRAPH_FAILED); | ||
| 341 | - | ||
| 342 | - // 校验 ksize/strides 范围(H/W > 0,int64_t → int32_t 安全窄化) | ||
| 343 | - OP_CHECK_IF(attr.ksizeArr[layout.hIdx] <= 0, | ||
| 344 | - OP_LOGE_FOR_INVALID_VALUE(context->GetNodeName(), "ksize", | ||
| 345 | - std::to_string(attr.ksizeArr[layout.hIdx]).c_str(), "> 0"), | ||
| 346 | - return ge::GRAPH_FAILED); | ||
| 347 | - tiling->kH = static_cast<int32_t>(attr.ksizeArr[layout.hIdx]); | ||
| 348 | - OP_CHECK_IF(attr.ksizeArr[layout.wIdx] <= 0, | ||
| 349 | - OP_LOGE_FOR_INVALID_VALUE(context->GetNodeName(), "ksize", | ||
| 350 | - std::to_string(attr.ksizeArr[layout.wIdx]).c_str(), "> 0"), | ||
| 351 | - return ge::GRAPH_FAILED); | ||
| 352 | - tiling->kW = static_cast<int32_t>(attr.ksizeArr[layout.wIdx]); | ||
| 353 | - // 校验 ksize N/C 维度必须为 1(与 TBE 对齐:ksize[N]/ksize[C]!=1 报错) | ||
| 354 | - OP_CHECK_IF(attr.ksizeArr[0] != 1, | ||
| 355 | - OP_LOGE_FOR_INVALID_VALUE(context->GetNodeName(), "ksize", std::to_string(attr.ksizeArr[0]).c_str(), | ||
| 356 | - "N dimension must be 1"), | ||
| 357 | - return ge::GRAPH_FAILED); | ||
| 358 | - OP_CHECK_IF(attr.ksizeArr[layout.cIdx] != 1, | ||
| 359 | - OP_LOGE_FOR_INVALID_VALUE(context->GetNodeName(), "ksize", | ||
| 360 | - std::to_string(attr.ksizeArr[layout.cIdx]).c_str(), "C dimension must be 1"), | ||
| 361 | - return ge::GRAPH_FAILED); | ||
| 362 | - OP_CHECK_IF(attr.stridesArr[layout.hIdx] <= 0, | ||
| 363 | - OP_LOGE_FOR_INVALID_VALUE(context->GetNodeName(), "strides", | ||
| 364 | - std::to_string(attr.stridesArr[layout.hIdx]).c_str(), "> 0"), | ||
| 365 | - return ge::GRAPH_FAILED); | ||
| 366 | - tiling->strideH = static_cast<int32_t>(attr.stridesArr[layout.hIdx]); | ||
| 367 | - OP_CHECK_IF(attr.stridesArr[layout.wIdx] <= 0, | ||
| 368 | - OP_LOGE_FOR_INVALID_VALUE(context->GetNodeName(), "strides", | ||
| 369 | - std::to_string(attr.stridesArr[layout.wIdx]).c_str(), "> 0"), | ||
| 370 | - return ge::GRAPH_FAILED); | ||
| 371 | - tiling->strideW = static_cast<int32_t>(attr.stridesArr[layout.wIdx]); | ||
| 372 | - // 校验 strides N/C 维度必须为 1(与 TBE 对齐:strides[N]/strides[C]!=1 报错) | ||
| 373 | - OP_CHECK_IF(attr.stridesArr[0] != 1, | ||
| 374 | - OP_LOGE_FOR_INVALID_VALUE(context->GetNodeName(), "strides", std::to_string(attr.stridesArr[0]).c_str(), | ||
| 375 | - "N dimension must be 1"), | ||
| 376 | - return ge::GRAPH_FAILED); | ||
| 377 | - OP_CHECK_IF( | ||
| 378 | - attr.stridesArr[layout.cIdx] != 1, | ||
| 379 | - OP_LOGE_FOR_INVALID_VALUE(context->GetNodeName(), "strides", | ||
| 380 | - std::to_string(attr.stridesArr[layout.cIdx]).c_str(), "C dimension must be 1"), | ||
| 381 | - return ge::GRAPH_FAILED); | ||
| 382 | - | ||
| 383 | - // outH*strideH/outW*strideW 溢出保护(保护 kernel 侧 idx*stride 运算) | ||
| 384 | - int64_t productH = 0; | ||
| 385 | - OP_CHECK_IF(__builtin_mul_overflow(tiling->outH, static_cast<int64_t>(tiling->strideH), &productH), | ||
| 386 | - OP_LOGE(context, "outH * strideH overflow: outH=%ld strideH=%d", tiling->outH, tiling->strideH), | ||
| 387 | - return ge::GRAPH_FAILED); | ||
| 388 | - int64_t productW = 0; | ||
| 389 | - OP_CHECK_IF(__builtin_mul_overflow(tiling->outW, static_cast<int64_t>(tiling->strideW), &productW), | ||
| 390 | - OP_LOGE(context, "outW * strideW overflow: outW=%ld strideW=%d", tiling->outW, tiling->strideW), | ||
| 391 | - return ge::GRAPH_FAILED); | ||
| 392 | - | ||
| 393 | - // padding_mode 解析 pads(三分支计算在 ComputePads 中) | ||
| 394 | - OP_CHECK_IF(ComputePads(context, attr.paddingMode, attr.ceilMode, attr.padsArr, tiling) != ge::GRAPH_SUCCESS, | ||
| 395 | - OP_LOGE(context, "ComputePads failed"), return ge::GRAPH_FAILED); | ||
| 396 | - | ||
| 397 | - return ge::GRAPH_SUCCESS; | ||
| 398 | -} | ||
| 399 | - | ||
| 400 | -static ge::graphStatus ComputeTiling(gert::TilingContext* context, AvgPoolUpdateTilingData* tiling, int64_t totalNum, | ||
| 401 | - int64_t coreNum) | ||
| 402 | -{ | ||
| 403 | - tiling->totalNum = totalNum; | ||
| 404 | - // totalNum + coreNum - 1 溢出保护 | ||
| 405 | - int64_t sum = 0; | ||
| 406 | - OP_CHECK_IF(__builtin_add_overflow(totalNum, coreNum - 1, &sum), | ||
| 407 | - OP_LOGE(context, "totalNum + coreNum - 1 overflow: totalNum=%ld coreNum=%ld", totalNum, coreNum), | ||
| 408 | - return ge::GRAPH_FAILED); | ||
| 409 | - int64_t blockFactor = sum / coreNum; | ||
| 410 | - if (blockFactor < PER_CORE_MIN) { | ||
| 411 | - blockFactor = PER_CORE_MIN; | ||
| 412 | - } | ||
| 413 | - // totalNum + (blockFactor - 1) 溢出保护 | ||
| 414 | - int64_t needCoreSum = 0; | ||
| 415 | - OP_CHECK_IF( | ||
| 416 | - __builtin_add_overflow(totalNum, blockFactor - 1, &needCoreSum), | ||
| 417 | - OP_LOGE(context, "totalNum + blockFactor - 1 overflow: totalNum=%ld blockFactor=%ld", totalNum, blockFactor), | ||
| 418 | - return ge::GRAPH_FAILED); | ||
| 419 | - // needCoreSum/blockFactor ≤ coreNum(典型值 ≤ 48),int32_t 安全 | ||
| 420 | - tiling->needCoreNum = static_cast<int32_t>(needCoreSum / blockFactor); | ||
| 421 | - if (tiling->needCoreNum > coreNum) { | ||
| 422 | - tiling->needCoreNum = static_cast<int32_t>(coreNum); | ||
| 423 | - } | ||
| 424 | - if (tiling->needCoreNum <= 0) { | ||
| 425 | - tiling->needCoreNum = 1; | ||
| 426 | - } | ||
| 427 | - return ge::GRAPH_SUCCESS; | ||
| 428 | -} | ||
| 429 | - | ||
| 430 | -static void DumpTilingData(gert::TilingContext* context, const AvgPoolUpdateTilingData* tiling) | ||
| 431 | -{ | ||
| 432 | - OP_LOGD(context, | ||
| 433 | - "AvgPoolUpdateTilingData: totalNum=%ld, needCoreNum=%d, outH=%ld, outW=%ld, inputH=%ld, inputW=%ld, " | ||
| 434 | - "kH=%d, kW=%d, strideH=%d, strideW=%d, padT=%ld, padB=%ld, padL=%ld, padR=%ld, " | ||
| 435 | - "isNhwc=%d, outC=%ld", | ||
| 436 | - tiling->totalNum, tiling->needCoreNum, tiling->outH, tiling->outW, tiling->inputH, tiling->inputW, | ||
| 437 | - tiling->kH, tiling->kW, tiling->strideH, tiling->strideW, tiling->padT, tiling->padB, tiling->padL, | ||
| 438 | - tiling->padR, tiling->isNhwc, tiling->outC); | ||
| 439 | -} | ||
| 440 | - | ||
| 441 | -// sysWorkspaceSize 由 GetPlatformInfo 统一获取后传入 | ||
| 442 | -static ge::graphStatus GetWorkspaceSize(gert::TilingContext* context, uint64_t sysWorkspaceSize) | ||
| 443 | -{ | ||
| 444 | - size_t userWorkspaceSize = WS_ARRAY_SIZE; | ||
| 445 | - size_t* currentWorkspace = context->GetWorkspaceSizes(1); | ||
| 446 | - OP_CHECK_NULL_WITH_CONTEXT(context, currentWorkspace); | ||
| 447 | - currentWorkspace[0] = userWorkspaceSize + sysWorkspaceSize; | ||
| 448 | - return ge::GRAPH_SUCCESS; | ||
| 449 | -} | ||
| 450 | - | ||
| 451 | -// 从 x1/x2 shape 提取 H/W/C 维度并校验,同时计算 totalNum | ||
| 452 | -static ge::graphStatus ExtractShapeInfo(gert::TilingContext* context, const DataFormatLayout& layout, | ||
| 453 | - AvgPoolUpdateTilingData* tiling, int64_t& totalNum) | ||
| 454 | -{ | ||
| 455 | - // x1 shape: 根据 data_format 获取 H/W/C | ||
| 456 | - auto x1Shape = context->GetInputShape(X1_IDX); | ||
| 457 | - OP_CHECK_NULL_WITH_CONTEXT(context, x1Shape); | ||
| 458 | - auto x1Storage = x1Shape->GetStorageShape(); | ||
| 459 | - tiling->outH = x1Storage.GetDim(layout.hIdx); | ||
| 460 | - tiling->outW = x1Storage.GetDim(layout.wIdx); | ||
| 461 | - tiling->outC = x1Storage.GetDim(layout.cIdx); | ||
| 462 | - totalNum = x1Storage.GetShapeSize(); | ||
| 463 | - OP_CHECK_IF(totalNum <= 0, | ||
| 464 | - OP_LOGE_FOR_INVALID_SHAPESIZE(context->GetNodeName(), "x1", std::to_string(totalNum).c_str(), "> 0"), | ||
| 465 | - return ge::GRAPH_FAILED); | ||
| 466 | - | ||
| 467 | - // x2 shape: 根据 data_format 获取 H/W | ||
| 468 | - auto x2Shape = context->GetInputShape(X2_IDX); | ||
| 469 | - OP_CHECK_NULL_WITH_CONTEXT(context, x2Shape); | ||
| 470 | - auto x2Storage = x2Shape->GetStorageShape(); | ||
| 471 | - tiling->inputH = x2Storage.GetDim(layout.hIdx); | ||
| 472 | - tiling->inputW = x2Storage.GetDim(layout.wIdx); | ||
| 473 | - // x2 空 shape 校验(与 TBE CheckUpdateZeroShape 对齐) | ||
| 474 | - int64_t x2TotalNum = x2Storage.GetShapeSize(); | ||
| 475 | - OP_CHECK_IF(x2TotalNum <= 0, | ||
| 476 | - OP_LOGE_FOR_INVALID_SHAPESIZE(context->GetNodeName(), "x2", std::to_string(x2TotalNum).c_str(), "> 0"), | ||
| 477 | - return ge::GRAPH_FAILED); | ||
| 478 | - | ||
| 479 | - return ge::GRAPH_SUCCESS; | ||
| 480 | -} | ||
| 481 | - | ||
| 482 | -static ge::graphStatus AvgPoolUpdateTilingFunc(gert::TilingContext* context) | ||
| 483 | -{ | ||
| 484 | - uint64_t ubSize; | ||
| 485 | - int64_t coreNum; | ||
| 486 | - uint64_t sysWorkspaceSize; | ||
| 487 | - OP_CHECK_IF(GetPlatformInfo(context, ubSize, coreNum, sysWorkspaceSize) != ge::GRAPH_SUCCESS, | ||
| 488 | - OP_LOGE(context, "GetPlatformInfo error"), return ge::GRAPH_FAILED); | ||
| 489 | - | ||
| 490 | - OP_CHECK_IF(ValidateInputs(context) != ge::GRAPH_SUCCESS, OP_LOGE(context, "ValidateInputs failed"), | ||
| 491 | - return ge::GRAPH_FAILED); | ||
| 492 | - | ||
| 493 | - AvgPoolUpdateTilingData* tiling = context->GetTilingData<AvgPoolUpdateTilingData>(); | ||
| 494 | - OP_CHECK_NULL_WITH_CONTEXT(context, tiling); | ||
| 495 | - OP_CHECK_IF(memset_s(tiling, sizeof(AvgPoolUpdateTilingData), 0, sizeof(AvgPoolUpdateTilingData)) != EOK, | ||
| 496 | - OP_LOGE(context, "set tiling data error"), return ge::GRAPH_FAILED); | ||
| 497 | - | ||
| 498 | - AvgPoolUpdateAttrs attr; | ||
| 499 | - OP_CHECK_IF(GetAttrPointers(context, &attr) != ge::GRAPH_SUCCESS, OP_LOGE(context, "GetAttrPointers failed"), | ||
| 500 | - return ge::GRAPH_FAILED); | ||
| 501 | - | ||
| 502 | - // 校验 data_format 合法值(与 TBE util_avgpool_update_dynamic.py 对齐,非 NCHW/NHWC 报错) | ||
| 503 | - OP_CHECK_IF(strcmp(attr.dataFormat, "NCHW") != 0 && strcmp(attr.dataFormat, "NHWC") != 0, | ||
| 504 | - OP_LOGE_FOR_INVALID_VALUE(context->GetNodeName(), "data_format", attr.dataFormat, "NCHW/NHWC"), | ||
| 505 | - return ge::GRAPH_FAILED); | ||
| 506 | - DataFormatLayout layout = ParseDataFormat(attr.dataFormat); | ||
| 507 | - tiling->isNhwc = layout.isNhwc ? 1 : 0; | ||
| 508 | - | ||
| 509 | - int64_t totalNum = 0; | ||
| 510 | - OP_CHECK_IF(ExtractShapeInfo(context, layout, tiling, totalNum) != ge::GRAPH_SUCCESS, | ||
| 511 | - OP_LOGE(context, "ExtractShapeInfo failed"), return ge::GRAPH_FAILED); | ||
| 512 | - | ||
| 513 | - OP_CHECK_IF(ParseAttrs(context, tiling, layout, attr) != ge::GRAPH_SUCCESS, OP_LOGE(context, "ParseAttrs error"), | ||
| 514 | - return ge::GRAPH_FAILED); | ||
| 515 | - | ||
| 516 | - OP_CHECK_IF(ComputeTiling(context, tiling, totalNum, coreNum) != ge::GRAPH_SUCCESS, | ||
| 517 | - OP_LOGE(context, "ComputeTiling error"), return ge::GRAPH_FAILED); | ||
| 518 | - | ||
| 519 | - DumpTilingData(context, tiling); | ||
| 520 | - | ||
| 521 | - OP_CHECK_IF(GetWorkspaceSize(context, sysWorkspaceSize) != ge::GRAPH_SUCCESS, | ||
| 522 | - OP_LOGE(context, "GetWorkspaceSize error"), return ge::GRAPH_FAILED); | ||
| 523 | - | ||
| 524 | - OP_CHECK_IF((ubSize <= DCACHE_SIZE + STATIC_UB_ESTIMATE), | ||
| 525 | - OP_LOGE(context, "ubSize %lu <= DCACHE_SIZE + STATIC_UB_ESTIMATE", ubSize), return ge::GRAPH_FAILED); | ||
| 526 | - // ascend950 UB ≤ 256KB,减去 DCACHE_SIZE+STATIC_UB_ESTIMATE 后远小于 UINT32_MAX,窄化安全 | ||
| 527 | - auto res = context->SetLocalMemorySize(static_cast<uint32_t>(ubSize - DCACHE_SIZE - STATIC_UB_ESTIMATE)); | ||
| 528 | - OP_CHECK_IF((res != ge::GRAPH_SUCCESS), | ||
| 529 | - OP_LOGE(context, "SetLocalMemorySize failed, ubSize=%lu, DCACHE_SIZE=%u, STATIC_UB_ESTIMATE=%u", ubSize, | ||
| 530 | - DCACHE_SIZE, STATIC_UB_ESTIMATE), | ||
| 531 | - return ge::GRAPH_FAILED); | ||
| 532 | - auto blockDimRet = context->SetBlockDim(static_cast<uint32_t>(tiling->needCoreNum)); | ||
| 533 | - OP_CHECK_IF(blockDimRet != ge::GRAPH_SUCCESS, | ||
| 534 | - OP_LOGE(context, "SetBlockDim failed, needCoreNum=%d", tiling->needCoreNum), return ge::GRAPH_FAILED); | ||
| 535 | - context->SetTilingKey(GET_TPL_TILING_KEY(static_cast<uint64_t>(AVG_POOL_UPDATE_SCH_MODE_ELEMWISE))); | ||
| 536 | - | ||
| 537 | - return ge::GRAPH_SUCCESS; | ||
| 538 | -} | ||
| 539 | - | ||
| 540 | -static ge::graphStatus TilingParseForAvgPoolUpdate(gert::TilingParseContext* context) | ||
| 541 | -{ | ||
| 542 | - (void)context; | ||
| 543 | - return ge::GRAPH_SUCCESS; | ||
| 544 | -} | ||
| 545 | - | ||
| 546 | -IMPL_OP_OPTILING(AvgPoolUpdate) | ||
| 547 | - .Tiling(AvgPoolUpdateTilingFunc) | ||
| 548 | - .TilingParse<AvgPoolUpdateCompileInfo>(TilingParseForAvgPoolUpdate); | ||
| 549 | -} // namespace optiling | ||
| @@ -1,74 +0,0 @@ | |||
| 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 | - * \file avg_pool_update_def.cpp | ||
| 13 | - * \brief Operator definition for avg_pool_update operator | ||
| 14 | - */ | ||
| 15 | - | ||
| 16 | - | ||
| 17 | - | ||
| 18 | -namespace ops { | ||
| 19 | -class AvgPoolUpdate : public OpDef { | ||
| 20 | -public: | ||
| 21 | - explicit AvgPoolUpdate(const char* name) : OpDef(name) | ||
| 22 | - { | ||
| 23 | - // 全笛卡尔积:x1∈{FP16,FP32} × x2∈{INT4,INT8,FP16,FP32} = 8 组 | ||
| 24 | - this->Input("x1") | ||
| 25 | - .ParamType(REQUIRED) | ||
| 26 | - .DataType({ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_FLOAT, | ||
| 27 | - ge::DT_FLOAT, ge::DT_FLOAT}) | ||
| 28 | - .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, | ||
| 29 | - ge::FORMAT_ND, ge::FORMAT_ND}) | ||
| 30 | - .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, | ||
| 31 | - ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) | ||
| 32 | - .AutoContiguous(); | ||
| 33 | - // x2 dtype 独立于 x1(x2 仅用于获取输入空间尺寸) | ||
| 34 | - this->Input("x2") | ||
| 35 | - .ParamType(REQUIRED) | ||
| 36 | - .DataType({ge::DT_INT4, ge::DT_INT8, ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_INT4, ge::DT_INT8, ge::DT_FLOAT16, | ||
| 37 | - ge::DT_FLOAT}) | ||
| 38 | - .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, | ||
| 39 | - ge::FORMAT_ND, ge::FORMAT_ND}) | ||
| 40 | - .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, | ||
| 41 | - ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) | ||
| 42 | - .AutoContiguous(); | ||
| 43 | - this->Output("y") | ||
| 44 | - .ParamType(REQUIRED) | ||
| 45 | - .DataType({ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_FLOAT, | ||
| 46 | - ge::DT_FLOAT, ge::DT_FLOAT}) | ||
| 47 | - .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, | ||
| 48 | - ge::FORMAT_ND, ge::FORMAT_ND}) | ||
| 49 | - .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, | ||
| 50 | - ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) | ||
| 51 | - .AutoContiguous(); | ||
| 52 | - | ||
| 53 | - // 属性注册(顺序与 REG_OP 一致,Tiling 侧按索引 0-6 获取) | ||
| 54 | - this->Attr("ksize").AttrType(REQUIRED).ListInt(); | ||
| 55 | - this->Attr("strides").AttrType(REQUIRED).ListInt(); | ||
| 56 | - this->Attr("padding_mode").AttrType(OPTIONAL).String("CALCULATED"); | ||
| 57 | - this->Attr("pads").AttrType(OPTIONAL).ListInt({0, 0, 0, 0}); | ||
| 58 | - this->Attr("data_format").AttrType(OPTIONAL).String("NHWC"); | ||
| 59 | - this->Attr("ceil_mode").AttrType(OPTIONAL).Bool(false); | ||
| 60 | - this->Attr("exclusive").AttrType(OPTIONAL).Bool(true); | ||
| 61 | - | ||
| 62 | - OpAICoreConfig aicoreConfig; | ||
| 63 | - aicoreConfig.DynamicCompileStaticFlag(true) | ||
| 64 | - .DynamicFormatFlag(false) | ||
| 65 | - .DynamicRankSupportFlag(true) | ||
| 66 | - .DynamicShapeSupportFlag(true) | ||
| 67 | - .NeedCheckSupportFlag(false) | ||
| 68 | - .PrecisionReduceFlag(true) | ||
| 69 | - .ExtendCfgInfo("opFile.value", "avg_pool_update"); | ||
| 70 | - this->AICore().AddConfig("ascend950", aicoreConfig); | ||
| 71 | - } | ||
| 72 | -}; | ||
| 73 | -OP_ADD(AvgPoolUpdate); | ||
| 74 | -} // namespace ops | ||
| @@ -1,42 +0,0 @@ | |||
| 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 | - * \file avg_pool_update_infershape.cpp | ||
| 13 | - * \brief Infershape implementation for avg_pool_update operator | ||
| 14 | - */ | ||
| 15 | - | ||
| 16 | - | ||
| 17 | - | ||
| 18 | - | ||
| 19 | -using namespace ge; | ||
| 20 | - | ||
| 21 | -namespace ops { | ||
| 22 | -static constexpr int64_t IDX_0 = 0; | ||
| 23 | - | ||
| 24 | -static ge::graphStatus InferShapeAvgPoolUpdate(gert::InferShapeContext* context) | ||
| 25 | -{ | ||
| 26 | - OP_LOGD(context->GetNodeName(), "Begin to do InferShapeAvgPoolUpdate"); | ||
| 27 | - | ||
| 28 | - const gert::Shape* x1Shape = context->GetInputShape(IDX_0); | ||
| 29 | - OP_CHECK_NULL_WITH_CONTEXT(context, x1Shape); | ||
| 30 | - | ||
| 31 | - gert::Shape* yShape = context->GetOutputShape(IDX_0); | ||
| 32 | - OP_CHECK_NULL_WITH_CONTEXT(context, yShape); | ||
| 33 | - | ||
| 34 | - // 输出 y shape 直接复制 x1 shape(正确处理动态 shape: rank 未知 -2 / dim 未知 -1) | ||
| 35 | - *yShape = *x1Shape; | ||
| 36 | - | ||
| 37 | - OP_LOGD(context->GetNodeName(), "End to do InferShapeAvgPoolUpdate"); | ||
| 38 | - return GRAPH_SUCCESS; | ||
| 39 | -} | ||
| 40 | - | ||
| 41 | -IMPL_OP_INFERSHAPE(AvgPoolUpdate).InferShape(InferShapeAvgPoolUpdate); | ||
| 42 | -} // namespace ops | ||