使用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
相关产品推荐
相关产品推荐

