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

PyTorch调用接收两个参数的自定义Module报错如何解决

错误原因
  • 第一个AttributeError: 'int' object has no attribute 'dim':
    你之前的VerySimple、VerySimple2传入数值可正常运行,大概率是这两个模块的forward逻辑中先对输入做了张量类型转换,而Simple2没有做该处理。PyTorch的nn.Module内置层、torch.cat等张量操作的输入必须是torch.Tensor类型,你直接传入int类型的2、3,运算逻辑尝试访问输入的dim属性做校验时,int类型没有该属性,直接触发报错。
  • 第二个IndexError: Dimension out of range:
    torch.cat的核心要求是:所有待拼接的张量必须存在你指定的拼接维度,且非拼接维度的尺寸完全一致。你传入的torch.tensor([2.0])和torch.tensor([3.0])都是1维张量,仅支持dim=0的维度索引,如果你调用torch.cat时指定的拼接维度≥1(比如常见的dim=1做特征维度拼接),就会因为目标维度不存在触发越界错误。
    你后续将拼接改为逐元素相加,仅要求两个张量shape一致即可运算,两个经过seq推理后的输出shape相同,所以运行成功。
最小可运行示例

错误复现代码

import torch
from torch import nn

class Simple2(nn.Module):
    def __init__(self):
        super().__init__()
        self.seq1 = nn.Sequential(nn.Linear(1, 4), nn.ReLU())
        self.seq2 = nn.Sequential(nn.Linear(1, 4), nn.ReLU())
    
    def forward(self, x, y):
        out1 = self.seq1(x)
        out2 = self.seq2(y)
        # 指定dim=1拼接,1维张量不存在该维度,触发IndexError
        return torch.cat([out1, out2], dim=1)

s2 = Simple2()
# 第一次调用触发AttributeError
# s2(2, 3)

# 改为张量输入触发IndexError
# x = torch.tensor([2.0])
# y = torch.tensor([3.0])
# s2(x, y)

修复版本1:适配torch.cat拼接逻辑

适合需要保留拼接操作的场景,通过给张量增加维度适配拼接要求:

class Simple2Fix1(nn.Module):
    def __init__(self):
        super().__init__()
        self.seq1 = nn.Sequential(nn.Linear(1, 4), nn.ReLU())
        self.seq2 = nn.Sequential(nn.Linear(1, 4), nn.ReLU())
    
    def forward(self, x, y):
        # 给输入增加batch维度,转为2维张量shape:[1,1],支持dim=1拼接
        x = x.unsqueeze(0)
        y = y.unsqueeze(0)
        out1 = self.seq1(x)
        out2 = self.seq2(y)
        return torch.cat([out1, out2], dim=1)

s2_fix1 = Simple2Fix1()
x = torch.tensor([2.0])
y = torch.tensor([3.0])
print(s2_fix1(x, y)) # 输出shape为[1,8]的张量,运行正常

修复版本2:保留逐元素相加逻辑

和你最后修改后的运行效果一致:

class Simple2Fix2(nn.Module):
    def __init__(self):
        super().__init__()
        self.seq1 = nn.Sequential(nn.Linear(1, 4), nn.ReLU())
        self.seq2 = nn.Sequential(nn.Linear(1, 4), nn.ReLU())
    
    def forward(self, x, y):
        out1 = self.seq1(x)
        out2 = self.seq2(y)
        return out1 + out2

s2_fix2 = Simple2Fix2()
x = torch.tensor([2.0])
y = torch.tensor([3.0])
print(s2_fix2(x, y)) # 输出shape为[4]的张量,运行正常

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.25 21:15:02