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

如何导出含变长张量循环的PyTorch模型至ONNX?

问题分析

你遇到的错误根源在于:代码中for point in points的显式循环在ONNX导出时会被固化为基于导出时输入尺寸的固定Split操作。当你修改输入batch size时,points.reshape(-1,2)的形状从导出时的(8,2)变为(16,2),但ONNX模型中的Split节点仍使用原固定参数(8个输出,每个尺寸1),导致拆分后尺寸总和(8)与当前轴实际尺寸(16)不匹配,触发报错。

解决方案(无需移出核心操作)

将显式循环替换为PyTorch向量化张量操作,让ONNX能正确追踪动态维度,避免生成固定参数的Split节点。修改后的模型代码如下:

import torch
from torch import nn
import onnx
import onnxruntime
import numpy as np


class Model(nn.Module):
    def __init__(self):
        super(Model, self).__init__()
        self.template = torch.randn((1000, 1000))
        # 预先定义切片尺寸,方便批量处理
        self.heatmap_h = 10
        self.heatmap_w = 20

    def forward(self, points):
        batch_size, num_points, _ = points.shape
        # 提取所有坐标并展平,形状为(batch*num_points,)
        x_coords = points[..., 0].flatten()
        y_coords = points[..., 1].flatten()
        
        # 生成批量切片的起止索引
        x_start = x_coords
        x_end = x_start + self.heatmap_h
        y_start = y_coords
        y_end = y_start + self.heatmap_w
        
        # 构建批量索引网格,实现一次性切片所有区域
        x_indices = torch.arange(self.heatmap_h, device=points.device)[None, :] + x_start[:, None]
        y_indices = torch.arange(self.heatmap_w, device=points.device)[None, :] + y_start[:, None]
        
        # 利用高级索引获取所有热力图,形状为(batch*num_points, heatmap_h, heatmap_w)
        heatmaps = self.template[x_indices[:, :, None], y_indices[:, None, :]]
        
        # 恢复为与输入batch对齐的形状:(batch_size, num_points, heatmap_h, heatmap_w)
        heatmaps = heatmaps.reshape(batch_size, num_points, self.heatmap_h, self.heatmap_w)
        return heatmaps


model = Model()
points = torch.randint(100, 200, (1, 8, 2))

# 导出ONNX,动态轴配置保持不变
torch.onnx.export(model, args=points, f='toy.onnx',
                  export_params=True,
                  opset_version=13,
                  do_constant_folding=True,
                  verbose=False,
                  input_names=['input1'],
                  output_names=['output1'],
                  dynamic_axes={'input1': {0: 'batch_size'},
                                'output1': {0: 'batch_size'},
                                }
                  )

# 测试不同batch size的推理
session = onnxruntime.InferenceSession("./toy.onnx")
# 测试batch size=2的输入
inputs = np.random.randint(100, 200, (2, 8, 2))
ort_inputs = {'input1': inputs.astype(np.int64)}  # 注意类型匹配,PyTorch导出的是int64
ort_outs = session.run(None, ort_inputs)
print(f"输出形状:{ort_outs[0].shape}")  # 应输出(2,8,10,20)
关键说明
  1. 向量化替代循环:通过构建批量索引网格,一次性完成所有坐标的切片操作,避免了Python循环遍历张量元素的行为,让ONNX能正确识别动态维度的变化。
  2. 维度对齐:最终输出形状与输入batch维度对齐,既保留了原逻辑的功能,又符合ONNX对动态张量的处理要求。
  3. 类型匹配:推理时需确保输入数据类型与导出时一致(这里是int64),避免额外报错。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.25 05:23:12