PyTorch实战:从零搭建前馈神经网络/CNN/RNN实现MNIST手写数字识别

1次阅读
没有评论

共计 3190 个字符,预计需要花费 8 分钟才能阅读完成。

image.webp

MNIST 数据集与图像分类基础

MNIST 是深度学习领域的 ”Hello World” 数据集,包含 60,000 张训练图片和 10,000 张测试图片,每张都是 28×28 像素的灰度手写数字(0-9)。这个看似简单的任务背后有几个关键技术挑战:

PyTorch 实战:从零搭建前馈神经网络 /CNN/RNN 实现 MNIST 手写数字识别

  • 图像预处理:需要将像素值归一化到 0 - 1 范围
  • 特征提取:如何从原始像素中学习有效特征
  • 模型泛化:防止记住训练样本但无法识别新样本

网络结构选型指南

  1. 前馈神经网络(FNN)
  2. 最基础的全连接网络
  3. 将 28×28 图像展平为 784 维向量
  4. 适合理解神经网络基本原理
  5. 参数量大,准确率约 98%

  6. 卷积神经网络(CNN)

  7. 通过卷积核自动提取空间特征
  8. 保留图像二维结构信息
  9. 参数量少,准确率可达 99% 以上
  10. 适合视觉类任务

  11. 循环神经网络(RNN)

  12. 将图像按行 / 列序列处理
  13. 理论上可以捕捉笔顺信息
  14. 实际效果通常不如 CNN
  15. 适合演示 RNN 在图像上的应用

实战代码详解

数据准备

import torch
from torchvision import datasets, transforms

# 数据增强策略
transform = transforms.Compose([transforms.ToTensor(),
    transforms.Normalize((0.1307,), (0.3081,))
])

# 自动下载并加载数据集
train_set = datasets.MNIST('./data', train=True, download=True, transform=transform)
test_set = datasets.MNIST('./data', train=False, transform=transform)

# 创建数据加载器
batch_size = 64
train_loader = torch.utils.data.DataLoader(train_set, batch_size=batch_size, shuffle=True)
test_loader = torch.utils.data.DataLoader(test_set, batch_size=batch_size)

FNN 模型实现

import torch.nn as nn
import torch.nn.functional as F

class FNN(nn.Module):
    def __init__(self):
        super().__init__()
        self.fc1 = nn.Linear(784, 512)  # 输入层→隐藏层
        self.fc2 = nn.Linear(512, 256)
        self.fc3 = nn.Linear(256, 10)   # 隐藏层→输出层

    def forward(self, x):
        x = x.view(-1, 784)  # 展平图像 [B,1,28,28]→[B,784]
        x = F.relu(self.fc1(x))
        x = F.relu(self.fc2(x))
        return F.log_softmax(self.fc3(x), dim=1)

CNN 模型实现

class CNN(nn.Module):
    def __init__(self):
        super().__init__()
        self.conv1 = nn.Conv2d(1, 32, 3, 1)  # [B,1,28,28]→[B,32,26,26]
        self.conv2 = nn.Conv2d(32, 64, 3, 1) # →[B,64,24,24]
        self.fc1 = nn.Linear(64*12*12, 128)
        self.fc2 = nn.Linear(128, 10)

    def forward(self, x):
        x = F.relu(self.conv1(x))
        x = F.max_pool2d(x, 2)       # →[B,64,12,12]
        x = F.relu(self.conv2(x))
        x = x.view(-1, 64*12*12)     # 展平
        x = F.relu(self.fc1(x))
        return F.log_softmax(self.fc2(x), dim=1)

RNN 模型实现

class RNN(nn.Module):
    def __init__(self):
        super().__init__()
        self.rnn = nn.LSTM(
            input_size=28,  # 每行像素数
            hidden_size=128,
            num_layers=2,
            batch_first=True
        )
        self.fc = nn.Linear(128, 10)

    def forward(self, x):
        # [B,1,28,28]→[B,28,28](去除通道)→[B,28,28](行序列)x = x.squeeze(1).permute(0, 2, 1) 
        _, (h_n, _) = self.rnn(x)    # h_n 形状[2,B,128]
        return F.log_softmax(self.fc(h_n[-1]), dim=1)  # 取最后一层

训练流程标准化

使用 PyTorch Lightning 规范训练循环:

import pytorch_lightning as pl

class LitModel(pl.LightningModule):
    def __init__(self, model):
        super().__init__()
        self.model = model

    def forward(self, x):
        return self.model(x)

    def training_step(self, batch, batch_idx):
        x, y = batch
        y_hat = self(x)
        loss = F.nll_loss(y_hat, y)
        self.log('train_loss', loss)
        return loss

    def configure_optimizers(self):
        return torch.optim.Adam(self.parameters(), lr=0.001)

# 实例化并训练
model = LitModel(CNN())
trainer = pl.Trainer(max_epochs=10, gpus=1)
trainer.fit(model, train_loader)

性能优化技巧

  1. 批处理大小
  2. GPU 显存充足:增大 batch size(如 256)加速训练
  3. 显存有限:减小 batch size(如 32)配合梯度累积

  4. 学习率调整

  5. 初始尝试:1e- 3 到 1e-4
  6. 使用 ReduceLROnPlateau 自动调整

    scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, 'min')

  7. 数据增强

  8. 添加随机旋转 / 平移
    transforms.RandomRotation(10),
    transforms.RandomAffine(0, translate=(0.1,0.1))

常见问题解决方案

  1. 维度不匹配错误
  2. 检查各层输入 / 输出形状
  3. 使用 print(x.shape) 调试

  4. GPU 内存不足

  5. 减小 batch size
  6. 使用torch.cuda.empty_cache()
  7. 混合精度训练

    trainer = pl.Trainer(precision=16)

  8. 过拟合识别

  9. 训练误差持续下降但测试误差上升
  10. 解决方案:
    • 增加 Dropout 层
    • 添加 L2 正则化
    • 早停(EarlyStopping)

模型保存与加载

# 保存
torch.save(model.state_dict(), 'mnist_cnn.pt')

# 加载
model = CNN()
model.load_state_dict(torch.load('mnist_cnn.pt'))
model.eval()

延伸思考

  1. 模型部署:如何用 Flask 将训练好的模型封装为 REST API?
  2. 非平衡数据:当某些数字样本过少时,该如何调整损失函数?
  3. 实时识别:如何扩展本项目实现摄像头实时手写数字识别?

通过这个完整的实践流程,相信你已经掌握了 PyTorch 实现图像分类的核心方法。建议尝试调整网络结构超参数,观察对模型性能的影响,这是提升深度学习实战能力的最佳途径。

正文完
 0
评论(没有评论)