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

PyTorch转ONNX遇scatter_max问题,推理时维度不匹配报错

PyTorch转ONNX时scatter_max算子的问题排查与解决

问题描述

  1. 初始转换错误:

raise errors.UnsupportedOperatorError(
torch.onnx.errors.UnsupportedOperatorError: ONNX export failed on an operator with unrecognized namespace 'torch_scatter::scatter_max'. If you are trying to export a custom operator, make sure you registered it with the right domain and version.

  1. 尝试注册算子的代码:
from typing import Optional, Tuple
import torch

def scatter_max(
        src: torch.Tensor, index: torch.Tensor, dim: int = -1,
        out: Optional[torch.Tensor] = None,
        dim_size: Optional[int] = None,fill_value=0) -> Tuple[torch.Tensor, torch.Tensor]:
    return torch.ops.torch_scatter.scatter_max(src, index, dim, out, dim_size)

# Register custom symbolic function
torch.onnx.register_custom_op_symbolic("torch_scatter::scatter_max", scatter_max,9)
  1. 生成ONNX后推理错误:

onnxruntime.capi.onnxruntime_pybind11_state.InvalidArgument: [ONNXRuntimeError] : 2 : INVALID_ARGUMENT : Non-zero status code returned while running ScatterElements node. Name:'/hsn/ScatterElements_1' Status Message: Indices vs updates dimensions differs at position=1 1 vs 64

问题原因与解决方法

你的错误出在自定义算子注册的逻辑完全错误——直接返回原torch_scatter的scatter_max,根本没帮ONNX理解这个算子的计算逻辑,ONNX只能用默认的ScatterElements节点去强行适配,但两者的维度要求不匹配,最终导致推理时维度冲突。

正确的做法是:给torch_scatter::scatter_max编写ONNX符号函数,用ONNX原生支持的算子组合出scatter_max的逻辑,明确告诉ONNX这个算子该怎么转换,而不是直接调用PyTorch算子。

正确的注册示例

针对scatter_max(返回最大值张量+对应索引张量),可以这样实现符号函数:

from typing import Optional, Tuple
import torch
from torch.onnx import symbolic_helper

def symbolic_scatter_max(g, src, index, dim=-1, out=None, dim_size=None, fill_value=0):
    # 处理负维度索引,转为正索引
    dim_val = symbolic_helper._get_const(dim, "i")
    if dim_val < 0:
        dim = g.op("Constant", value_t=torch.tensor(src.dim() + dim_val, dtype=torch.int64))
    else:
        dim = g.op("Constant", value_t=torch.tensor(dim_val, dtype=torch.int64))
    
    # 自动计算dim_size(如果未指定)
    if dim_size is None:
        max_idx = g.op("ReduceMax", index, keepdims=False)
        dim_size = g.op("Add", max_idx, g.op("Constant", value_t=torch.tensor(1, dtype=torch.int64)))
    
    # 创建初始填充张量,用fill_value初始化
    src_shape = g.op("Shape", src)
    new_shape = []
    for i in range(src.dim()):
        if i == dim_val:
            new_shape.append(dim_size)
        else:
            new_shape.append(g.op("Gather", src_shape, g.op("Constant", value_t=torch.tensor(i, dtype=torch.int64))))
    init_out = g.op("Constant", value_t=torch.full(tuple(symbolic_helper._get_const(s, "i") for s in new_shape), fill_value, dtype=src.dtype))
    
    # 用ScatterElements的max归约实现最大值聚合
    scattered_max = g.op("ScatterElements", init_out, index, src, axis_i=dim_val, reduction_s="max")
    
    # 处理argmax部分:记录最大值对应的索引
    init_argmax = g.op("Constant", value_t=torch.full(tuple(symbolic_helper._get_const(s, "i") for s in new_shape), -1, dtype=index.dtype))
    scattered_argmax = g.op("ScatterElements", init_argmax, index, index, axis_i=dim_val, reduction_s="max")
    
    return scattered_max, scattered_argmax

# 注册自定义符号函数
torch.onnx.register_custom_op_symbolic("torch_scatter::scatter_max", symbolic_scatter_max, 9)

关键说明

  • 符号函数必须用ONNX原生算子(如ScatterElements、ReduceMax等)构建逻辑,不能直接调用PyTorch的算子接口。
  • 必须手动处理负维度、默认dim_size这些参数,确保所有张量的维度完全对齐。
  • scatter_max返回两个张量,符号函数需要同时处理这两个输出的ONNX转换逻辑。

另外,建议使用PyTorch 1.12+和ONNX 1.13+版本,避免版本兼容问题。

内容的提问来源于stack exchange,提问作者FatemeZamanian

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 07:35:07