已合并
Revert AvgPoolUpdate、NPUClearFloatStatus算子950支持 #10779
吴成文创建于 14 天前
Revert AvgPoolUpdate、NPUClearFloatStatus算子950支持 #10779
已合并
吴成文创建于 14 天前
共 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+#ifndef OPS_PROTO_DEF_NPUCLEARFLOATSTATUS
1726+#define OPS_PROTO_DEF_NPUCLEARFLOATSTATUS
1727+ REG_OP(NPUClearFloatStatus)
1728+ .INPUT(addr, TensorType({DT_FLOAT}))
1729+ .OUTPUT(data, TensorType({DT_FLOAT}))
1730+ .OP_END_FACTORY_REG(NPUClearFloatStatus)
1731+#endif
1732+ 
1725/**1733/**
1726* @brief Gather slices from "params" according to "indices"."indices" must be1734* @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+#ifndef OPS_PROTO_DEF_AVGPOOLUPDATE
2805+#define OPS_PROTO_DEF_AVGPOOLUPDATE
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+#endif
2841+ 
2796 /**2842 /**
2797 *@brief Finds unique elements in a 1D tensor. \n2843 *@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-#include <cstddef>
15-#include <cstdint>
16-#include <iostream>
17-#include <cstring>
18-#include <vector>
19-#include <string>
20-#include <map>
21-#include <memory>
22- 
23-#include "graph.h"
24-#include "types.h"
25-#include "tensor.h"
26-#include "ge_error_codes.h"
27-#include "ge_api_types.h"
28-#include "ge_api.h"
29-#include "array_ops.h"
30-#include "ge_ir_build.h"
31- 
32-#include "experiment_ops.h"
33-#include "nn_other.h"
34-#include "../op_graph/npu_clear_float_status_proto.h"
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-#define ADD_INPUT_MODE(inputIndex, inputName, inputDtype, inputShape, mode) \
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-#define ADD_CONST_INPUT(inputIndex, inputName, inputDtype, inputShape) \
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-#define ADD_OUTPUT_MODE(outputIndex, outputName, outputDtype, outputShape, mode) \
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-#include "register/op_impl_registry.h"
18-#include "log/log.h"
19-#include "../op_host/npu_clear_float_status_common.h"
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-#ifndef OPS_OP_PROTO_INC_NPUCLEARFLOATSTATUS_H_
17-#define OPS_OP_PROTO_INC_NPUCLEARFLOATSTATUS_H_
18- 
19-#include "graph/operator_reg.h"
20-#include "graph/types.h"
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-#ifndef OPS_PROTO_DEF_NPUCLEARFLOATSTATUS
34-#define OPS_PROTO_DEF_NPUCLEARFLOATSTATUS
35-REG_OP(NPUClearFloatStatus)
36- .INPUT(addr, TensorType({DT_FLOAT}))
37- .OUTPUT(data, TensorType({DT_FLOAT}))
38- .OP_END_FACTORY_REG(NPUClearFloatStatus)
39-#endif
40- 
41-} // namespace ge
42- 
43-#endif // OPS_OP_PROTO_INC_NPUCLEARFLOATSTATUS_H_
@@ -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-#include "graph/types.h"
14-#include "graph/utils/type_utils.h"
15-#include "log/log.h"
16-#include "platform/platform_ascendc.h"
17-#include "register/op_impl_registry.h"
18-#include "../../op_kernel/arch35/npu_clear_float_status_tiling_data.h"
19-#include "../../op_kernel/arch35/npu_clear_float_status_tiling_key.h"
20-#include "../npu_clear_float_status_common.h"
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-#ifndef NPU_CLEAR_FLOAT_STATUS_COMMON_H_
14-#define NPU_CLEAR_FLOAT_STATUS_COMMON_H_
15- 
16-#include <cstddef>
17-#include <cstdint>
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-#endif // NPU_CLEAR_FLOAT_STATUS_COMMON_H_
@@ -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-#include "register/op_def_registry.h"
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-#include "graph/types.h"
25-#include "graph/utils/type_utils.h"
26-#include "register/op_impl_registry.h"
27-#include "log/log.h"
28-#include "npu_clear_float_status_common.h"
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-#ifndef NPU_CLEAR_FLOAT_STATUS_SIMT_H_
14-#define NPU_CLEAR_FLOAT_STATUS_SIMT_H_
15- 
16-#include "kernel_operator.h"
17-#include "kernel_tiling/kernel_tiling.h"
18-#include "simt_api/common_functions.h"
19-#include "simt_api/asc_simt.h"
20-#include "npu_clear_float_status_tiling_data.h"
21-#include "npu_clear_float_status_tiling_key.h"
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-#endif // NPU_CLEAR_FLOAT_STATUS_SIMT_H_
@@ -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-#ifndef NPU_CLEAR_FLOAT_STATUS_TILING_DATA_H_
14-#define NPU_CLEAR_FLOAT_STATUS_TILING_DATA_H_
15- 
16-#include <cstdint>
17- 
18-struct NPUClearFloatStatusTilingData {
19- int32_t needCoreNum = 0; // 需要启动的核数(= 物理核数)
20-};
21- 
22-#endif
@@ -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-#ifndef NPU_CLEAR_FLOAT_STATUS_TILING_KEY_H_
14-#define NPU_CLEAR_FLOAT_STATUS_TILING_KEY_H_
15- 
16-#include "ascendc/host_api/tiling/template_argument.h"
17- 
18-#define NPU_CLEAR_FLOAT_STATUS_SCH_MODE_DEFAULT 0
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-#endif
@@ -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-#include "arch35/npu_clear_float_status_simt.h"
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- @staticmethod
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-#include <iostream>
14-#include <map>
15-#include <string>
16-#include <vector>
17-#include <gtest/gtest.h>
18-#include "kernel_run_context_facker.h"
19-#include "test_cube_util.h"
20-#include "platform/platform_infos_def.h"
21-#include "register/op_impl_registry.h"
22-#include "exe_graph/runtime/storage_format.h"
23-#include "exe_graph/runtime/storage_shape.h"
24-#include "ut_op_util.h"
25-#include "../../../../op_kernel/arch35/npu_clear_float_status_tiling_data.h"
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-#include <gtest/gtest.h>
14-#include <iostream>
15-#include <vector>
16-#include "kernel_run_context_facker.h"
17-#include "register/op_impl_registry.h"
18-#include "exe_graph/runtime/storage_shape.h"
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-#include <cstddef>
15-#include <cstdint>
16-#include <iostream>
17-#include <fstream>
18-#include <vector>
19-#include <string>
20-#include <map>
21- 
22-#include "graph.h"
23-#include "types.h"
24-#include "tensor.h"
25-#include "ge_error_codes.h"
26-#include "ge_api_types.h"
27-#include "ge_api.h"
28-#include "ops_proto_legacy.h"
29-#include "ge_ir_build.h"
30- 
31-#include "../../op_graph/avg_pool_update_proto.h"
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-#define ADD_INPUT_MODE(inputIndex, inputName, intputDtype, inputShape, mode) \
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-#define ADD_OUTPUT_MODE(outputIndex, outputName, outputDtype, outputShape, mode) \
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-#include "register/op_impl_registry.h"
19-#include "log/log.h"
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-#ifndef OPS_POOLING_AVG_POOL_UPDATE_PROTO_H_
18-#define OPS_POOLING_AVG_POOL_UPDATE_PROTO_H_
19- 
20-#include "graph/operator_reg.h"
21-#include "graph/types.h"
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-#endif // OPS_POOLING_AVG_POOL_UPDATE_PROTO_H_
@@ -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-#include "log/log.h"
17-#include "platform/platform_ascendc.h"
18-#include "register/op_impl_registry.h"
19-#include "op_host/tiling_templates_registry.h"
20-#include "../../op_kernel/arch35/avg_pool_update_tiling_data.h"
21-#include "../../op_kernel/arch35/avg_pool_update_tiling_key.h"
22- 
23-#include <cstdint>
24-#include <cstring>
25-#include <string>
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-#include "register/op_def_registry.h"
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-#include "register/op_impl_registry.h"
17-#include "log/log.h"
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