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

使用pytorch2keras转换PyTorch权重至Keras遇RuntimeError求助

解决PyTorch转Keras时的RuntimeError断言错误

这个Assertion 'var_state == state' failed错误,本质是PyTorch JIT追踪机制在处理模型或输入时遇到了状态不匹配的问题,咱们一步步拆解解决:

1. 先修正输入相关的核心问题

(1)更新输入变量的创建方式

在新版本PyTorch里,Variable已经被废弃,直接用torch.FloatTensor就行,不需要额外包装:

input_np = np.random.uniform(0, 1, (3, 12, 512, 512))  # 维度:batch, channels, H, W
input_tensor = torch.FloatTensor(input_np)

(2)修正input_shape参数

你传给pytorch_to_keras的input_shape错了——它需要的是单样本的输入形状(去掉batch维度),也就是(12, 512, 512),不是(3, 512, 512)。另外Keras默认是channels_last格式,建议加上change_ordering=True来自动转换维度顺序:

k_model = pytorch_to_keras(model, input_tensor, (12, 512, 512), 
                           change_ordering=True, verbose=True)

2. 强制模型进入评估模式

转换前必须把模型切到评估状态,避免Dropout、BatchNorm的训练行为干扰JIT追踪:

model.eval()  # 一定要加这行,再执行转换代码

3. 替换模型中JIT不友好的操作

你的forward里用了函数式的F.log_softmax,JIT追踪对层式操作的支持更稳定,把它改成nn.LogSoftmax层加入到分类器序列里:

修改ImageWiseNetwork的代码:

class ImageWiseNetwork(BaseNetwork):
    def __init__(self, channels=1):
        super(ImageWiseNetwork, self).__init__('iw' + str(channels), channels)
        self.features = nn.Sequential(
            # 原features层代码保持不变
        )
        self.classifier = nn.Sequential(
            nn.Linear(1 * 16 * 16, 128),
            nn.ReLU(inplace=True),
            nn.Dropout(0.5, inplace=True),
            nn.Linear(128, 128),
            nn.ReLU(inplace=True),
            nn.Dropout(0.5, inplace=True),
            nn.Linear(128, 64),
            nn.ReLU(inplace=True),
            nn.Dropout(0.5, inplace=True),
            nn.Linear(64, 4),
            nn.LogSoftmax(dim=1)  # 新增这一层替代函数式操作
        )
        self.initialize_weights()

    def forward(self, x):
        x = self.features(x)
        x = x.view(x.size(0), -1)
        x = self.classifier(x)
        # 移除原来的F.log_softmax(x, dim=1)
        return x

4. 验证权重加载的正确性

确保权重文件和模型结构完全匹配,加载后可以打印部分权重确认:

model = ImageWiseNetwork()
state_dict = torch.load('Path\to\weights\weights_iw1.pth')
model.load_state_dict(state_dict)
# 打印卷积层的小部分权重,确认加载成功
print(model.features[0].weight[:2, :2, :2, :2])

5. 用小尺寸输入测试(可选)

如果还是报错,可以先缩小输入尺寸(比如(1,12,64,64))来测试,排除大尺寸输入导致的内存或追踪异常:

input_np_small = np.random.uniform(0, 1, (1, 12, 64, 64))
input_tensor_small = torch.FloatTensor(input_np_small)
k_model = pytorch_to_keras(model, input_tensor_small, (12, 64, 64), 
                           change_ordering=True, verbose=True)

按上面的步骤调整后,应该能解决这个JIT追踪的断言错误。

内容的提问来源于stack exchange,提问作者Suvidha

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.29 08:00:40