共计 2995 个字符,预计需要花费 8 分钟才能阅读完成。
背景与痛点
手写体数字识别是计算机视觉领域的经典问题,广泛应用于邮政编码识别、银行支票处理、表单信息录入等场景。尽管问题看似简单,但实际开发中常遇到以下挑战:

- 手写数字的书写风格差异大(如倾斜、连笔、大小不一)
- 传统图像处理方法(如模板匹配)泛化能力差
- 模型在小型数据集(如 MNIST)上易过拟合
- 训练过程中容易出现梯度消失或爆炸
技术选型
前馈神经网络(FNN) vs 卷积神经网络(CNN)
- FNN 优势:
- 结构简单,训练速度快
- 适合入门理解神经网络基础原理
-
对 MNIST 等简单数据集效果尚可(可达 98%+ 准确率)
-
CNN 优势:
- 自动学习局部特征(如边缘、角点)
- 参数共享机制更适合图像数据
- 在复杂数据集上表现更优
建议:新手建议从 FNN 开始掌握基础,再过渡到 CNN
核心实现
环境准备
import torch
import torch.nn as nn
import torchvision
from torchvision import transforms
数据加载与预处理
-
下载 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) -
可视化样本
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 数字分类
训练流程
-
初始化损失函数与优化器
criterion = nn.CrossEntropyLoss() optimizer = torch.optim.SGD(model.parameters(), lr=0.01) -
训练循环
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}')
性能优化
超参数调优
- 学习率选择
- 太大(如 >0.1):损失震荡不收敛
- 太小(如 <0.001):收敛速度慢
-
推荐尝试:0.01→0.001 阶梯下降
-
批量大小(Batch Size)
- 较小值(如 32):梯度估计噪声大
- 较大值(如 256):内存占用高
-
推荐值:64 或 128
-
隐藏层节点数
- 太少:模型容量不足
- 太多:易过拟合
- 经验公式:输入层与输出层节点数的几何平均数(如√(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 减少过拟合
避坑指南
常见问题及解决方案
- 损失不下降
- 检查数据是否正常加载(可视化样本)
- 检查学习率是否过小
-
检查网络结构是否正确(如忘记加激活函数)
-
过拟合
- 增加 Dropout 层
- 使用 L2 正则化
-
早停(Early Stopping)
-
梯度爆炸
- 使用梯度裁剪(
torch.nn.utils.clip_grad_norm_) - 改用更稳定的激活函数(如 ReLU 替代 Sigmoid)
延伸思考
模型部署方案
-
导出为 TorchScript
traced_model = torch.jit.script(model) traced_model.save('mnist_fnn.pt') -
Web API 服务
- 使用 Flask/FastAPI 封装预测接口
-
示例请求处理:
@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))} -
移动端部署
- 通过 ONNX 转换为平台兼容格式
- 使用 PyTorch Mobile 在 Android/iOS 端运行
结语
通过本文的实践,我们完成了从数据加载到模型部署的完整流程。虽然前馈神经网络在 MNIST 上表现尚可,但要处理更复杂的图像任务(如 CIFAR-10),建议转向 CNN 架构。后续可尝试:
- 改用卷积神经网络 (CNN) 提升准确率
- 使用数据增强 (Data Augmentation) 增加样本多样性
- 尝试迁移学习 (Transfer Learning) 加速训练
完整的代码已上传至 GitHub 仓库(虚构地址):
github.com/username/mnist-fnn-demo
正文完
发表至: 未分类
近一天内
