PyTorch JIT脚本化自定义模型失败,求助解决方法
解决PyTorch JIT脚本化MaxUnpool1D相关错误
问题分析
调用torch.jit.script(modelFull())时出现两类错误:
RuntimeError: Can't redefine method: forward on class: __torch__.torch.nn.modules.pooling.MaxUnpool1d:旧版本PyTorch JIT对MaxUnpool1d的内部方法处理存在兼容性缺陷。- 移除Unpool层后出现
RuntimeError: isIntList() INTERNAL ASSERT FAILED...:输入张量尺寸与层参数不匹配,导致JIT类型推断失败。
解决方案
1. 升级PyTorch到稳定版本
优先升级到PyTorch 1.12及以上版本,新版本JIT对池化/反池化层的支持更完善,可直接解决方法重定义类错误。
2. 严格匹配MaxPool与MaxUnpool参数
确保MaxUnpool1d的kernel_size、stride参数与对应MaxPool1d完全一致,同时手动指定output_size参数,避免JIT自动推断出错。
3. 添加类型注解辅助JIT推断
给自定义模块的forward方法添加输入输出类型注解,帮助JIT精准识别张量类型与维度。
修改后的完整代码
convBlock类
import torch import torch.nn as nn class convBlock(nn.Module): def __init__(self): super(convBlock, self).__init__() self.conv = nn.Conv1d(1, 64, kernel_size=3, stride=1, padding=1, bias=False) self.batch = nn.BatchNorm1d(64) self.relu = nn.ReLU() self.maxPool = nn.MaxPool1d(kernel_size=3, stride=2, padding=1, return_indices=True) def forward(self, input_1D: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: input_1D = self.conv(input_1D) input_1D = self.relu(self.batch(input_1D)) input_1D, indx_mat = self.maxPool(input_1D) return input_1D, indx_mat
deconvBlock类
class deconvBlock(nn.Module): def __init__(self): super(deconvBlock, self).__init__() self.deconv = nn.ConvTranspose1d(64, 32, kernel_size=3, stride=1, padding=1, bias=False) self.batchNorm = nn.BatchNorm1d(32) self.relu = nn.ReLU() # 与MaxPool1d参数严格匹配 self.unpool = nn.MaxUnpool1d(kernel_size=3, stride=2, padding=1) def forward(self, input_1D: torch.Tensor, idmat: torch.Tensor) -> torch.Tensor: # 手动计算output_size:对应MaxPool stride=2,输出长度为输入长度*2 output_size = input_1D.size(2) * 2 input_1D = self.unpool(input_1D, idmat, output_size=output_size) input_1D = self.deconv(input_1D) input_1D = self.batchNorm(input_1D) input_1D = self.relu(input_1D) return input_1D
modelFull类
class modelFull(nn.Module): def __init__(self): super(modelFull, self).__init__() self.bll = convBlock() self.deconv = deconvBlock() def forward(self, x: torch.Tensor) -> torch.Tensor: xx, y = self.bll(x) xz = self.deconv(xx, y) return xz
验证脚本化
执行以下代码验证修复效果:
# 创建模型并脚本化 model = modelFull() scripted_model = torch.jit.script(model) # 测试输入示例(batch_size=2,通道数=1,序列长度=16) test_input = torch.randn(2, 1, 16) output = scripted_model(test_input) print(output.shape)
内容的提问来源于stack exchange,提问作者Newbie
相关产品推荐
相关产品推荐

