PyTorch数字分类代码报错NotImplementedError:CNN模块缺失forward函数
解决PyTorch中
NotImplementedError: Module [CNN] is missing the required "forward" function报错 问题原因
你自定义的CNN类继承了torch.nn.Module,但没有实现必须的forward方法——这个方法是PyTorch模型的核心,定义了数据通过模型的前向传播逻辑,没有它模型无法运行。
解决步骤
- 找到你的
CNN类定义代码,确保它包含forward方法。以下是一个适用于数字分类(比如MNIST)的基础CNN类示例:
import torch.nn as nn import torch.nn.functional as F class CNN(nn.Module): def __init__(self): super(CNN, self).__init__() # 定义卷积层、池化层、全连接层 self.conv1 = nn.Conv2d(1, 32, kernel_size=3, stride=1, padding=1) self.pool = nn.MaxPool2d(2, 2) self.conv2 = nn.Conv2d(32, 64, kernel_size=3, stride=1, padding=1) self.fc1 = nn.Linear(64 * 7 * 7, 128) # 适配28x28的MNIST输入图像 self.fc2 = nn.Linear(128, 10) # 对应10个数字分类类别 def forward(self, x): # 定义前向传播逻辑 x = self.pool(F.relu(self.conv1(x))) x = self.pool(F.relu(self.conv2(x))) x = x.view(-1, 64 * 7 * 7) # 将特征图展平为一维张量 x = F.relu(self.fc1(x)) x = self.fc2(x) return x
- 确认你实例化模型时使用的是这个完整的
CNN类:
model = CNN().to(device)
- 重新运行训练代码,此时模型可正常执行前向传播,不会再触发该报错。
注意事项
- 所有继承自
nn.Module的自定义模型类,必须实现forward方法,这是PyTorch的强制要求。 - 若你是跟随教程编写代码,检查是否漏抄了
forward方法的代码,或是类定义存在语法错误导致方法未被正确识别。
内容的提问来源于stack exchange,提问作者Narimene Abdelli
相关产品推荐
相关产品推荐

