9.数据增强与自定义 Transform
第 9 章 — 数据增强与自定义 Transform
第 8 章结束时我们撞上了一面墙:训练集 acc 已经冲到 99%,验证集却只有 72%——模型把训练集"背"下来了。本章引入数据增强:训练时给每张图生成大量"变体",逼模型去学"通用的规律"而不是"特定的像素"。同时把增强做成
data.augmentation配置段,让"关增强 / 开增强"从改代码变成改一行 YAML。
一、本章要解决的问题
- 之前(第 8 章结尾):训练 acc 99%、验证 acc 72%,训练-验证 gap 高达 27 个百分点,这是教科书级的过拟合。
- 现在:用数据增强给训练集"注水",扩大样本的多样性,压缩 gap、提升验证/测试准确率。
- 判断标准:开启增强后 val_acc 明显高于关闭增强;并且"开不开增强"只改配置文件、不动任何代码。
二、核心概念速览
下面 7 个概念是本章的全部"生词"。先花 5 分钟读完,再看代码会轻松很多。
1. 过拟合的表现:背题 vs 理解
想象一个备考的学生:把题库答案背得滚瓜烂熟,碰到一模一样的原题拿满分;可题目换一张配图、换一种问法,立刻就不会了。这就是"背题"和"理解"的区别。神经网络过拟合是同一回事——它把训练集的图"背"下来了,第 8 章结尾训练 99%、验证 72%,中间差的 27 个点就是"背题但没理解"的铁证。
2. 数据增强:给同一张图拍各种"变体照片"
数据增强就是在训练时把每张图随机"做手脚":平移几个像素、左右翻转、调亮调暗、改一点颜色……相当于给同一个物体拍了一堆不同角度、不同光线下的"变体照片"。模型见过的样本从 5 万张变成"无穷多张",它没法再靠"这张图我背过"蒙混过关,只能去学真正稳定可靠的规律。注意:这是训练技巧,跟数据本身的"质量"无关。
3. RandomCrop(padding=4):先放大再裁回,模拟平移
CIFAR-10 原图是 32x32。RandomCrop(32, padding=4) 先给图片四周各补 4 像素(默认补 0),图片变成 40x40;再随机从里面裁一个 32x32 的窗口。因为每次裁的窗口位置都不同,同一张图每次进模型时"主体"都偏移了几个像素——等于免费的平移增强,逼模型学会"物体不在正中间也能认出来"。
4. RandomHorizontalFlip:随机左右翻转
以 50% 概率把图片水平镜像。对 CIFAR-10 这类自然图片为什么安全?因为"船在左边还是右边"跟"它是不是船"没有关系。翻转让有效训练样本翻倍,还让模型对左右方向不敏感。但要注意适用范围:数字"6"、字母"b"这类左右翻转会变含义的图不能这么干。
5. ColorJitter:亮度 / 对比度 / 饱和度 / 色相
ColorJitter(brightness, contrast, saturation, hue) 四个参数分别控制亮度、对比度、饱和度、色相的随机扰动幅度,值越大颜色变化越夸张。特别要注意:hue(色相)的取值范围是 [-0.5, 0.5],单位是"色相环转的圈数"(0.5 = 转半圈);配置里给 0.1 表示最多转 10% 的色相。给太大会让"红苹果"变"绿苹果",分类反而更难。
6. 增强只加在训练集:验证/测试集必须保持"干净"
验证集和测试集的使命是"模拟真实世界的考试",而真实世界里不会因为你考试时换了个姿势就给加分。所以增强永远只加在训练集,验证/测试集只用 ToTensor + Normalize。如果测试集也做增强,指标就不代表真实的泛化能力,等于"考试时偷看答案"。后面 build_transforms(cfg, train=False) 里的 train 开关就是为此设计的。
7. Transform pipeline 的组合顺序:ToTensor 必须在 Normalize 之前
transforms.Compose 按列表顺序依次处理图片。ToTensor 把 PIL 图片(0~255 的整数)转成 0~1 的 float 张量;Normalize 是张量数学运算(减均值、除方差),它只认识张量。顺序反了,Normalize 拿到 PIL 图片会直接报错。另外,RandomCrop / RandomHorizontalFlip / ColorJitter 这些随机增强也要放在 ToTensor 之前(它们操作的是 PIL 图片),而 Normalize 永远放最后。
三、解决思路
- 复用现成算子:三个随机增强(RandomCrop、RandomHorizontalFlip、ColorJitter)全是
torchvision.transforms自带的,不自己写随机逻辑——官方实现经过大量测试,边界情况都处理好了。 - 配置化:在
config/config.yaml的data段下新增augmentation子段,三个开关各自独立;把color_jitter设为null即关闭颜色抖动,把padding设为0即不裁剪。 - 训练/评估分家:写一个
build_transforms(cfg, train=False)函数,train=True时拼接增强,train=False时只保留ToTensor + Normalize。从代码结构上"逼"你区分训练变换和评估变换,杜绝"测试集也被增强"的经典错误。 - 对比实验:分别以"关闭增强"和"全增强"跑两组实验,用 val_acc 和 gap 的大小说话——增强有没有用,数据说了算。
trade-off:增强是用"多样性"换"保真度"。太强的增强(比如 hue 给 0.5)会让训练样本失真到"这根本不是原来的物体",反而更难收敛;增强只产生"变体",不会带来全新的真实样本。所以增强强度要靠实验调,不是越大越好。
四、代码变更
相对第 8 章的改动点:
# config/config.yaml:data 段新增 augmentation 子段
data:
root: ./data
batch_size: 64
num_workers: 2
val_ratio: 0.1
+ augmentation:
+ random_crop_padding: 4 # 四周各补 4 像素后随机裁剪 32x32
+ horizontal_flip: true # 50% 概率左右翻转
+ color_jitter: [0.2, 0.2, 0.2, 0.1] # 亮度/对比度/饱和度/色相扰动幅度
# train.py:新增两个依赖(进度条 + 彩色日志)
+ from tqdm import tqdm
+ from rich.logging import RichHandler
# train.py:日志输出改用 RichHandler(终端更好看,文件日志不变)
def setup_logger(log_dir):
...
- console = logging.StreamHandler()
- console.setFormatter(fmt)
+ console = RichHandler(rich_tracebacks=True, markup=True)
+ console.setLevel(logging.INFO)
logger.addHandler(console)
...
# train.py:固定的 transform 改为 build_transforms 函数,训练/评估分家
- transform = transforms.Compose([
- transforms.ToTensor(),
- transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)),
- ])
- train_dataset = datasets.CIFAR10(..., transform=transform)
- test_dataset = datasets.CIFAR10(..., transform=transform)
+ # 数据增强:训练集随机扰动,验证/测试集只做归一化
+ def build_transforms(cfg, train=False):
+ aug = cfg.data.augmentation
+ transforms_list = []
+ if train:
+ if aug.random_crop_padding > 0:
+ transforms_list.append(
+ transforms.RandomCrop(32, padding=aug.random_crop_padding))
+ if aug.horizontal_flip:
+ transforms_list.append(transforms.RandomHorizontalFlip())
+ if aug.color_jitter is not None:
+ transforms_list.append(transforms.ColorJitter(*aug.color_jitter))
+ transforms_list += [
+ transforms.ToTensor(),
+ transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)),
+ ]
+ return transforms.Compose(transforms_list)
+
+ train_dataset = datasets.CIFAR10(..., transform=build_transforms(cfg, train=True))
+ test_dataset = datasets.CIFAR10(..., transform=build_transforms(cfg, train=False))
# train.py:训练/验证循环加上 tqdm 进度条(循环体逻辑不变)
- for images, labels in loader:
+ for images, labels in tqdm(loader, desc="Train", leave=False):其余部分(Config 类、EarlyStopping、checkpoint、resume、调度器)与第 8 章完全一致,改动只在"数据怎么喂"这一条线上。
五、完整代码
创建 config/config.yaml:
# CIFAR-10 图像分类实验配置
device: auto
data:
root: ./data
batch_size: 64
num_workers: 2
val_ratio: 0.1
augmentation:
random_crop_padding: 4
horizontal_flip: true
color_jitter: [0.2, 0.2, 0.2, 0.1]
model:
num_classes: 10
train:
epochs: 30
lr: 0.001
lr_scheduler: step
lr_step_size: 15
lr_gamma: 0.1
patience: 7
checkpoint:
dir: ./checkpoints
log:
dir: ./runs/cifar10创建 train.py:
import os
import logging
import argparse
import yaml
import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader, random_split
from torch.utils.tensorboard import SummaryWriter
from torchvision import datasets, transforms
from collections import defaultdict
from tqdm import tqdm
from rich.logging import RichHandler
def parse_args():
parser = argparse.ArgumentParser(description="CIFAR-10 图像分类训练")
parser.add_argument("--config", type=str, default="config/config.yaml",
help="YAML 配置文件路径")
parser.add_argument("--lr", type=float, default=None, help="覆盖配置文件中的 train.lr")
parser.add_argument("--epochs", type=int, default=None, help="覆盖配置文件中的 train.epochs")
parser.add_argument("--batch-size", type=int, default=None,
help="覆盖配置文件中的 data.batch_size")
parser.add_argument("--resume", type=str, default=None,
help="从指定 checkpoint 恢复训练")
return parser.parse_args()
args = parse_args()
class Config(dict):
def __getattr__(self, key):
try:
return self[key]
except KeyError as e:
raise AttributeError(key) from e
def _to_config(obj):
if isinstance(obj, dict):
return Config({k: _to_config(v) for k, v in obj.items()})
if isinstance(obj, list):
return [_to_config(v) for v in obj]
return obj
def load_config(path):
with open(path, "r", encoding="utf-8") as f:
return _to_config(yaml.safe_load(f))
cfg = load_config(args.config)
if args.lr is not None:
cfg.train.lr = args.lr
if args.epochs is not None:
cfg.train.epochs = args.epochs
if args.batch_size is not None:
cfg.data.batch_size = args.batch_size
def setup_logger(log_dir):
os.makedirs(log_dir, exist_ok=True)
logger = logging.getLogger("train")
logger.handlers.clear()
logger.setLevel(logging.INFO)
console = RichHandler(rich_tracebacks=True, markup=True)
console.setLevel(logging.INFO)
logger.addHandler(console)
file_handler = logging.FileHandler(os.path.join(log_dir, "train.log"),
encoding="utf-8")
file_handler.setFormatter(logging.Formatter(
"%(asctime)s | %(levelname)s | %(message)s", datefmt="%Y-%m-%d %H:%M:%S"))
logger.addHandler(file_handler)
return logger
logger = setup_logger(cfg.log.dir)
writer = SummaryWriter(log_dir=cfg.log.dir)
if cfg.device == "auto":
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
else:
device = torch.device(cfg.device)
logger.info(f"Using device: {device}")
class SimpleCNN(nn.Module):
def __init__(self, num_classes=10):
super().__init__()
self.features = nn.Sequential(
nn.Conv2d(3, 32, kernel_size=3, padding=1),
nn.BatchNorm2d(32),
nn.ReLU(inplace=True),
nn.MaxPool2d(2),
nn.Conv2d(32, 64, kernel_size=3, padding=1),
nn.BatchNorm2d(64),
nn.ReLU(inplace=True),
nn.MaxPool2d(2),
nn.Conv2d(64, 128, kernel_size=3, padding=1),
nn.BatchNorm2d(128),
nn.ReLU(inplace=True),
nn.MaxPool2d(2),
)
self.classifier = nn.Sequential(
nn.AdaptiveAvgPool2d(1),
nn.Flatten(),
nn.Linear(128, num_classes),
)
def forward(self, x):
return self.classifier(self.features(x))
# 数据增强:训练集随机扰动,验证/测试集只做归一化
def build_transforms(cfg, train=False):
aug = cfg.data.augmentation
transforms_list = []
if train:
if aug.random_crop_padding > 0:
transforms_list.append(
transforms.RandomCrop(32, padding=aug.random_crop_padding))
if aug.horizontal_flip:
transforms_list.append(transforms.RandomHorizontalFlip())
if aug.color_jitter is not None:
transforms_list.append(transforms.ColorJitter(*aug.color_jitter))
transforms_list += [
transforms.ToTensor(),
transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)),
]
return transforms.Compose(transforms_list)
train_dataset = datasets.CIFAR10(root=cfg.data.root, train=True, download=True,
transform=build_transforms(cfg, train=True))
test_dataset = datasets.CIFAR10(root=cfg.data.root, train=False, download=True,
transform=build_transforms(cfg, train=False))
val_size = int(len(train_dataset) * cfg.data.val_ratio)
train_dataset, val_dataset = random_split(
train_dataset, [len(train_dataset) - val_size, val_size]
)
train_loader = DataLoader(train_dataset, batch_size=cfg.data.batch_size,
shuffle=True, num_workers=cfg.data.num_workers)
val_loader = DataLoader(val_dataset, batch_size=cfg.data.batch_size,
shuffle=False, num_workers=cfg.data.num_workers)
test_loader = DataLoader(test_dataset, batch_size=cfg.data.batch_size,
shuffle=False, num_workers=cfg.data.num_workers)
model = SimpleCNN(num_classes=cfg.model.num_classes).to(device)
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=cfg.train.lr)
if cfg.train.lr_scheduler == "step":
scheduler = optim.lr_scheduler.StepLR(
optimizer, step_size=cfg.train.lr_step_size, gamma=cfg.train.lr_gamma)
elif cfg.train.lr_scheduler == "cosine":
scheduler = optim.lr_scheduler.CosineAnnealingLR(
optimizer, T_max=cfg.train.epochs)
else:
scheduler = None
def train_one_epoch(model, loader, criterion, optimizer, device):
model.train()
total_loss, correct, total = 0.0, 0, 0
for images, labels in tqdm(loader, desc="Train", leave=False):
images, labels = images.to(device), labels.to(device)
optimizer.zero_grad()
outputs = model(images)
loss = criterion(outputs, labels)
loss.backward()
optimizer.step()
total_loss += loss.item() * images.size(0)
correct += (outputs.argmax(dim=1) == labels).sum().item()
total += images.size(0)
return total_loss / total, correct / total
def validate(model, loader, criterion, device):
model.eval()
total_loss, correct, total = 0.0, 0, 0
with torch.no_grad():
for images, labels in tqdm(loader, desc="Val", leave=False):
images, labels = images.to(device), labels.to(device)
outputs = model(images)
loss = criterion(outputs, labels)
total_loss += loss.item() * images.size(0)
correct += (outputs.argmax(dim=1) == labels).sum().item()
total += images.size(0)
return total_loss / total, correct / total
def save_checkpoint(state, path):
os.makedirs(os.path.dirname(path), exist_ok=True)
torch.save(state, path)
logger.info(f"[Checkpoint] 已保存到 {path}")
class EarlyStopping:
def __init__(self, patience=7, delta=0.0, verbose=False):
self.patience = patience
self.delta = delta
self.verbose = verbose
self.best_score = None
self.counter = 0
self.early_stop = False
def __call__(self, val_acc):
score = val_acc
if self.best_score is None:
self.best_score = score
elif score < self.best_score + self.delta:
self.counter += 1
if self.verbose:
logger.info(f"[EarlyStopping] 连续 {self.counter}/{self.patience} 轮未提升")
if self.counter >= self.patience:
self.early_stop = True
else:
self.best_score = score
self.counter = 0
metrics = defaultdict(list)
start_epoch, best_acc = 1, 0.0
if args.resume:
ckpt = torch.load(args.resume, map_location=device, weights_only=False)
model.load_state_dict(ckpt["model_state"])
optimizer.load_state_dict(ckpt["optimizer_state"])
if scheduler is not None and ckpt.get("scheduler_state") is not None:
scheduler.load_state_dict(ckpt["scheduler_state"])
start_epoch = ckpt["epoch"] + 1
best_acc = ckpt.get("best_acc", 0.0)
metrics = defaultdict(list, ckpt.get("metrics", {}))
logger.info(f"[Resume] 从 epoch {ckpt['epoch']} 恢复,历史最优 val_acc={best_acc:.4f}")
early_stopping = EarlyStopping(patience=cfg.train.patience, verbose=True)
logger.info("=" * 60)
logger.info("最终生效配置:")
logger.info(yaml.safe_dump(cfg, allow_unicode=True, sort_keys=False))
logger.info("=" * 60)
for epoch in range(start_epoch, cfg.train.epochs + 1):
train_loss, train_acc = train_one_epoch(model, train_loader, criterion, optimizer, device)
val_loss, val_acc = validate(model, val_loader, criterion, device)
if scheduler is not None:
scheduler.step()
current_lr = optimizer.param_groups[0]["lr"]
metrics["train_loss"].append(train_loss)
metrics["train_acc"].append(train_acc)
metrics["val_loss"].append(val_loss)
metrics["val_acc"].append(val_acc)
logger.info(f"Epoch {epoch:02d}/{cfg.train.epochs} | "
f"Train loss {train_loss:.4f} acc {train_acc:.4f} | "
f"Val loss {val_loss:.4f} acc {val_acc:.4f} | "
f"lr {current_lr:.2e}")
writer.add_scalar("train/loss", train_loss, epoch)
writer.add_scalar("train/acc", train_acc, epoch)
writer.add_scalar("val/loss", val_loss, epoch)
writer.add_scalar("val/acc", val_acc, epoch)
writer.add_scalar("lr", current_lr, epoch)
state = {
"model_state": model.state_dict(),
"optimizer_state": optimizer.state_dict(),
"scheduler_state": scheduler.state_dict() if scheduler else None,
"epoch": epoch,
"best_acc": best_acc,
"metrics": dict(metrics),
}
save_checkpoint(state, os.path.join(cfg.checkpoint.dir, "last.pth"))
if val_acc > best_acc:
best_acc = val_acc
state["best_acc"] = best_acc
save_checkpoint(state, os.path.join(cfg.checkpoint.dir, "best.pth"))
early_stopping(val_acc)
if early_stopping.early_stop:
logger.info(f"[EarlyStopping] 连续 {cfg.train.patience} 轮无提升,在 epoch {epoch} 停止")
break
test_loss, test_acc = validate(model, test_loader, criterion, device)
logger.info(f"[Test] loss {test_loss:.4f} acc {test_acc:.4f}")
writer.close()build_transforms 逐行解读(本章唯一的新函数):
aug = cfg.data.augmentation:拿到配置里的增强参数,后面三个if全是"配置说了算"——这就是第 8 章配置化的威力,改代码变成改 YAML。if train::训练时才拼随机增强;train=False时transforms_list直接是空的,最后只接上ToTensor + Normalize,验证/测试集保持"干净"。RandomCrop(32, padding=aug.random_crop_padding):注意裁剪目标尺寸写死 32,因为 CIFAR-10 原图就是 32x32;padding=4先放大成 40x40 再随机裁回。if aug.color_jitter is not None:配置里给null就跳过颜色抖动;transforms.ColorJitter(*aug.color_jitter)的*把列表[0.2, 0.2, 0.2, 0.1]展开成四个位置参数(对应 亮度/对比度/饱和度/色相)。- 顺序固定为:随机增强 → ToTensor → Normalize。随机增强必须在
ToTensor之前(操作 PIL 图片),Normalize永远在最后(只认张量)。
运行示例:
# 安装新增依赖
pip install tqdm rich
# 全增强(默认配置):训练/验证/test 自动分家
python train.py
# 临时调参仍可用命令行覆盖
python train.py --lr 1e-4 --epochs 50六、对比实验:关闭增强 vs 全增强
增强有没有用,空口无凭。我们把配置里的 augmentation 段调成"全关",与默认的"全开"各跑一遍:
# config/no_aug.yaml:在默认配置基础上,只改 augmentation 段
augmentation:
random_crop_padding: 0 # 0 = 不裁剪
horizontal_flip: false # 关闭翻转
color_jitter: null # 关闭颜色抖动# 实验 A:关闭增强
python train.py --config config/no_aug.yaml
# 实验 B:全增强(默认配置,即 config/config.yaml)
python train.py两组实验(30 epoch、早停 patience=7,其余完全相同)的典型结果:
| 实验 | 训练 acc(末轮) | 验证 acc(最优) | 训练-验证 gap |
|---|---|---|---|
| 关闭增强 | ~99% | ~72% | ~27 个点 |
| 全增强 | ~88% | ~78% | ~10 个点 |
读这张表,最反直觉的一点是:开启增强后训练 acc 反而从 99% 降到了 88%。这不是退步,恰恰说明模型不再"背题"了——训练集变难(每张图都随机变形),模型没法靠记忆刷高分,只能去学泛化规律;而验证 acc 从 72% 涨到 78%,gap 从 27 个点缩到 10 个点,这才是我们真正要的。请记住这句话:训练 acc 高不是目标,验证 acc 高才是。
七、本章小结
- 学到了什么
- 数据增强 = 训练时给每张图生成随机"变体",扩大有效样本多样性,是压制过拟合最便宜有效的手段。
- 三大基础增强:
RandomCrop(padding=4)(先放大再裁回、模拟平移)、RandomHorizontalFlip(50% 翻转)、ColorJitter(亮度/对比度/饱和度/色相四参数,hue 范围 [-0.5, 0.5])。 - 铁律两条:增强只加在训练集,验证/测试集必须干净;Transform pipeline 顺序固定为 随机增强 → ToTensor → Normalize(ToTensor 必须在 Normalize 之前)。
- 用
build_transforms(cfg, train)一个函数 + 一个train开关,从代码结构上杜绝"测试集被增强"。
- 常见坑
- 把增强也套在验证/测试集上:val/test 指标被"美化",不再代表真实泛化能力。
ToTensor与Normalize顺序颠倒:Normalize收到 PIL 图片直接报错(它只认张量)。- 给不适用于翻转的数据(数字、文字)开
horizontal_flip:类别语义被破坏。 hue超出 [-0.5, 0.5] 会直接抛ValueError,不是"自动截断"。- 增强太强导致训练不收敛:val_acc 不升反降时,先想到"增强过头了"而不是"模型坏了"。
- 下一章预告:数据增强解决的是"模型记不住",而显存不够、训练太慢是另一道坎——第 10 章用 PyTorch 混合精度(AMP)在不动模型结构的前提下省显存、提速。
八、动手练习
- 调裁剪强度:把
random_crop_padding从 4 改成 8 再改成 0,分别记录 val_acc。想想为什么 padding 太大(主体被裁掉一部分)反而变差。 - 调色相:把
color_jitter的 hue 从 0.1 改成 0.3,观察 val_acc 变化,并解释这与"颜色失真"的关系。 - 反序实验:把
build_transforms里的ToTensor和Normalize对调顺序跑一次,记录报错信息,确认自己理解了"顺序为什么必须这样"。 - 拆开关:保持其他增强不变,只关掉
horizontal_flip,对比开启时 val_acc 的差距——验证"翻转对 CIFAR-10 有没有贡献"。 - 脑内推导:不运行代码,回答——
RandomCrop(32, padding=4)后图片尺寸是多少?为什么训练集能因此"看见"平移后的物体?
