PyTorch使用Flatten层报tuple不可调用错误的解决方法
PyTorch自定义模型调用Flatten层报错排查
问题复现
自定义模型代码
import torch import torch.nn as nn import torch.nn.functional as F class MyModel(nn.Module): def __init__(self, input_size, num_classes): super(MyModel, self).__init__() self.layer_1 = nn.Conv1d(1, 16, 3, bias=False, stride=2) self.activation_1 = F.relu self.adap = nn.AdaptiveAvgPool1d(1) self.flatten = nn.Flatten(), self.layer_2 = torch.nn.Linear(2249, 500) self.activation_2 = F.relu self.layer_3 = torch.nn.Linear(500, 2) def forward(self, x, labels=None): x = x.reshape(256, 1, -1) x = self.layer_1(x) x = self.activation_1(x) x = self.flatten(x) return x
测试调用代码
model = MyModel(input_size=4500, num_classes=2) torchinfo.summary(model, (256, 4500))
触发的报错信息
Input In [101], in MyModel.forward(self, x, labels) 30 x = self.activation_1(x) —> 31 x = self.flatten(x) 32 return x TypeError: ‘tuple’ object is not callable
问题解答
1. 错误产生原因
报错的直接诱因是__init__方法中定义self.flatten时,行尾多写了一个半角逗号。
Python语法规则下,单个对象赋值时末尾加逗号,会自动将该对象包装为单元素元组。也就是说此时self.flatten不是预期的nn.Flatten()可调用层实例,而是一个仅包含一个nn.Flatten实例的元组对象。元组本身是不可调用的,当在forward里写self.flatten(x)尝试像调用函数/层一样传入参数运行时,就会触发tuple object is not callable的类型错误。
2. 代码修改方案
- 第一步先修复直接触发报错的语法问题:删除
self.flatten赋值行末尾的多余逗号,修正为:
self.flatten = nn.Flatten()
- 同步修复代码中其他会导致后续运行错误的逻辑问题:
- 不要在forward方法中把batch维度硬编码为256,改为动态读取输入的batch size,否则更换batch size输入时会触发维度不匹配错误,将
x = x.reshape(256, 1, -1)修改为:x = x.reshape(x.shape[0], 1, -1) __init__中定义的自适应平均池化层self.adap、全连接层self.layer_2、激活函数self.activation_2、输出层self.layer_3目前没有加入前向传播流程,如果需要实现完整的分类逻辑,需要按网络结构顺序将这些层补充到forward方法中,否则模型经过卷积、激活、展平后就直接返回输出,无法完成分类任务。
- 不要在forward方法中把batch维度硬编码为256,改为动态读取输入的batch size,否则更换batch size输入时会触发维度不匹配错误,将
内容的提问来源于stack exchange,提问作者user3668129
相关产品推荐
相关产品推荐

