PyTorch转ONNX遇scatter_max问题,推理时维度不匹配报错
问题描述
- 初始转换错误:
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.
- 尝试注册算子的代码:
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)
- 生成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

