PyTorch DeepLabV3 Mobile版本导出ONNX报错求助
DeepLabV3-Mobile版导出ONNX报错解决
问题场景
基于PyTorch官方DeepLabV3实现微调多个模型后,其他版本均可正常导出ONNX,但Mobile版本导出时触发如下错误:
E:\Documents\Florian\Programmation\QuickTestModel\venv\lib\site-packages\torchvision\io\image.py:11: UserWarning: Failed to load image Python extension: Could not find module 'E:\Documents\Florian\Programmation\QuickTestModel\venv\Lib\site-packages\torchvision\image.pyd' (or one of its dependencies). Try using the full path with constructor syntax. warn(f"Failed to load image Python extension: {e}") Traceback (most recent call last): File "E:\Documents\Florian\Programmation\QuickTestModel\export_onnx.py", line 31, in <module> torch_out = model(input_batch) File "E:\Documents\Florian\Programmation\QuickTestModel\venv\lib\site-packages\torch\nn\modules\module.py", line 1102, in _call_impl return forward_call(*input, **kwargs) File "E:\Documents\Florian\Programmation\QuickTestModel\venv\lib\site-packages\torchvision\models\segmentation\_utils.py", line 25, in forward features = self.backbone(x) File "E:\Documents\Florian\Programmation\QuickTestModel\venv\lib\site-packages\torch\nn\modules\module.py", line 1102, in _call_impl return forward_call(*input, **kwargs) File "E:\Documents\Florian\Programmation\QuickTestModel\venv\lib\site-packages\torchvision\models\_utils.py", line 62, in forward x = module(x) File "E:\Documents\Florian\Programmation\QuickTestModel\venv\lib\site-packages\torch\nn\modules\module.py", line 1102, in _call_impl return forward_call(*input, **kwargs) File "E:\Documents\Florian\Programmation\QuickTestModel\venv\lib\site-packages\torchvision\models\mobilenetv3.py", line 89, in forward result = self.block(input) File "E:\Documents\Florian\Programmation\QuickTestModel\venv\lib\site-packages\torch\nn\modules\module.py", line 1102, in _call_impl return forward_call(*input, **kwargs) File "E:\Documents\Florian\Programmation\QuickTestModel\venv\lib\site-packages\torch\nn\modules\container.py", line 141, in forward input = module(input) File "E:\Documents\Florian\Programmation\QuickTestModel\venv\lib\site-packages\torch\nn\modules\module.py", line 1102, in _call_impl return forward_call(*input, **kwargs) File "E:\Documents\Florian\Programmation\QuickTestModel\venv\lib\site-packages\torchvision\ops\misc.py", line 151, in forward scale = self._scale(input) File "E:\Documents\Florian\Programmation\QuickTestModel\venv\lib\site-packages\torchvision\ops\misc.py", line 144, in _scale scale = self.avgpool(input) File "E:\Documents\Florian\Programmation\QuickTestModel\venv\lib\site-packages\torch\nn\modules\module.py", line 1177, in __getattr__ raise AttributeError("'{}' object has no attribute '{}'".format( AttributeError: 'SqueezeExcitation' object has no attribute 'avgpool'
导出代码如下:
import torch import torch.onnx if __name__ == '__main__': model = torch.load(r"result/deepllabv3_trace_hum_v2_mobile/weights.pt") model = model.to("cpu") x = torch.randn(3, 960, 540) input_batch = x.unsqueeze(0) with torch.no_grad(): torch_out = model(input_batch) torch.onnx.export(model, input_batch, f="model.onnx", export_params=True, opset_version=14, do_constant_folding=True, input_names=['input'], output_names=['output'], dynamic_axes={'input' : {0 : 'batch_size'}, # variable length axes 'output' : {0 : 'batch_size'}})
解决方案
1. 统一torchvision版本
这个错误本质是训练环境和导出环境的torchvision版本不兼容:旧版本SqueezeExcitation模块使用avgpool属性,新版本已将其改名为avg_pool或直接内嵌自适应池化逻辑。
- 操作:确保导出环境的torch、torchvision版本和训练时完全一致;如果是用新版本训练的,建议将导出环境的torchvision升级到
>=0.13.0。
2. 手动修复模型的SE模块
如果无法更换环境,可以通过代码替换模型中的SqueezeExcitation模块,兼容新旧版本的属性差异:
import torch import torch.nn as nn from torchvision.ops.misc import SqueezeExcitation as OriginalSE # 重写SE模块,适配属性名差异 class FixedSqueezeExcitation(nn.Module): def __init__(self, se_module): super().__init__() self.original_se = se_module # 自动匹配avgpool/avg_pool属性,或手动创建自适应池化 if hasattr(se_module, 'avgpool'): self.avgpool = se_module.avgpool elif hasattr(se_module, 'avg_pool'): self.avgpool = se_module.avg_pool else: self.avgpool = nn.AdaptiveAvgPool2d(1) # 复用原模块的其他组件 self.fc1 = se_module.fc1 self.fc2 = se_module.fc2 self.activation = se_module.activation self.sigmoid = se_module.sigmoid def forward(self, x): scale = self.avgpool(x) scale = self.fc1(scale) scale = self.activation(scale) scale = self.fc2(scale) scale = self.sigmoid(scale) return x * scale # 递归遍历模型,替换所有SE模块 def fix_se_modules(model): for name, module in model.named_children(): if isinstance(module, OriginalSE): setattr(model, name, FixedSqueezeExcitation(module)) else: fix_se_modules(module) return model # 加载模型后先执行修复 model = torch.load(r"result/deepllabv3_trace_hum_v2_mobile/weights.pt") model = model.to("cpu") model = fix_se_modules(model) model.eval() # 务必设置为评估模式
3. 先转TorchScript再导出ONNX
动态图模型直接导出可能存在属性绑定问题,先将模型转为TorchScript再导出,能规避此类动态属性错误:
# 加载模型并设置为评估模式 model = torch.load(r"result/deepllabv3_trace_hum_v2_mobile/weights.pt") model = model.to("cpu").eval() x = torch.randn(3, 960, 540) input_batch = x.unsqueeze(0) # 先trace模型为TorchScript with torch.no_grad(): traced_model = torch.jit.trace(model, input_batch) # 导出traced模型到ONNX torch.onnx.export(traced_model, input_batch, f="model.onnx", export_params=True, opset_version=14, do_constant_folding=True, input_names=['input'], output_names=['output'], dynamic_axes={'input' : {0 : 'batch_size'}, 'output' : {0 : 'batch_size'}})
内容的提问来源于stack exchange,提问作者FlorianDufour
相关产品推荐
相关产品推荐

