如何导出含变长张量循环的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)
关键说明
- 向量化替代循环:通过构建批量索引网格,一次性完成所有坐标的切片操作,避免了Python循环遍历张量元素的行为,让ONNX能正确识别动态维度的变化。
- 维度对齐:最终输出形状与输入batch维度对齐,既保留了原逻辑的功能,又符合ONNX对动态张量的处理要求。
- 类型匹配:推理时需确保输入数据类型与导出时一致(这里是
int64),避免额外报错。
内容的提问来源于stack exchange,提问作者sunny
相关产品推荐
相关产品推荐

