MXNet导出ONNX报错pick算子无转换函数,如何实现对应转换逻辑
MXNet pick算子转ONNX实现方案
问题原因
MXNet的pick算子暂未被官方ONNX导出模块支持,该算子功能为沿指定轴抽取输入张量对应索引位置的元素,可直接映射为ONNX标准算子GatherElements(ONNX opset 11及以上版本支持)。
转换函数实现
你可以在python/mxnet/contrib/onnx/mx2onnx/_op_translations.py文件中添加如下代码:
@mx_op.register("pick") def convert_pick(node, **kwargs): from onnx import helper # 读取算子属性,默认axis=-1,keepdims=0 axis = int(node.attrs.get("axis", -1)) keepdims = int(node.attrs.get("keepdims", 0)) # 输入顺序:第0位是待取值张量,第1位是索引张量 data_input, index_input = node.inputs[0], node.inputs[1] # 构造GatherElements节点 gather_node = helper.make_node( op_type="GatherElements", inputs=[data_input, index_input], outputs=[f"{node.name}_gather_out"] if keepdims == 0 else node.outputs, axis=axis, name=node.name ) # keepdims为0时需追加Squeeze节点删除对应轴 if keepdims == 0: squeeze_node = helper.make_node( op_type="Squeeze", inputs=[f"{node.name}_gather_out"], outputs=node.outputs, axes=[axis], name=f"{node.name}_squeeze" ) return [gather_node, squeeze_node] return [gather_node]
参考依据
该实现完全对齐MXNet官方算子转换的编码规范,参考了同文件中gather、take等索引类算子的已实现转换逻辑,功能1:1匹配MXNet pick算子的行为。
注意事项
- 导出时需指定ONNX opset版本≥11,修改你的导出代码如下:
converted_model_path = onnx_mxnet.export_model(sym, params, [input_shape], np.float32, onnx_file, opset_version=11) - 验证转换正确性时,可生成随机输入分别运行MXNet原始模型和导出后的ONNX模型,确认两者输出误差小于1e-5即可。
内容的提问来源于stack exchange,提问作者ee44552
相关产品推荐
相关产品推荐

