共计 1581 个字符,预计需要花费 4 分钟才能阅读完成。
CIFAR-10 数据集与挑战
CIFAR-10 是经典的彩色图像分类数据集,包含 10 个类别的 32×32 小尺寸图片,训练集 50000 张、测试集 10000 张。对初学者而言存在以下典型挑战:

- 小尺寸图像特征提取困难 :32×32 分辨率导致传统视觉特征(如 SIFT)失效
- 数据不平衡风险 :部分类别(如 ” 猫 ”)存在天然数据量偏少
- 过拟合高发 :模型参数量易远大于训练样本数
- 硬件资源限制 :深层网络训练显存占用大
模型架构对比
| 模型 | 参数量 | FLOPs | 核心创新点 |
|---|---|---|---|
| LeNet-5 | 60k | 0.4M | 基础卷积 - 池化堆叠 |
| VGG-16 | 138M | 15.5G | 3×3 小卷积核深层架构 |
| ResNet-18 | 11M | 1.8G | 残差连接解决梯度消失 |
LeNet- 5 实现要点
class LeNet5(nn.Module):
def __init__(self):
super().__init__()
self.conv1 = nn.Conv2d(3, 6, 5) # 输入通道 3,输出 6,卷积核 5x5
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 = torch.flatten(x, 1)
x = F.relu(self.fc1(x))
x = F.relu(self.fc2(x))
x = self.fc3(x)
return x
完整训练流程
数据预处理
transform_train = transforms.Compose([transforms.RandomCrop(32, padding=4),
transforms.RandomHorizontalFlip(),
transforms.ToTensor(),
transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616))
])
优化器配置
optimizer = torch.optim.SGD(model.parameters(), lr=0.1, momentum=0.9)
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=200)
实验结果对比
| 指标 | LeNet-5 | VGG-16 | ResNet-18 |
|---|---|---|---|
| 训练时间 /epoch | 45s | 210s | 180s |
| 测试准确率 | 68.2% | 72.8% | 76.5% |
| 显存占用 | 1.2GB | 5.8GB | 3.4GB |
关键问题解决方案
BatchNorm 小批量陷阱
- 当 batch_size<32 时,建议:
- 使用 GroupNorm 替代 BatchNorm
- 冻结 BN 层的 running_mean/var
梯度爆炸对策
- 梯度裁剪:
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=2.0) - 权重初始化:
for m in model.modules(): if isinstance(m, nn.Conv2d): nn.init.kaiming_normal_(m.weight, mode='fan_out')
验证集波动调参
- 检查学习率是否过大
- 增加验证集 batch_size
- 尝试 Label Smoothing 正则化
延伸思考
- 残差连接如何解决深层网络退化问题?
- 为什么 VGG 采用连续的 3 ×3 卷积?
- 针对 CIFAR-10 的特定优化策略有哪些?
建议读者结合原论文《Deep Residual Learning for Image Recognition》深入理解残差结构设计思想。
正文完
