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

PyTorch JIT脚本化自定义模型失败,求助解决方法

解决PyTorch JIT脚本化MaxUnpool1D相关错误

问题分析

调用torch.jit.script(modelFull())时出现两类错误:

  1. RuntimeError: Can't redefine method: forward on class: __torch__.torch.nn.modules.pooling.MaxUnpool1d:旧版本PyTorch JIT对MaxUnpool1d的内部方法处理存在兼容性缺陷。
  2. 移除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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.14 22:59:50