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

PyTorch JIT脚本报错:Sequential容器传入Tuple输入时出错

PyTorch JIT下Sequential处理Tuple输入类型推断错误的修复

问题背景

实现了一个简易网络,子模块MyBatchNorm的forward方法接受Tuple[Tensor, int]作为输入并返回同类型结果,但使用nn.Sequential组合这些模块后,JIT将Sequential的forward输入类型推断为Tensor而非Tuple,导致运行报错。

复现代码

from typing import Tuple
import torch
import torch.nn as nn

class MyBatchNorm(nn.Module):
    def __init__(self, output_size, d_ids):
        super().__init__()
        self.d_ids = d_ids
        self.net = nn.ModuleDict({f"{d}": nn.BatchNorm1d(output_size) for d in d_ids})
    
    def forward(self, input_tuple: Tuple[torch.Tensor, int]) -> Tuple[torch.Tensor, int]:
        input_tensor, d = input_tuple
        output_tensor = torch.tensor([])
        for d_name, d_norm in self.net.items():
            if f"{d}" == d_name:
                output_tensor = d_norm(input_tensor)
        if len(output_tensor) == 0:
            raise ValueError(f"invalid d {d}, must be {self.d_ids}")
        return output_tensor, d

class MyNet(nn.Module):
    def __init__(self, output_size, d_ids):
        super().__init__()
        dense_layers = [
            MyBatchNorm(output_size, d_ids),
            MyBatchNorm(output_size, d_ids)
        ]
        self.net = torch.nn.Sequential(*dense_layers)
        
    def forward(self, input_tensor: torch.Tensor, d_tensor: torch.Tensor) -> torch.Tensor:
        d = d_tensor.squeeze()[0].item()
        output_tensor, _ = self.net((input_tensor, d))
        return torch.squeeze(output_tensor)

报错信息

RuntimeError: 

forward(__torch__.___torch_mangle_16.MyBatchNorm self, (Tensor, int) input_tuple) -> ((Tensor, int)):
Expected a value of type 'Tuple[Tensor, int]' for argument 'input_tuple' but instead found type 'Tensor (inferred)'.
Inferred the value for argument 'input_tuple' to be of type 'Tensor' because it was not annotated with an explicit type.
:
  File "/home/ec2-user/anaconda3/envs/pytorch_latest_p36/lib/python3.6/site-packages/torch/nn/modules/container.py", line 117
    def forward(self, input):
        for module in self:
            input = module(input)
                    ~~~~~~ <--- HERE
        return input

修复方案

原因分析

nn.Sequential的默认forward方法没有添加类型注解,JIT无法正确推断其输入为Tuple类型。当第一个MyBatchNorm返回Tuple后,Sequential会错误地将其当成Tensor传给下一个模块,导致类型不匹配。

方法1:自定义支持Tuple的Sequential容器

继承nn.Sequential并给forward方法添加明确的类型注解,让JIT能正确识别输入输出类型:

from typing import Tuple
import torch
import torch.nn as nn

class TupleSequential(nn.Sequential):
    def forward(self, input: Tuple[torch.Tensor, int]) -> Tuple[torch.Tensor, int]:
        for module in self:
            input = module(input)
        return input

# 修改MyNet中的容器为自定义的TupleSequential
class MyNet(nn.Module):
    def __init__(self, output_size, d_ids):
        super().__init__()
        dense_layers = [
            MyBatchNorm(output_size, d_ids),
            MyBatchNorm(output_size, d_ids)
        ]
        self.net = TupleSequential(*dense_layers)
        
    def forward(self, input_tensor: torch.Tensor, d_tensor: torch.Tensor) -> torch.Tensor:
        d = d_tensor.squeeze()[0].item()
        output_tensor, _ = self.net((input_tensor, d))
        return torch.squeeze(output_tensor)

方法2:手动调用子模块

放弃使用nn.Sequential,在MyNet的forward中依次调用每个子模块,明确传递Tuple输入:

class MyNet(nn.Module):
    def __init__(self, output_size, d_ids):
        super().__init__()
        self.bn1 = MyBatchNorm(output_size, d_ids)
        self.bn2 = MyBatchNorm(output_size, d_ids)
        
    def forward(self, input_tensor: torch.Tensor, d_tensor: torch.Tensor) -> torch.Tensor:
        d = d_tensor.squeeze()[0].item()
        x, d = self.bn1((input_tensor, d))
        x, d = self.bn2((x, d))
        return torch.squeeze(x)

两种方法都能解决JIT类型推断错误的问题,可根据实际场景选择。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.29 00:35:37