CIFAR100过拟合实战:从数据增强到模型正则化的新手避坑指南

1次阅读
没有评论

共计 1907 个字符,预计需要花费 5 分钟才能阅读完成。

image.webp

为什么 CIFAR100 容易过拟合?

CIFAR100 是一个经典的图像分类数据集,包含 100 个细粒度类别(比如 20 种不同品种的狗),每张图片只有 32×32 像素。这么小的分辨率意味着:

CIFAR100 过拟合实战:从数据增强到模型正则化的新手避坑指南

  • 特征提取难度大:模糊的像素块需要模型具备更强的局部特征捕捉能力
  • 类别间差异小:不同品种的鸟类可能只有羽毛颜色的细微差别
  • 数据量有限:每类仅 500 张训练图片,是 CIFAR10 的 1 /10

我刚开始训练时,发现训练准确率很快达到 80%,但测试集卡在 45% 不动——这就是典型的过拟合:模型记住了训练数据的噪声,而非学习通用特征。

四招破解过拟合组合拳

1. 数据增强:低成本扩容训练集

transform_train = transforms.Compose([transforms.RandomCrop(32, padding=4),  # 随机裁剪保留主体
    transforms.RandomHorizontalFlip(),  # 水平镜像增加视角变化
    transforms.ColorJitter(brightness=0.2, contrast=0.2),  # 颜色扰动
    transforms.ToTensor(),
    transforms.Normalize((0.507, 0.487, 0.441), (0.267, 0.256, 0.276))
])

避坑点
– 避免使用 RandomRotation(30)这种大角度旋转,可能让飞机变成 ” 侧翻 ” 的无效样本
– 测试阶段必须用简单 transform(仅 Resize+Normalize)

2. Dropout:随机关闭神经元防死记硬背

在全连接层前加入:

self.dropout = nn.Dropout(0.5)  # 推荐 0.3-0.5 比例

黄金法则
– 越靠近输出层的 Dropout 比率可以越大
– 卷积层后通常用 BatchNorm 而非 Dropout

3. L2 正则化:限制权重不要过大

在优化器中加入 weight_decay 参数:

optimizer = torch.optim.Adam(model.parameters(), lr=0.001, weight_decay=1e-4)

经验值
– 1e- 4 到 1e- 2 之间效果较好
– 搭配学习率衰减效果更佳

4. Early Stopping:及时喊停训练

if val_acc > best_acc:
    best_acc = val_acc
    patience = 0  # 重置计数器
else:
    patience += 1
    if patience >= 5:  # 连续 5 轮未提升
        break  # 停止训练

完整代码结构

# 模型定义示例
class Net(nn.Module):
    def __init__(self):
        super().__init__()
        self.conv1 = nn.Conv2d(3, 64, 3, padding=1)
        self.bn1 = nn.BatchNorm2d(64)
        self.dropout = nn.Dropout(0.3)
        self.fc = nn.Linear(256, 100)

    def forward(self, x):
        x = F.relu(self.bn1(self.conv1(x)))
        x = self.dropout(x)
        return self.fc(x)

# 训练循环关键代码
for epoch in range(100):
    model.train()
    for inputs, labels in train_loader:
        outputs = model(inputs)
        loss = criterion(outputs, labels)

        optimizer.zero_grad()
        loss.backward()
        optimizer.step()

    # 验证阶段
    model.eval()
    with torch.no_grad():
        # ... 计算验证集准确率

效果对比

方法 测试准确率 过拟合程度
原始模型 43.2% 严重
+ 数据增强 51.7% 中等
全部技术组合 58.3% 轻微

避坑经验总结

  1. 数据增强不是越多越好:发现验证集准确率反而下降时,检查增强是否破坏了图像语义
  2. BatchNorm 和 Dropout 别打架 :训练时 model.train() 会同时激活两者,测试时 model.eval()会关闭 Dropout 但保持 BatchNorm
  3. 正则化要循序渐进:先用小 weight_decay(1e-5),观察 loss 曲线再调整

进阶方向

当上述方法效果饱和时,可以尝试:
模型剪枝:移除对结果影响小的神经元(适合部署到资源受限设备)
知识蒸馏:用大模型(教师模型)指导小模型(适合模型压缩)

建议在自己的数据集上:
1. 先复现本文基线方法
2. 记录不同组合的效果
3. 逐步加入更复杂的技术

过拟合就像考试死记硬背,而我们要培养的是模型的 ” 举一反三 ” 能力。希望这些实战经验能帮你少走弯路!

正文完
 0
评论(没有评论)