You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.10.05 02:51:04