如何在PyTorch转ONNX模型时实现输出的动态批量大小?
PyTorch转ONNX输出动态批量大小的解决方法
核心问题排查与解决步骤
修正输入样本
input_sample:你当前传入的(x, None)可能导致ONNX无法正确跟踪输出维度关联。如果模型仅接受image作为输入,直接传递符合预期维度的张量x即可(比如x = torch.randn(2, 3, 640, 640),这里的2是测试用的批量大小),不要加None。ONNX需要基于有效的输入样本来推导输出维度的动态性。统一动态轴的关联命名:修改
dynamic_axes的定义,将输入和输出的第0轴绑定到同一个动态维度名称(比如batch_size),让ONNX明确它们的维度联动关系:dynamic_axes={ "image": {0: "batch_size"}, "boxes": {0: "batch_size"}, "scores": {0: "batch_size"}, "labels": {0: "batch_size"}, }检查模型forward函数的硬编码操作:如果模型内部存在固定批量大小的操作(比如
torch.reshape(out, (1, -1, 4))这种硬写1的情况),必须修改为动态获取批量维度。比如用input_image.shape[0]来获取当前输入的批量大小,替换所有硬编码的数值,确保输出的批量维度始终跟随输入变化。临时禁用常量折叠:尝试将
do_constant_folding=False,部分场景下常量折叠会把动态维度相关的操作优化为固定值,禁用后可能恢复动态性。验证模型输出的动态性:导出前先用不同批量大小的输入(比如batch=2和batch=4)运行模型forward,确认输出
boxes、scores、labels的第0轴确实对应输入的批量大小。如果输出维度固定,说明模型内部存在维度固化逻辑,必须修改forward函数。
内容的提问来源于stack exchange,提问作者brad
相关产品推荐
相关产品推荐

