共计 3651 个字符,预计需要花费 10 分钟才能阅读完成。
背景介绍
图像分类是计算机视觉的基础任务,LeNet- 5 作为早期卷积神经网络代表,特别适合处理手写数字等简单图像。MNIST 数据集包含 60,000 张 28×28 灰度手写数字图像,是验证模型效果的经典选择。

基础实现
环境准备
首先安装必要依赖:
!pip install torch torchvision matplotlib
数据加载与预处理
import torch
from torchvision import datasets, transforms
# 标准化处理能加速模型收敛
transform = transforms.Compose([transforms.ToTensor(),
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)
LeNet- 5 模型定义
import torch.nn as nn
import torch.nn.functional as F
class LeNet5(nn.Module):
def __init__(self):
super(LeNet5, self).__init__()
self.conv1 = nn.Conv2d(1, 6, 5, padding=2) # 保持 28x28 尺寸
self.pool = nn.MaxPool2d(2, 2)
self.conv2 = nn.Conv2d(6, 16, 5)
self.fc1 = nn.Linear(16*5*5, 120)
self.fc2 = nn.Linear(120, 84)
self.fc3 = nn.Linear(84, 10)
def forward(self, x):
x = self.pool(F.relu(self.conv1(x)))
x = self.pool(F.relu(self.conv2(x)))
x = x.view(-1, 16*5*5)
x = F.relu(self.fc1(x))
x = F.relu(self.fc2(x))
x = self.fc3(x)
return x
训练与评估
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.cross_entropy(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}')
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.cross_entropy(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 correct / len(test_loader.dataset)
# 训练流程
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = LeNet5().to(device)
optimizer = torch.optim.SGD(model.parameters(), lr=0.01, momentum=0.9)
for epoch in range(1, 11):
train(model, device, train_loader, optimizer, epoch)
test(model, device, test_loader)
初始 baseline 准确率通常在 98% 左右,这将成为我们优化的基准。
三重优化实验
维度 1:batch size 对比
batch size 影响内存使用和梯度更新频率:
- batch_size=32:更新频繁,可能震荡但泛化性好
- batch_size=64:平衡选择
- batch_size=128:内存占用大但训练快
实验数据:
| Batch Size | 训练时间 (秒) | 测试准确率 |
|---|---|---|
| 32 | 142 | 98.56% |
| 64 | 128 | 98.72% |
| 128 | 115 | 98.63% |
维度 2:优化器对比
修改 optimizer 定义部分即可测试不同优化器:
# SGD with momentum
optimizer = torch.optim.SGD(model.parameters(), lr=0.01, momentum=0.9)
# Adam
optimizer = torch.optim.Adam(model.parameters(), lr=0.001)
收敛曲线显示:
– Adam 初期收敛更快
– SGD 最终准确率略高 (98.7% vs 98.5%)
– 学习率对 Adam 更敏感
维度 3:激活函数对比
修改模型定义中的 ReLU 为 SiLU(也称为 Swish):
# 在 forward 方法中替换
x = self.pool(F.silu(self.conv1(x)))
实验结果:
– SiLU 准确率提升 0.3%
– 训练时间增加约 15%
– 更适合深层网络
结果分析
综合优化后的最佳组合:
| 参数 | 选择 | 准确率提升 |
|---|---|---|
| Batch Size | 64 | +0.1% |
| 优化器 | SGD | +0.2% |
| 激活函数 | SiLU | +0.3% |
可视化工具推荐:
from matplotlib import pyplot as plt
# 记录训练过程中的 loss 和 acc
loss_history = []
acc_history = []
# 在训练循环中添加记录
loss_history.append(loss.item())
acc_history.append(100. * correct / len(test_loader.dataset))
# 绘制曲线
plt.figure(figsize=(12,4))
plt.subplot(1,2,1)
plt.plot(loss_history)
plt.title('Training Loss')
plt.subplot(1,2,2)
plt.plot(acc_history)
plt.title('Test Accuracy')
plt.show()
避坑指南
- 学习率与 batch size
- 大 batch size 需相应增大学习率
-
推荐线性缩放规则:lr_new = lr_default * (batch_size / 64)
-
验证集注意事项
- 应从训练集划分 10-20% 作为验证集
-
确保各类别分布均衡
-
GPU 内存不足
- 减小 batch size
- 使用梯度累积技术
- 尝试混合精度训练
延伸思考
- 其他优化方向:
- 数据增强:旋转、平移、缩放
- 模型压缩:知识蒸馏、量化
-
正则化:Dropout、权重衰减
-
迁移学习尝试:
- 在 CIFAR-10 上微调 LeNet-5
- 注意调整输入通道数为 3
- 可能需要加深网络
完整代码已上传 GitHub 仓库,包含所有实验配置。建议读者尝试调整其他超参数如学习率调度器,观察模型表现变化。记住:没有放之四海而皆准的最优参数,理解原理比盲目调参更重要!
正文完
发表至: 未分类
近三天内
