You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.08 05:45:37