基于PyTorch实现手写数字识别:从全连接网络到卷积神经网络的实战对比

1次阅读
没有评论

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

image.webp

问题背景

MNIST 数据集包含 60,000 张训练图像和 10,000 张测试图像,每张都是 28×28 像素的灰度手写数字(0-9)。在图像分类任务中,我们通常使用准确率作为主要评价指标。

基于 PyTorch 实现手写数字识别:从全连接网络到卷积神经网络的实战对比

全连接网络 (FNN) 在处理图像数据时会面临维度灾难问题。例如,对于 MNIST 的 28×28 图像,输入层就需要 784 个节点。如果第一隐藏层有 500 个节点,仅这一层就会产生 784*500=392,000 个权重参数。这种全连接方式不仅参数量爆炸,而且忽略了图像的局部空间结构信息。

技术方案对比

FNN 与 CNN 结构差异

  • FNN:每个神经元与上一层的所有神经元全连接,适合处理向量化数据
  • CNN:通过卷积核实现局部感受野,具有空间层次化特征提取能力

卷积操作的数学表达:

O[i,j] = ∑∑ (K[m,n] * I[i+m, j+n]) + b

其中 K 是卷积核,I 是输入图像,b 是偏置项。这种局部连接方式大幅减少了参数量。

PyTorch Lightning 实现框架

使用 PyTorch Lightning 可以规范训练流程,以下是一个模板类:

class LitModel(pl.LightningModule):
    def __init__(self):
        super().__init__()
        self.example_input_array = torch.rand(1, 1, 28, 28)

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

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

核心代码实现

FNN 网络构建

class FNN(nn.Module):
    def __init__(self):
        super().__init__()
        self.net = nn.Sequential(nn.Flatten(),
            nn.Linear(784, 512),
            nn.ReLU(),
            nn.Dropout(0.2),  # 防止过拟合
            nn.Linear(512, 256),
            nn.BatchNorm1d(256),  # 批归一化
            nn.ReLU(),
            nn.Linear(256, 10)
        )

CNN 网络构建

class CNN(nn.Module):
    def __init__(self):
        super().__init__()
        self.net = nn.Sequential(nn.Conv2d(1, 32, 3, padding=1),  # 保持空间维度
            nn.ReLU(),
            nn.MaxPool2d(2),
            nn.Conv2d(32, 64, 3, padding=1),
            nn.BatchNorm2d(64),  # 2D 批归一化
            nn.ReLU(),
            nn.MaxPool2d(2),
            nn.Flatten(),
            nn.Linear(64*7*7, 128),
            nn.Dropout(0.5),
            nn.Linear(128, 10)
        )

训练优化技巧

  1. 梯度裁剪防止爆炸
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
  1. 早停机制实现
early_stop = EarlyStopping(
    monitor='val_loss',
    patience=5,
    mode='min'
)

性能优化

训练速度对比

模型类型 每 epoch 时间(GPU) 参数量
FNN 45s 550K
CNN 60s 120K

虽然 CNN 单次迭代稍慢,但达到相同准确率需要的 epoch 更少。

显存优化策略

  • 使用混合精度训练
trainer = Trainer(precision=16)  # 自动混合精度
  • 梯度累积
trainer = Trainer(accumulate_grad_batches=4)  # 相当于增大 batch size

避坑指南

  1. 张量格式问题

PyTorch 默认使用 NCHW 格式,但某些数据加载器可能输出 NHWC。转换方法:

x = x.permute(0, 3, 1, 2)  # NHWC -> NCHW
  1. 学习率调整

批量大小增大 N 倍时,学习率也应增大√N 倍(线性缩放规则)。使用学习率预热:

lr_scheduler = torch.optim.lr_scheduler.LinearLR(optimizer, start_factor=0.1, total_iters=5)
  1. 多 GPU 训练

注意 DataLoader 的 num_workers 设置:

dataloader = DataLoader(..., num_workers=min(4, os.cpu_count()//2))

完整代码与扩展

Colab 运行链接

推荐进一步学习:
–《Deep Learning with PyTorch》书籍
– PyTorch 官方 Tutorials
– torch.profiler 性能分析工具

通过本次实践可以看到,CNN 在图像任务上具有明显优势。虽然实现稍复杂,但通过合理的网络设计和训练技巧,可以轻松达到 99%+ 的准确率。

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