PyTorch模型导出动态Batch Size的ONNX时出现段错误且未生成文件求助
我尝试将PyTorch模型导出为ONNX格式,同时确保批处理大小保持动态,但结果总是出现
Segmentation fault (core dumped),以下是我使用的代码:
import torch # Import PyTorch library # Create a dummy input tensor of shape (1, 3, 256, 256) and move it to the appropriate device dummy_input = torch.randn(1, 3, 256, 256).to(device) # Create dummy camera and view labels, initialized to zeros, and move them to the device dummy_cam = torch.zeros(1, dtype=torch.long).to(device) dummy_view = torch.zeros(1, dtype=torch.long).to(device) # Export the model to ONNX format torch.onnx.export( model, # The PyTorch model to be converted (dummy_input, dummy_cam, dummy_view), # The input tuple "deit_transreid_veri.onnx", # The output file name input_names=["input", "cam_label", "view_label"], # Naming the input tensors output_names=["output"], # Naming the output tensor dynamic_axes={ # Define dynamic batch size for inputs and outputs 'input': {0: 'batch_size'}, 'cam_label': {0: 'batch_size'}, 'view_label': {0: 'batch_size'}, 'output': {0: 'batch_size'} }, opset_version=13, # Specify the ONNX opset version verbose=True, # Print detailed information during export ) print("Model has been successfully exported to ONNX format!") # Confirmation message
遇到这种段错误的情况确实挺头疼的,我之前也碰到过类似的问题,咱们一步步来排查和解决:
先确认设备一致性:你这里用到了
device但没看到具体定义,一定要确保model和所有dummy输入(dummy_input、dummy_cam、dummy_view)完全在同一个设备上——比如如果device是cuda,必须确认模型已经通过model.to(device)移到GPU,不能出现模型在CPU但输入在GPU(或反过来)的混合设备情况,这种操作很容易触发导出崩溃。先关闭动态轴测试基础导出:暂时注释掉
dynamic_axes参数,尝试导出固定batch size的ONNX文件。如果能成功导出,说明问题大概率出在动态batch的配置或者模型对动态维度的支持上;如果还是崩溃,那问题可能在模型本身或者PyTorch/ONNX的版本兼容上。检查模型中的特殊操作:从模型名称看是DEIT+TransReID结构,这类模型如果包含PyTorch特有的、ONNX不支持的自定义操作或动态控制流(比如基于batch维度的条件判断、自定义注意力层实现),需要做适配——比如把动态控制流改成ONNX支持的形式,或者用
torch.onnx.symbolic_registry注册自定义操作的ONNX映射。调整dummy输入和导出参数:
- 试试把dummy输入的batch size改成2(比如
torch.randn(2, 3, 256, 256)),有时候batch size=1会触发特殊的张量形状问题; - 关闭
verbose=True改成verbose=False,大量日志输出可能导致内存异常; - 添加
do_constant_folding=False参数,禁用常量折叠,某些情况下常量折叠会和动态维度冲突导致崩溃。
- 试试把dummy输入的batch size改成2(比如
检查版本兼容性:确保PyTorch和ONNX版本匹配,opset_version=13对应的PyTorch版本建议在1.9.0及以上,太旧的版本对动态batch的ONNX导出支持不完善,可以尝试升级到稳定版本后再试。
尝试切换trace/script模式:默认
torch.onnx.export用trace模式,如果模型有动态控制流,trace模式会出问题,可以尝试用torch.jit.script(model)代替原来的model,用script模式导出试试。
备注:内容来源于stack exchange,提问作者mootaz haddad

