共计 1818 个字符,预计需要花费 5 分钟才能阅读完成。
CIFAR-10 数据集简介
CIFAR-10 是计算机视觉领域的经典基准数据集,包含 10 个类别的 6 万张 32×32 彩色图像(5 万训练 + 1 万测试)。其特点包括:

- 图像尺寸小(32×32),适合快速验证模型原型
- 类别均衡(每类 6000 张),避免数据倾斜问题
- 涵盖飞机、汽车、鸟类等常见物体,具有现实意义
方法对比
1. 传统 CNN(LeNet-5)
特点:
- 浅层网络(2 卷积 + 3 全连接)
- 使用 tanh 激活函数
- 最大池化进行下采样
优势:
- 结构简单,训练速度快
- 适合教学和基线测试
局限:
- 对深层特征提取能力有限
- 准确率天花板明显(约 65%)
2. ResNet-18
特点:
- 引入残差连接解决梯度消失
- 基础块包含两个 3 ×3 卷积
- 全局平均池化替代全连接
优势:
- 训练稳定性显著提升
- 准确率可达 90%+
局限:
- 参数量较大(约 11M)
3. EfficientNet-B0
特点:
- 复合系数统一缩放深度 / 宽度 / 分辨率
- 使用 MBConv 模块
- Swish 激活函数
优势:
- 参数效率高(约 5M)
- 准确率与 ResNet 相当
局限:
- 训练需要更多 epoch
4. ViT-Tiny
特点:
- 将图像分为 9 个 16×16 的 patch
- 多头自注意力机制
- 可学习的位置编码
优势:
- 长距离依赖建模能力强
- 无卷积归纳偏置
局限:
- 需要大量数据预训练
- 显存占用较高
完整 PyTorch 实现
数据增强
transform_train = transforms.Compose([transforms.RandomCrop(32, padding=4),
transforms.RandomHorizontalFlip(),
transforms.ToTensor(),
transforms.Normalize((0.4914, 0.4822, 0.4465), (0.247, 0.243, 0.261))
])
ViT 模型定义(片段)
class PatchEmbed(nn.Module):
"""将图像分割为 patch 并线性投影"""
def __init__(self, img_size=32, patch_size=16, in_chans=3, embed_dim=192):
super().__init__()
self.proj = nn.Conv2d(in_chans, embed_dim,
kernel_size=patch_size,
stride=patch_size)
def forward(self, x):
x = self.proj(x).flatten(2).transpose(1, 2)
return x
训练循环
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=200)
for epoch in range(epochs):
model.train()
for data, target in train_loader:
optimizer.zero_grad()
with torch.cuda.amp.autocast(): # 混合精度
output = model(data)
loss = criterion(output, target)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
scheduler.step()
性能对比(Tesla T4)
| 模型 | 准确率 | 训练时间 /epoch | 显存占用 |
|---|---|---|---|
| LeNet-5 | 68.2% | 45s | 1.2GB |
| ResNet-18 | 93.5% | 2.3min | 3.8GB |
| EfficientNet | 92.1% | 3.1min | 2.9GB |
| ViT-Tiny | 88.7% | 4.5min | 5.4GB |
最佳实践
小图像处理技巧
- 使用更大的 patch 尺寸(如 16×16)
- 减少下采样次数
- 适当增加通道数
防过拟合策略
- 添加 CutMix 数据增强
- 使用 Label Smoothing
- 早停法 + 模型保存
混合精度训练
- 初始化 scaler
scaler = torch.cuda.amp.GradScaler() - 包装前向计算
with torch.cuda.amp.autocast(): outputs = model(inputs) - 缩放梯度
scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()
开放问题
- 如何在只有 1000 样本的情况下提升 ViT 性能?
-
可能的方案:知识蒸馏、迁移学习
-
移动端部署时如何压缩模型?
- 量化方案:动态 8bit 量化
- 剪枝策略:基于重要性的通道剪枝
正文完
