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
相关产品推荐
相关产品推荐

