已合并
refactor: 迁移 Resize ONNX 插件至 cv 仓并补充 Slice 算子桩定义 #1406
refactor: 迁移 Resize ONNX 插件至 cv 仓并补充 Slice 算子桩定义 #1406
已合并
chenfeng创建于 8月27日
共 2 个文件变更+394-0
@@ -0,0 +1,372 @@
1+/**
2+ * Copyright (c) 2025 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+#include "onnx_common.h"
12+#include "op_cv_proto_extend.h"
13+#include "resize_nearest_neighbor_v2_proto.h"
14+ 
15+using namespace ge;
16+namespace domi {
17+int INPUT_SIZES_IS_FOUR = 4;
18+int INPUT_SIZES_IS_THREE = 3;
19+int INPUT_SIZES_IS_TWO = 2;
20+ 
21+static ge::Operator CreateSliceForResize(const std::string& ori_name, ge::Operator& sizes)
22+{
23+ int32_t offsets = 2;
24+ int32_t size_num = 2;
25+ ge::Tensor scalar_offsets = CreateScalar(offsets, ge::DT_INT32);
26+ ge::Tensor scalar_size_num = CreateScalar(size_num, ge::DT_INT32);
27+ 
28+ auto data_offsets = op::Const((ori_name + "_data_offsets").c_str()).set_attr_value(scalar_offsets);
29+ auto data_size = op::Const((ori_name + "_data_size").c_str()).set_attr_value(scalar_size_num);
30+ 
31+ return op::Slice((ori_name + "_Slice").c_str())
32+ .set_input_x(sizes)
33+ .set_input_offsets(data_offsets)
34+ .set_input_size(data_size);
35+}
36+ 
37+static Status ParseParamsResize(const Message* op_src, Operator& op_dst)
38+{
39+ const ge::onnx::NodeProto* node = reinterpret_cast<const ge::onnx::NodeProto*>(op_src);
40+ if (node == nullptr) {
41+ OP_LOGE(GetOpName(op_dst).c_str(), "Dynamic cast op_src to NodeProto failed.");
42+ return FAILED;
43+ }
44+ 
45+ std::string coordinate_transformation_mode_value = "half_pixel";
46+ std::string mode_value = "nearest";
47+ for (auto attr : node->attribute()) {
48+ if (attr.name() == "coordinate_transformation_mode" && attr.type() == ge::onnx::AttributeProto::STRING) {
49+ coordinate_transformation_mode_value = attr.s();
50+ } else if (attr.name() == "mode" && attr.type() == ge::onnx::AttributeProto::STRING) {
51+ mode_value = attr.s();
52+ } else if (attr.name() == "antialias" && attr.type() == ge::onnx::AttributeProto::INT) {
53+ OP_LOGW(GetOpName(op_dst).c_str(), "Current antialias unsupported; expected type: int.");
54+ } else if (attr.name() == "axes" && attr.type() == ge::onnx::AttributeProto::INTS) {
55+ OP_LOGW(GetOpName(op_dst).c_str(), "Current axes unsupported; expected type: ints.");
56+ } else if (attr.name() == "keep_aspect_ratio_policy" && attr.type() == ge::onnx::AttributeProto::STRING) {
57+ OP_LOGW(GetOpName(op_dst).c_str(), "Current keep_aspect_ratio_policy unsupported; expected type: string.");
58+ }
59+ }
60+ 
61+ if (node->input_size() < INPUT_SIZES_IS_THREE) {
62+ OP_LOGE(GetOpName(op_dst).c_str(), "Input size is less than 3, cannot access input_roi and input_scales.");
63+ return FAILED;
64+ }
65+ 
66+ auto input_roi = node->input(1);
67+ op_dst.SetAttr("input_roi", input_roi);
68+ auto input_scales = node->input(INPUT_SIZES_IS_TWO);
69+ op_dst.SetAttr("input_scales", input_scales);
70+ op_dst.SetAttr("name", node->name());
71+ int input_size = node->input_size();
72+ if (input_size == INPUT_SIZES_IS_FOUR && node->input(INPUT_SIZES_IS_THREE).empty()) {
73+ input_size = INPUT_SIZES_IS_THREE;
74+ }
75+ op_dst.SetAttr("input_size", input_size);
76+ op_dst.SetAttr("coordinate_transformation_mode", coordinate_transformation_mode_value);
77+ op_dst.SetAttr("mode", mode_value);
78+ 
79+ op_dst.DynamicInputRegister("x", input_size);
80+ op_dst.DynamicOutputRegister("y", 1);
81+ op_dst.SetAttr("original_type", "ai.onnx::11::Resize");
82+ return SUCCESS;
83+}
84+ 
85+static Status BuildNearestResize(const std::string& ori_name, Operator& resize_x, Operator& sizes, Operator& resize_roi,
86+ Operator& resize_scales, const std::string& input_roi, const std::string& input_scales,
87+ int input_size, bool align_corners, bool half_pixel_centers,
88+ std::vector<Operator>& inputs,
89+ std::vector<std::pair<Operator, std::vector<size_t>>>& output_indexs)
90+{
91+ auto size = CreateSliceForResize(ori_name, sizes);
92+ inputs.push_back(size);
93+ auto ret_resize_x = ChangeFormatFromOnnx(resize_x, 0, ge::FORMAT_NCHW, false);
94+ if (ret_resize_x != ge::GRAPH_SUCCESS) {
95+ OP_LOGE(ori_name.c_str(), "update resize_x format failed.");
96+ return FAILED;
97+ }
98+ auto resizeout_1 = op::ResizeNearestNeighborV2((ori_name + "_ResizeNearestNeighborV2").c_str())
99+ .set_input_x(resize_x)
100+ .set_input_size(size)
101+ .set_attr_align_corners(align_corners)
102+ .set_attr_half_pixel_centers(half_pixel_centers);
103+ if (!input_roi.empty()) {
104+ resizeout_1.AddControlInput(resize_roi);
105+ }
106+ if ((input_size == INPUT_SIZES_IS_FOUR) && (!input_scales.empty())) {
107+ resizeout_1.AddControlInput(resize_scales);
108+ }
109+ ChangeFormatFromOnnx(resizeout_1, 0, ge::FORMAT_NCHW, false);
110+ ChangeFormatFromOnnx(resizeout_1, 0, ge::FORMAT_NCHW, true);
111+ output_indexs.emplace_back(resizeout_1, vector<std::size_t>{0});
112+ return SUCCESS;
113+}
114+ 
115+static Status BuildInterpolatingResize(const std::string& ori_name, Operator& resize_x, Operator& sizes,
116+ Operator& resize_roi, Operator& resize_scales, Operator& data1, Operator& data2,
117+ const std::string& input_roi, const std::string& input_scales,
118+ const std::string& coordinate_transformation_mode_value,
119+ const std::string& mode_value, std::vector<Operator>& inputs,
120+ std::vector<std::pair<Operator, std::vector<size_t>>>& output_indexs)
121+{
122+ auto resizeout_2 = op::Resize((ori_name + "_Resize").c_str())
123+ .set_input_x(resize_x)
124+ .set_input_sizes(sizes)
125+ .set_attr_coordinate_transformation_mode(coordinate_transformation_mode_value)
126+ .set_attr_mode(mode_value);
127+ inputs.push_back(sizes);
128+ std::vector<float> empty_vector = {};
129+ ge::Tensor empty_tensor = Vec2Tensor(empty_vector, {0}, ge::DT_FLOAT, ge::FORMAT_ND);
130+ if (!input_roi.empty()) {
131+ inputs.push_back(data1);
132+ resizeout_2.set_input_roi(resize_roi);
133+ } else {
134+ auto roi = op::Const((ori_name + "_roi").c_str()).set_attr_value(empty_tensor);
135+ inputs.push_back(roi);
136+ resizeout_2.set_input_roi(roi);
137+ }
138+ if (!input_scales.empty()) {
139+ inputs.push_back(data2);
140+ resizeout_2.set_input_scales(resize_scales);
141+ } else {
142+ auto scales = op::Const((ori_name + "_scales").c_str()).set_attr_value(empty_tensor);
143+ inputs.push_back(scales);
144+ resizeout_2.set_input_scales(scales);
145+ }
146+ resizeout_2.SetAttr("resize_original_type", "onnx_resize");
147+ output_indexs.emplace_back(resizeout_2, vector<std::size_t>{0});
148+ return SUCCESS;
149+}
150+ 
151+static Status ParseResizeModeAttrs(const Operator& op, std::string& coordinate_transformation_mode_value,
152+ bool& half_pixel_centers, bool& align_corners, std::string& mode_value)
153+{
154+ coordinate_transformation_mode_value = "pytorch_half_pixel";
155+ if (op.GetAttr("coordinate_transformation_mode", coordinate_transformation_mode_value) != GRAPH_SUCCESS) {
156+ OP_LOGW(GetOpName(op).c_str(), "Get attr coordinate transformation mode failed, set to default.");
157+ }
158+ half_pixel_centers = false;
159+ align_corners = false;
160+ if (coordinate_transformation_mode_value == "pytorch_half_pixel" ||
161+ coordinate_transformation_mode_value == "half_pixel") {
162+ half_pixel_centers = true;
163+ } else if (coordinate_transformation_mode_value == "align_corners") {
164+ align_corners = true;
165+ }
166+ 
167+ if (op.GetAttr("mode", mode_value) != GRAPH_SUCCESS) {
168+ OP_LOGE(GetOpName(op).c_str(), "Get attr mode failed, set to default.");
169+ return FAILED;
170+ }
171+ return SUCCESS;
172+}
173+ 
174+static Status BuildResizeSizesFromInputs(const std::string& ori_name, int input_size, ge::Operator& resize_x,
175+ ge::Operator& resize_scales, ge::Operator& resize_sizes, ge::Operator& sizes)
176+{
177+ if (input_size == INPUT_SIZES_IS_FOUR) {
178+ sizes = op::Cast((ori_name + "_Cast0").c_str()).set_input_x(resize_sizes).set_attr_dst_type(ge::DT_INT32);
179+ } else if (input_size == INPUT_SIZES_IS_THREE) {
180+ int dtype = 0; // DT_INT32
181+ int sizes_type = 3; // e.g., DT_INT64
182+ auto resize = op::Shape((ori_name + "_Shape").c_str()).set_input_x(resize_x);
183+ auto resize_cast = op::Cast((ori_name + "_Cast0").c_str()).set_input_x(resize).set_attr_dst_type(dtype);
184+ auto mul_sizes = op::Mul((ori_name + "_Mul").c_str()).set_input_x1(resize_cast).set_input_x2(resize_scales);
185+ sizes = op::Cast((ori_name + "_Cast1").c_str()).set_input_x(mul_sizes).set_attr_dst_type(sizes_type);
186+ } else {
187+ OP_LOGE("ParseOpToGraphResize", "The input_size is error.");
188+ return FAILED;
189+ }
190+ return SUCCESS;
191+}
192+ 
193+static Status ParseOpToGraphResize(const Operator& op, Graph& graph)
194+{
195+ std::string ori_name;
196+ if (op.GetAttr("name", ori_name) != SUCCESS) {
197+ OP_LOGE(GetOpName(op).c_str(), "get name from op failed.");
198+ return FAILED;
199+ }
200+ std::string input_roi, input_scales;
201+ op.GetAttr("input_roi", input_roi); // intentionally unchecked
202+ op.GetAttr("input_scales", input_scales);
203+ auto data0 = op::Data((ori_name + "_data0").c_str()).set_attr_index(0);
204+ auto resize_x = op::Identity((ori_name + "_x").c_str()).set_input_x(data0);
205+ auto data1 = op::Data((ori_name + "_data1").c_str()).set_attr_index(1);
206+ auto resize_roi = op::Identity((ori_name + "_roi").c_str()).set_input_x(data1);
207+ auto data2 = op::Data((ori_name + "_data2").c_str()).set_attr_index(2);
208+ auto resize_scales = op::Identity((ori_name + "_scales").c_str()).set_input_x(data2);
209+ auto data3 = op::Data((ori_name + "_data3").c_str()).set_attr_index(3);
210+ auto resize_sizes = op::Identity((ori_name + "_sizes").c_str()).set_input_x(data3);
211+ int input_size = 0;
212+ if (op.GetAttr("input_size", input_size) != SUCCESS) {
213+ OP_LOGE(GetOpName(op).c_str(), "get input_size from op failed");
214+ return FAILED;
215+ }
216+ std::string coordinate_transformation_mode_value;
217+ bool half_pixel_centers = false, align_corners = false;
218+ std::string mode_value;
219+ if (ParseResizeModeAttrs(op, coordinate_transformation_mode_value, half_pixel_centers, align_corners, mode_value) !=
220+ SUCCESS) {
221+ return FAILED;
222+ }
223+ ge::Operator sizes;
224+ if (BuildResizeSizesFromInputs(ori_name, input_size, resize_x, resize_scales, resize_sizes, sizes) != SUCCESS) {
225+ return FAILED;
226+ }
227+ std::vector<Operator> inputs{data0};
228+ std::vector<std::pair<Operator, std::vector<size_t>>> output_indexs;
229+ if (mode_value == "nearest") {
230+ auto status = BuildNearestResize(ori_name, resize_x, sizes, resize_roi, resize_scales, input_roi, input_scales,
231+ input_size, align_corners, half_pixel_centers, inputs, output_indexs);
232+ if (status != SUCCESS)
233+ return status;
234+ } else if (mode_value == "linear" || mode_value == "cubic") {
235+ auto status = BuildInterpolatingResize(ori_name, resize_x, sizes, resize_roi, resize_scales, data1, data2,
236+ input_roi, input_scales, coordinate_transformation_mode_value,
237+ mode_value, inputs, output_indexs);
238+ if (status != SUCCESS)
239+ return status;
240+ } else {
241+ OP_LOGE(GetOpName(op).c_str(), "Unsupported interpolation mode.");
242+ return FAILED;
243+ }
244+ graph.SetInputs(inputs).SetOutputs(output_indexs);
245+ return SUCCESS;
246+}
247+ 
248+static Status ParseParamsResizeV10(const Message* op_src, Operator& op_dst)
249+{
250+ const ge::onnx::NodeProto* node = reinterpret_cast<const ge::onnx::NodeProto*>(op_src);
251+ if (node == nullptr) {
252+ OP_LOGE(GetOpName(op_dst).c_str(), "Dynamic cast op_src to NodeProto failed.");
253+ return FAILED;
254+ }
255+ 
256+ std::string mode_value = "nearest";
257+ for (auto attr : node->attribute()) {
258+ if (attr.name() == "mode" && attr.type() == ge::onnx::AttributeProto::STRING) {
259+ mode_value = attr.s();
260+ }
261+ }
262+ op_dst.SetAttr("mode", mode_value);
263+ op_dst.SetAttr("name", node->name());
264+ op_dst.DynamicInputRegister("x", INPUT_SIZES_IS_TWO);
265+ op_dst.DynamicOutputRegister("y", 1);
266+ op_dst.SetAttr("original_type", "ai.onnx::10::Resize");
267+ return SUCCESS;
268+}
269+ 
270+static Status BuildNearestResizeV10(const std::string& ori_name, Operator& resize_x, Operator& sizes,
271+ std::vector<Operator>& inputs,
272+ std::vector<std::pair<Operator, std::vector<size_t>>>& output_indexs)
273+{
274+ auto size = CreateSliceForResize(ori_name, sizes);
275+ inputs.push_back(size);
276+ auto ret_resize_x = ChangeFormatFromOnnx(resize_x, 0, ge::FORMAT_NCHW, false);
277+ if (ret_resize_x != ge::GRAPH_SUCCESS) {
278+ OP_LOGE(ori_name.c_str(), "update resize_x format failed.");
279+ return FAILED;
280+ }
281+ bool half_pixel_centers = false;
282+ bool align_corners = false;
283+ auto resizeout_1 = op::ResizeNearestNeighborV2((ori_name + "_ResizeNearestNeighborV2").c_str())
284+ .set_input_x(resize_x)
285+ .set_input_size(size)
286+ .set_attr_align_corners(align_corners)
287+ .set_attr_half_pixel_centers(half_pixel_centers);
288+ ChangeFormatFromOnnx(resizeout_1, 0, ge::FORMAT_NCHW, false);
289+ ChangeFormatFromOnnx(resizeout_1, 0, ge::FORMAT_NCHW, true);
290+ output_indexs.emplace_back(resizeout_1, vector<std::size_t>{0});
291+ return SUCCESS;
292+}
293+ 
294+static Status BuildLinearResizeV10(const std::string& ori_name, Operator& resize_x, Operator& sizes,
295+ const std::string& mode_value, std::vector<Operator>& inputs,
296+ std::vector<std::pair<Operator, std::vector<size_t>>>& output_indexs)
297+{
298+ auto data2 = op::Const((ori_name + "_data2").c_str()).set_attr_value(0);
299+ auto data3 = op::Const((ori_name + "_data3").c_str()).set_attr_value(0);
300+ auto resizeout_2 = op::Resize((ori_name + "_Resize").c_str())
301+ .set_input_x(resize_x)
302+ .set_input_sizes(sizes)
303+ .set_input_roi(data2)
304+ .set_input_scales(data3)
305+ .set_attr_mode(mode_value);
306+ inputs.push_back(data2);
307+ inputs.push_back(data3);
308+ resizeout_2.SetAttr("resize_original_type", "onnx_resize");
309+ output_indexs.emplace_back(resizeout_2, vector<std::size_t>{0});
310+ return SUCCESS;
311+}
312+ 
313+static Status ParseOpToGraphResizeV10(const Operator& op, Graph& graph)
314+{
315+ std::string ori_name;
316+ if (op.GetAttr("name", ori_name) != SUCCESS) {
317+ OP_LOGE(GetOpName(op).c_str(), "get name from op failed.");
318+ return FAILED;
319+ }
320+ 
321+ auto data0 = op::Data((ori_name + "_data0").c_str()).set_attr_index(0);
322+ auto data1 = op::Data((ori_name + "_data1").c_str()).set_attr_index(1);
323+ auto resize_x = op::Identity((ori_name + "_Identity").c_str()).set_input_x(data0);
324+ 
325+ int dtype = 0;
326+ int sizes_type = 3;
327+ auto resize = op::Shape((ori_name + "_Shape").c_str()).set_input_x(resize_x);
328+ auto resize_cast = op::Cast((ori_name + "_Cast0").c_str()).set_input_x(resize).set_attr_dst_type(dtype);
329+ auto mul_sizes = op::Mul((ori_name + "_Mul").c_str()).set_input_x1(resize_cast).set_input_x2(data1);
330+ auto sizes = op::Cast((ori_name + "_Cast1").c_str()).set_input_x(mul_sizes).set_attr_dst_type(sizes_type);
331+ 
332+ std::string mode_value;
333+ if (op.GetAttr("mode", mode_value) != GRAPH_SUCCESS) {
334+ OP_LOGE(GetOpName(op).c_str(), "Get attr mode failed, set to default.");
335+ return FAILED;
336+ }
337+ 
338+ std::vector<Operator> inputs{data0, data1};
339+ std::vector<std::pair<Operator, std::vector<size_t>>> output_indexs;
340+ if (mode_value == "nearest") {
341+ auto status = BuildNearestResizeV10(ori_name, resize_x, sizes, inputs, output_indexs);
342+ if (status != SUCCESS)
343+ return status;
344+ } else if (mode_value == "linear") {
345+ auto status = BuildLinearResizeV10(ori_name, resize_x, sizes, mode_value, inputs, output_indexs);
346+ if (status != SUCCESS)
347+ return status;
348+ } else {
349+ OP_LOGE(GetOpName(op).c_str(), "Unsupported interpolation mode.");
350+ return FAILED;
351+ }
352+ graph.SetInputs(inputs).SetOutputs(output_indexs);
353+ return SUCCESS;
354+}
355+ 
356+REGISTER_CUSTOM_OP("PartitionedCall")
357+ .FrameworkType(ONNX)
358+ .OriginOpType({ge::AscendString("ai.onnx::11::Resize"), ge::AscendString("ai.onnx::12::Resize"),
359+ ge::AscendString("ai.onnx::13::Resize"), ge::AscendString("ai.onnx::14::Resize"),
360+ ge::AscendString("ai.onnx::15::Resize"), ge::AscendString("ai.onnx::16::Resize"),
361+ ge::AscendString("ai.onnx::17::Resize"), ge::AscendString("ai.onnx::18::Resize")})
362+ .ParseParamsFn(ParseParamsResize)
363+ .ParseOpToGraphFn(ParseOpToGraphResize)
364+ .ImplyType(ImplyType::TVM);
365+ 
366+REGISTER_CUSTOM_OP("PartitionedCall")
367+ .FrameworkType(ONNX)
368+ .OriginOpType({ge::AscendString("ai.onnx::10::Resize")})
369+ .ParseParamsFn(ParseParamsResizeV10)
370+ .ParseOpToGraphFn(ParseOpToGraphResizeV10)
371+ .ImplyType(ImplyType::TVM);
372+} // namespace domi
@@ -282,6 +282,28 @@ REG_OP(Empty)
282 .ATTR(dtype, Int, DT_INT32)282 .ATTR(dtype, Int, DT_INT32)
283 .ATTR(init, Bool, false)283 .ATTR(init, Bool, false)
284 .OP_END_FACTORY_REG(Empty)284 .OP_END_FACTORY_REG(Empty)
285+ 
286+/**
287+*@brief Extracts a slice from a tensor. \n
288+ 
289+*@par Inputs:
290+*Three inputs, including:
291+*@li x: A tensor. \n
292+*@li offsets: The starting location for the slice. Must be one of the following types: int32, int64.
293+*@li size: The tensor shape. Must be one of the following types: int32, int64. \n
294+ 
295+*@par Outputs:
296+*y: A tensor with the same type as "x". \n
297+ 
298+*@par Third-party framework compatibility
299+*Compatible with the TensorFlow operator Slice.
300+*/
301+REG_OP(Slice)
302+ .INPUT(x, TensorType({BasicType(), DT_HIFLOAT8, DT_FLOAT8_E5M2, DT_FLOAT8_E4M3FN}))
303+ .INPUT(offsets, TensorType::IndexNumberType())
304+ .INPUT(size, TensorType::IndexNumberType())
305+ .OUTPUT(y, TensorType({BasicType(), DT_HIFLOAT8, DT_FLOAT8_E5M2, DT_FLOAT8_E4M3FN}))
306+ .OP_END_FACTORY_REG(Slice)
285} // namespace ge307} // namespace ge
286 308 
287#endif // CV_COMMON_STUB_OPS_H309#endif // CV_COMMON_STUB_OPS_H