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

PyTorch神经网络报错:linear()参数input应为Tensor而非Flatten

问题解决:TypeError: linear(): argument 'input' must be Tensor, not Flatten

错误原因

你在forward方法里直接调用nn.Flatten(x, 1)是错误用法。nn.Flatten是PyTorch的模块类,这样写会返回Flatten模块实例而非张量,导致后续全连接层接收到非张量输入,触发类型错误。

两种修复方案

方案1:使用函数式的torch.flatten

修改forward中的展平代码,替换为函数式API:

def forward(self, x):
    x = self.pool(F.relu(self.conv1(x)))
    x = self.pool(F.relu(self.conv2(x)))
    x = torch.flatten(x, 1)  # 替换原nn.Flatten(x,1)
    x = F.relu(self.fc1(x))
    x = F.relu(self.fc2(x))
    x = self.fc3(x)
    return x

方案2:在__init__中实例化Flatten模块

先在模型初始化时创建Flatten模块实例,再在forward中调用:

class Network(nn.Module):
    def __init__(self):
        super(Network, self).__init__()
        self.conv1 = nn.Conv2d(3, 6, 5)
        self.pool = nn.MaxPool2d(2, 2)
        self.conv2 = nn.Conv2d(6, 16, 5)
        self.flatten = nn.Flatten(1)  # 新增Flatten模块实例
        self.fc1 = nn.Linear(600, 120)
        self.fc2 = nn.Linear(120, 2)
        self.fc3 = nn.Linear(2, 10)
    
    def forward(self, x):
        x = self.pool(F.relu(self.conv1(x)))
        x = self.pool(F.relu(self.conv2(x)))
        x = self.flatten(x)  # 调用实例化后的模块
        x = F.relu(self.fc1(x))
        x = F.relu(self.fc2(x))
        x = self.fc3(x)
        return x

额外注意点

你当前fc1的输入维度设为600,需确保卷积池化后的张量展平后恰好是600维。可在forward中临时添加print(x.shape)查看展平前的张量形状,计算正确的展平维度,避免后续触发维度不匹配错误。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.13 17:01:19