使用EfficientNet backbone时,PyTorch JIT Trace存模型遇SwishImplementation错误
解决Segmentation Models PyTorch中EfficientNet backbone的JIT追踪保存问题
错误原因
EfficientNet backbone内置的Swish激活使用了纯Python实现的SwishImplementation类,而TorchScript追踪模式无法处理未被脚本化的原生Python函数调用,因此触发导出失败。
解决方案
方法一:替换为PyTorch原生SiLU激活
PyTorch官方提供了torch.nn.SiLU(即Swish激活的原生实现),直接替换模型中所有的SwishImplementation实例即可:
import torch import segmentation_models_pytorch as smp # 递归替换模型中的SwishImplementation为原生SiLU def replace_swish(module): for name, child in module.named_children(): if isinstance(child, smp.encoders.efficientnet.SwishImplementation): setattr(module, name, torch.nn.SiLU(inplace=True)) else: replace_swish(child) # 初始化模型并替换激活 model = smp.Unet('efficientnet-b7') replace_swish(model) model.eval() # 准备输入并完成追踪保存 input_tensor = torch.randn((1, 3, 224, 224)) with torch.no_grad(): traced_model = torch.jit.trace(model, input_tensor) traced_model.save('model.pt') # 验证模型可正常加载使用 loaded_model = torch.jit.load('model.pt') output = loaded_model(input_tensor) print(f"输出形状:{output.shape}")
方法二:脚本化Swish实现(猴子补丁)
通过猴子补丁将原SwishImplementation替换为带@torch.jit.script装饰的版本,让TorchScript能识别:
import torch import segmentation_models_pytorch as smp from segmentation_models_pytorch.encoders.efficientnet import SwishImplementation # 定义脚本化的Swish类 @torch.jit.script class ScriptedSwish(torch.nn.Module): def __init__(self, inplace: bool = False): super().__init__() self.inplace = inplace def forward(self, x: torch.Tensor) -> torch.Tensor: return x * torch.sigmoid(x) # 猴子补丁替换原实现 smp.encoders.efficientnet.SwishImplementation = ScriptedSwish # 后续流程与原代码一致 model = smp.Unet('efficientnet-b7') model.eval() input_tensor = torch.randn((1, 3, 224, 224)) with torch.no_grad(): traced_model = torch.jit.trace(model, input_tensor) traced_model.save('model.pt')
注意事项
- 替换激活后,建议用
torch.allclose对比原模型与替换后模型的输出,确保精度一致 - 替换操作必须在模型初始化后、JIT追踪前执行
- 追踪过程包裹在
torch.no_grad()中,避免不必要的梯度计算,提升效率
内容的提问来源于stack exchange,提问作者Julia.T
相关产品推荐
相关产品推荐

