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

1次阅读
没有评论

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

image.webp

1. MNIST 数据集与手写数字识别

MNIST 数据集包含 6 万张 28×28 像素的手写数字灰度图像,是计算机视觉领域的经典入门数据集。每张图片都标注了 0 - 9 对应的数字标签。这个数据集虽然简单,但很好地体现了图像分类任务的核心挑战:

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

  • 不同人的书写风格差异大(如数字 ’7’ 带横线或不带)
  • 数字在图像中的位置和大小略有变化
  • 笔画粗细和倾斜角度各不相同

手写数字识别技术在实际中有广泛应用场景,比如银行支票识别、快递单号自动录入等。通过这个案例,我们可以学习如何将真实世界的图像转换为计算机可以理解的数值表示,并训练模型从中提取规律。

2. FNN 与 CNN 的架构对比

前馈神经网络(FNN)

  • 结构特点:全连接层堆叠,前一层的每个神经元都与后一层所有神经元相连
  • 优势
  • 结构简单,易于实现
  • 对计算资源要求较低
  • 劣势
  • 参数量大(28×28 图像展平后输入层就有 784 个节点)
  • 忽略图像的空间局部性特征
  • 对平移、旋转等变化敏感

卷积神经网络(CNN)

  • 结构特点:交替使用卷积层和池化层,最后接全连接层
  • 优势
  • 通过卷积核提取局部特征
  • 参数共享大幅减少参数量
  • 对平移、缩放有一定不变性
  • 劣势
  • 计算复杂度较高
  • 需要调整的超参数更多(卷积核大小、步长等)

3. 完整实现步骤

3.1 数据准备

import torch
from torchvision import datasets, transforms

# 定义数据预处理
transform = transforms.Compose([transforms.ToTensor(),  # 转换为 Tensor 并归一化到[0,1]
    transforms.Normalize((0.1307,), (0.3081,))  # MNIST 的均值和标准差
])

# 加载数据集
train_dataset = datasets.MNIST('./data', 
                              train=True, 
                              download=True,
                              transform=transform)
test_dataset = datasets.MNIST('./data',
                             train=False,
                             transform=transform)

# 创建数据加载器
train_loader = torch.utils.data.DataLoader(train_dataset,
                                         batch_size=64,
                                         shuffle=True)
test_loader = torch.utils.data.DataLoader(test_dataset,
                                        batch_size=1000,
                                        shuffle=False)

3.2 FNN 模型实现

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

class FNN(nn.Module):
    def __init__(self):
        super(FNN, self).__init__()
        self.fc1 = nn.Linear(784, 512)  # 输入层到隐藏层
        self.fc2 = nn.Linear(512, 256)
        self.fc3 = nn.Linear(256, 10)   # 输出层
        self.dropout = nn.Dropout(0.2)  # 防止过拟合

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

3.3 CNN 模型实现

class CNN(nn.Module):
    def __init__(self):
        super(CNN, self).__init__()
        self.conv1 = nn.Conv2d(1, 32, 3, 1)  # 输入通道, 输出通道, 卷积核大小, 步长
        self.conv2 = nn.Conv2d(32, 64, 3, 1)
        self.dropout1 = nn.Dropout2d(0.25)
        self.dropout2 = nn.Dropout2d(0.5)
        self.fc1 = nn.Linear(9216, 128)  # 64*12*12=9216
        self.fc2 = nn.Linear(128, 10)

    def forward(self, x):
        x = self.conv1(x)
        x = F.relu(x)
        x = self.conv2(x)
        x = F.relu(x)
        x = F.max_pool2d(x, 2)  # 池化层
        x = self.dropout1(x)
        x = torch.flatten(x, 1)
        x = self.fc1(x)
        x = F.relu(x)
        x = self.dropout2(x)
        x = self.fc2(x)
        return F.log_softmax(x, dim=1)

3.4 训练流程

def train(model, device, train_loader, optimizer, epoch):
    model.train()
    for batch_idx, (data, target) in enumerate(train_loader):
        data, target = data.to(device), target.to(device)
        optimizer.zero_grad()
        output = model(data)
        loss = F.nll_loss(output, target)  # 负对数似然损失
        loss.backward()
        optimizer.step()
        if batch_idx % 100 == 0:
            print(f'Train Epoch: {epoch} [{batch_idx * len(data)}/{len(train_loader.dataset)}]\tLoss: {loss.item():.6f}')

3.5 测试评估

def test(model, device, test_loader):
    model.eval()
    test_loss = 0
    correct = 0
    with torch.no_grad():
        for data, target in test_loader:
            data, target = data.to(device), target.to(device)
            output = model(data)
            test_loss += F.nll_loss(output, target, reduction='sum').item()
            pred = output.argmax(dim=1, keepdim=True)
            correct += pred.eq(target.view_as(pred)).sum().item()

    test_loss /= len(test_loader.dataset)
    print(f'\nTest set: Average loss: {test_loss:.4f}, Accuracy: {correct}/{len(test_loader.dataset)} ({100. * correct / len(test_loader.dataset):.2f}%)\n')
    return test_loss, correct / len(test_loader.dataset)

4. 性能对比与分析

指标 FNN 模型 CNN 模型
测试准确率 97.8% 99.2%
训练时间(秒 /epoch) 45 68
参数量 669,706 119,978

从结果可以看出:

  1. CNN 在准确率上明显优于 FNN,特别是在处理笔画变形、位置偏移等情况时
  2. FNN 训练更快,但这是以牺牲准确率为代价的
  3. 虽然 CNN 结构更复杂,但由于参数共享机制,实际参数量反而更少

5. 常见问题与解决方案

学习率设置

  • 问题表现:损失值震荡不收敛或下降非常缓慢
  • 解决方案
  • 初始学习率通常设为 0.001-0.01
  • 使用学习率调度器如torch.optim.lr_scheduler.StepLR
  • 监控训练损失曲线调整

内存不足

  • 原因:Batch Size 设置过大
  • 建议
  • 从较小的 batch size(如 32)开始
  • 使用 torch.cuda.empty_cache() 清理缓存
  • 考虑梯度累积技术

过拟合处理

  • 识别方法:训练准确率远高于测试准确率
  • 应对策略
  • 增加 Dropout 层
  • 添加 L2 正则化(weight decay)
  • 使用数据增强(如随机旋转、平移)

6. 拓展方向

  1. 模型部署:使用 Flask 或 FastAPI 将训练好的模型封装为 Web API
  2. 性能优化
  3. 尝试更复杂的 CNN 架构如 ResNet
  4. 使用注意力机制提升特征提取能力
  5. 实际应用
  6. 扩展到更复杂的字符识别(如 EMNIST)
  7. 应用到验证码识别等场景

通过这个项目,我们不仅学习了 PyTorch 的基本用法,更重要的是理解了不同神经网络架构的设计思想。建议读者可以尝试调整网络结构、超参数,观察对模型性能的影响,这是提升深度学习实践能力的最佳途径。

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