PyTorch手写数字识别模型无训练提升问题排查求助
MNIST模型训练异常问题排查与修复
核心问题分析
- 损失函数不匹配:
F.nll_loss要求输入是经过log_softmax处理后的概率对数,但你的模型最后一层直接输出原始logits,这会导致损失计算错误,出现负值。 - 学习率过低:SGD优化器使用
lr=0.0001太小,参数更新幅度微乎其微,模型无法有效学习。 - 测试函数笔误:
ouput拼写错误,会导致测试阶段无法正确获取模型输出。
修复后的代码
import torch import torch.nn as nn import torch.nn.functional as F import torch.optim as optim import torchvision class Net(nn.Module): def __init__(self): super(Net, self).__init__() self.conv1 = nn.Conv2d(1,5,5) self.conv2 = nn.Conv2d(5,10,4) self.fc1 = nn.Linear(160,120) self.fc2 = nn.Linear(120,84) self.fc3 = nn.Linear(84,10) def forward(self, x): x = self.conv1(x) x = F.relu(x) x = F.max_pool2d(x, (2,2)) x = self.conv2(x) x = F.relu(x) x = F.max_pool2d(x,2) x = torch.flatten(x,1) x = self.fc1(x) x = F.relu(x) x = self.fc2(x) x = F.relu(x) x = self.fc3(x) # 适配nll_loss,添加log_softmax处理 return F.log_softmax(x, dim=1) net = Net() # 提高学习率,调整momentum至MNIST常用值 optimizer = optim.SGD(net.parameters(), lr=0.01, momentum=0.9) train_loader = torch.utils.data.DataLoader( torchvision.datasets.MNIST('~/files/', train=True, download=True, transform=torchvision.transforms.Compose([ torchvision.transforms.ToTensor(), torchvision.transforms.Normalize( (0.1307,), (0.3081,)) ])), batch_size=64, shuffle=True) test_loader = torch.utils.data.DataLoader( torchvision.datasets.MNIST('~/files/', train=False, download=True, transform=torchvision.transforms.Compose([ torchvision.transforms.ToTensor(), torchvision.transforms.Normalize( (0.1307,), (0.3081,)) ])), batch_size=64, shuffle=True) def train(epoch): net.train() for batch, (data, target) in enumerate(train_loader): optimizer.zero_grad() output = net(data) loss = F.nll_loss(output, target) print(loss.item()) loss.backward() optimizer.step() if batch % 10 == 0: torch.save(net.state_dict(), 'model.pth') torch.save(optimizer.state_dict(), 'optimizer.pth') def test(): net.eval() num_correct = 0 with torch.no_grad(): for data,target in test_loader: # 修正拼写错误 output = net(data) pred = output.data.max(1, keepdim=True)[1] num_correct += pred.eq(target.data.view_as(pred)).sum() print(f"测试集正确数: {num_correct.item()}/{len(test_loader.dataset)}") print(f"测试集准确率: {num_correct.item()/len(test_loader.dataset):.4f}") for i in range(1,15): print(f"===== 第{i}轮训练 =====") train(i) test()
额外优化建议
- 可以给卷积层添加
padding,避免特征图尺寸过度缩小,比如nn.Conv2d(1,5,5, padding=2),这样输入28x28经过卷积后仍为28x28,池化后得到14x14的特征图,保留更多细节。 - 训练时可添加训练集准确率统计,方便同步观察训练效果。
- 尝试使用
Adam优化器(学习率设为0.001),收敛速度通常比SGD更快。
内容的提问来源于stack exchange,提问作者oli2
相关产品推荐
相关产品推荐

