基于PyTorch搭建前馈神经网络实现手写体数字识别:从原理到实战

1次阅读
没有评论

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

image.webp

背景与痛点

手写体数字识别是计算机视觉领域的经典问题,广泛应用于邮政编码识别、银行支票处理、表单信息录入等场景。尽管问题看似简单,但实际开发中常遇到以下挑战:

基于 PyTorch 搭建前馈神经网络实现手写体数字识别:从原理到实战

  • 手写数字的书写风格差异大(如倾斜、连笔、大小不一)
  • 传统图像处理方法(如模板匹配)泛化能力差
  • 模型在小型数据集(如 MNIST)上易过拟合
  • 训练过程中容易出现梯度消失或爆炸

技术选型

前馈神经网络(FNN) vs 卷积神经网络(CNN)

  1. FNN 优势
  2. 结构简单,训练速度快
  3. 适合入门理解神经网络基础原理
  4. 对 MNIST 等简单数据集效果尚可(可达 98%+ 准确率)

  5. CNN 优势

  6. 自动学习局部特征(如边缘、角点)
  7. 参数共享机制更适合图像数据
  8. 在复杂数据集上表现更优

建议:新手建议从 FNN 开始掌握基础,再过渡到 CNN

核心实现

环境准备

import torch
import torch.nn as nn
import torchvision
from torchvision import transforms

数据加载与预处理

  1. 下载 MNIST 数据集

    transform = transforms.Compose([transforms.ToTensor(),
        transforms.Normalize((0.5,), (0.5,))  # 归一化到[-1,1]
    ])
    
    trainset = torchvision.datasets.MNIST(
        root='./data', 
        train=True,
        download=True, 
        transform=transform)
    
    # 批量加载数据(建议 batch_size=64)trainloader = torch.utils.data.DataLoader(
        trainset, 
        batch_size=64,
        shuffle=True)

  2. 可视化样本

    import matplotlib.pyplot as plt
    
    def show_images(images, labels):
        fig, axes = plt.subplots(1, 5, figsize=(12,3))
        for i, ax in enumerate(axes):
            ax.imshow(images[i].numpy().squeeze(), cmap='gray')
            ax.set_title(f'Label: {labels[i]}')
        plt.show()
    
    # 获取一个 batch 的数据
    images, labels = next(iter(trainloader))
    show_images(images, labels)

网络结构定义

class SimpleNN(nn.Module):
    def __init__(self):
        super().__init__()
        self.flatten = nn.Flatten()
        self.fc1 = nn.Linear(28*28, 128)  # 输入层→隐藏层
        self.relu = nn.ReLU()
        self.fc2 = nn.Linear(128, 10)     # 隐藏层→输出层

    def forward(self, x):
        x = self.flatten(x)
        x = self.fc1(x)
        x = self.relu(x)
        x = self.fc2(x)
        return x

model = SimpleNN()
print(model)

关键点说明
nn.Flatten()将 28×28 图像展平成 784 维向量
– 隐藏层使用 ReLU 激活函数避免梯度消失
– 输出层 10 个节点对应 0 - 9 数字分类

训练流程

  1. 初始化损失函数与优化器

    criterion = nn.CrossEntropyLoss()
    optimizer = torch.optim.SGD(model.parameters(), lr=0.01)

  2. 训练循环

    for epoch in range(10):  # 训练 10 个 epoch
        running_loss = 0.0
    
        for images, labels in trainloader:
            # 前向传播
            outputs = model(images)
            loss = criterion(outputs, labels)
    
            # 反向传播
            optimizer.zero_grad()
            loss.backward()
            optimizer.step()
    
            running_loss += loss.item()
    
        print(f'Epoch {epoch+1}, Loss: {running_loss/len(trainloader):.4f}')

性能优化

超参数调优

  1. 学习率选择
  2. 太大(如 >0.1):损失震荡不收敛
  3. 太小(如 <0.001):收敛速度慢
  4. 推荐尝试:0.01→0.001 阶梯下降

  5. 批量大小(Batch Size)

  6. 较小值(如 32):梯度估计噪声大
  7. 较大值(如 256):内存占用高
  8. 推荐值:64 或 128

  9. 隐藏层节点数

  10. 太少:模型容量不足
  11. 太多:易过拟合
  12. 经验公式:输入层与输出层节点数的几何平均数(如√(784*10)≈89)

改进网络结构

class ImprovedNN(nn.Module):
    def __init__(self):
        super().__init__()
        self.net = nn.Sequential(nn.Flatten(),
            nn.Linear(28*28, 256),
            nn.BatchNorm1d(256),  # 添加批归一化
            nn.ReLU(),
            nn.Dropout(0.2),      # 添加 Dropout
            nn.Linear(256, 10)
        )

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

优化点
– 批归一化 (BatchNorm) 加速收敛
– Dropout 减少过拟合

避坑指南

常见问题及解决方案

  1. 损失不下降
  2. 检查数据是否正常加载(可视化样本)
  3. 检查学习率是否过小
  4. 检查网络结构是否正确(如忘记加激活函数)

  5. 过拟合

  6. 增加 Dropout 层
  7. 使用 L2 正则化
  8. 早停(Early Stopping)

  9. 梯度爆炸

  10. 使用梯度裁剪(torch.nn.utils.clip_grad_norm_
  11. 改用更稳定的激活函数(如 ReLU 替代 Sigmoid)

延伸思考

模型部署方案

  1. 导出为 TorchScript

    traced_model = torch.jit.script(model)
    traced_model.save('mnist_fnn.pt')

  2. Web API 服务

  3. 使用 Flask/FastAPI 封装预测接口
  4. 示例请求处理:

    @app.route('/predict', methods=['POST'])
    def predict():
        img = request.files['image'].read()
        img = preprocess(img)  # 转换为 Tensor
        with torch.no_grad():
            output = model(img)
        return {'prediction': int(torch.argmax(output))}

  5. 移动端部署

  6. 通过 ONNX 转换为平台兼容格式
  7. 使用 PyTorch Mobile 在 Android/iOS 端运行

结语

通过本文的实践,我们完成了从数据加载到模型部署的完整流程。虽然前馈神经网络在 MNIST 上表现尚可,但要处理更复杂的图像任务(如 CIFAR-10),建议转向 CNN 架构。后续可尝试:

  • 改用卷积神经网络 (CNN) 提升准确率
  • 使用数据增强 (Data Augmentation) 增加样本多样性
  • 尝试迁移学习 (Transfer Learning) 加速训练

完整的代码已上传至 GitHub 仓库(虚构地址):
github.com/username/mnist-fnn-demo

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