共计 3190 个字符,预计需要花费 8 分钟才能阅读完成。
MNIST 数据集与图像分类基础
MNIST 是深度学习领域的 ”Hello World” 数据集,包含 60,000 张训练图片和 10,000 张测试图片,每张都是 28×28 像素的灰度手写数字(0-9)。这个看似简单的任务背后有几个关键技术挑战:

- 图像预处理:需要将像素值归一化到 0 - 1 范围
- 特征提取:如何从原始像素中学习有效特征
- 模型泛化:防止记住训练样本但无法识别新样本
网络结构选型指南
- 前馈神经网络(FNN)
- 最基础的全连接网络
- 将 28×28 图像展平为 784 维向量
- 适合理解神经网络基本原理
-
参数量大,准确率约 98%
-
卷积神经网络(CNN)
- 通过卷积核自动提取空间特征
- 保留图像二维结构信息
- 参数量少,准确率可达 99% 以上
-
适合视觉类任务
-
循环神经网络(RNN)
- 将图像按行 / 列序列处理
- 理论上可以捕捉笔顺信息
- 实际效果通常不如 CNN
- 适合演示 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)
性能优化技巧
- 批处理大小
- GPU 显存充足:增大 batch size(如 256)加速训练
-
显存有限:减小 batch size(如 32)配合梯度累积
-
学习率调整
- 初始尝试:1e- 3 到 1e-4
-
使用 ReduceLROnPlateau 自动调整
scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, 'min') -
数据增强
- 添加随机旋转 / 平移
transforms.RandomRotation(10), transforms.RandomAffine(0, translate=(0.1,0.1))
常见问题解决方案
- 维度不匹配错误
- 检查各层输入 / 输出形状
-
使用
print(x.shape)调试 -
GPU 内存不足
- 减小 batch size
- 使用
torch.cuda.empty_cache() -
混合精度训练
trainer = pl.Trainer(precision=16) -
过拟合识别
- 训练误差持续下降但测试误差上升
- 解决方案:
- 增加 Dropout 层
- 添加 L2 正则化
- 早停(EarlyStopping)
模型保存与加载
# 保存
torch.save(model.state_dict(), 'mnist_cnn.pt')
# 加载
model = CNN()
model.load_state_dict(torch.load('mnist_cnn.pt'))
model.eval()
延伸思考
- 模型部署:如何用 Flask 将训练好的模型封装为 REST API?
- 非平衡数据:当某些数字样本过少时,该如何调整损失函数?
- 实时识别:如何扩展本项目实现摄像头实时手写数字识别?
通过这个完整的实践流程,相信你已经掌握了 PyTorch 实现图像分类的核心方法。建议尝试调整网络结构超参数,观察对模型性能的影响,这是提升深度学习实战能力的最佳途径。
正文完
发表至: 未分类
近一天内
