共计 2246 个字符,预计需要花费 6 分钟才能阅读完成。
背景与评价指标
图像分类是计算机视觉的基础任务,常用评价指标包括:

- 准确率 (Accuracy):正确预测样本数 / 总样本数,直观但受类别分布影响
- F1-score:精确率与召回率的调和平均,适合类别不平衡场景
CIFAR-10 数据集包含:
- 10 类彩色图像(飞机、汽车、鸟等)
- 50k 训练 +10k 测试样本
- 32×32 低分辨率特性(需注意信息密度)
模型选型对比
| 模型 | 参数量 | FLOPs | 适用场景 |
|---|---|---|---|
| 全连接网络 | ~1.2M | ~2.4M | MNIST 级简单任务 |
| LeNet-5 | ~60k | ~0.4M | 早期手写数字识别 |
| ResNet-18 | ~11M | ~1.8G | 现代 CV 任务基准 |
残差结构优势:
- 解决深层网络梯度消失
- 恒等映射保留原始特征
- 参数量与性能的平衡点
核心实现
数据预处理
# torchvision 标准处理流程
train_transform = transforms.Compose([transforms.RandomCrop(32, padding=4), # 随机裁剪留白
transforms.RandomHorizontalFlip(), # 水平翻转增强
transforms.ToTensor(),
transforms.Normalize((0.4914, 0.4822, 0.4465),
(0.2023, 0.1994, 0.2010)) # CIFAR10 统计值
])
# 验证集无需增强
test_transform = transforms.Compose([transforms.ToTensor(),
transforms.Normalize((0.4914, 0.4822, 0.4465),
(0.2023, 0.1994, 0.2010))
])
残差块实现关键
class BasicBlock(nn.Module):
def __init__(self, in_planes, planes, stride=1):
super().__init__()
self.conv1 = nn.Conv2d(in_planes, planes,
kernel_size=3, stride=stride,
padding=1, bias=False)
self.bn1 = nn.BatchNorm2d(planes)
self.conv2 = nn.Conv2d(planes, planes,
kernel_size=3, stride=1,
padding=1, bias=False)
self.bn2 = nn.BatchNorm2d(planes)
# 下采样时调整维度
self.shortcut = nn.Sequential()
if stride != 1 or in_planes != planes:
self.shortcut = nn.Sequential(
nn.Conv2d(in_planes, planes,
kernel_size=1, stride=stride, bias=False),
nn.BatchNorm2d(planes)
)
def forward(self, x):
out = F.relu(self.bn1(self.conv1(x)))
out = self.bn2(self.conv2(out))
out += self.shortcut(x) # 残差连接核心
return F.relu(out)
三大优化策略
1. Batch Size 对比
| 批次大小 | 显存占用 | 迭代速度 | 最终准确率 |
|---|---|---|---|
| 32 | 2.1GB | 120it/s | 92.3% |
| 128 | 4.8GB | 380it/s | 91.7% |
经验 :小 batch 更易收敛但耗时,大 batch 需配合学习率升温 (warmup)
2. 优化器选择
# SGD with momentum
optimizer = torch.optim.SGD(model.parameters(),
lr=0.1,
momentum=0.9,
weight_decay=5e-4)
# Adam
optimizer = torch.optim.Adam(model.parameters(),
lr=0.001,
betas=(0.9, 0.999))
对比曲线显示 :
– Adam 初期收敛快但易震荡
– SGD+ 动量最终精度更高(需配合学习率衰减)
3. 数据增强效果
| 策略组合 | 训练准确率 | 测试准确率 | 过拟合程度 |
|---|---|---|---|
| 仅标准化 | 99.8% | 89.2% | 严重 |
| 裁剪 + 翻转 | 95.1% | 92.7% | 显著改善 |
| 裁剪 + 翻转 +Cutout | 93.6% | 93.9% | 最佳 |
避坑指南
学习率与 Batch Size
- 线性缩放原则:当 batch 扩大 k 倍,lr 也应扩大 k 倍
- 实际调整公式:
new_lr = base_lr * (new_bs / base_bs)
梯度爆炸处理
# 监控梯度范数
total_norm = torch.nn.utils.clip_grad_norm_(model.parameters(),
max_norm=2.0) # 阈值根据任务调整
if total_norm > 10:
print(f'梯度异常: {total_norm.item()}')
早停策略
# 当验证损失连续 5 轮不下降时停止
early_stopper = EarlyStopper(patience=5, min_delta=0.01)
for epoch in range(EPOCHS):
val_loss = validate()
if early_stopper.early_stop(val_loss):
break
部署思考题
- 类别不平衡处理 :
- 重加权交叉熵损失
- Focal Loss 应对难样本
-
过采样 / 欠采样策略
-
轻量化取舍 :
- 剪枝:保留重要连接(需微调)
- 量化:8bit 推理(精度损失约 1 -2%)
- 知识蒸馏:小模型模仿大模型
完整代码见:[GitHub 仓库链接]
正文完
发表至: 未分类
近两天内
