一、引言:让模型长见识,让成果留下来
在前两篇博客中,我们完成了从自定义数据集到CNN模型训练的完整流程。但如果你仔细回顾,会发现一个潜在的问题:训练数据太单一。
模型只见过正着摆放的、亮度固定的、同一角度的物品。一旦测试图片稍有旋转、翻转或颜色变化,模型就可能傻眼——这就是所谓的过拟合。
解决这个问题的利器,就是数据增强(Data Augmentation)。它通过对训练图片进行随机变换(旋转、翻转、调色等),人为地制造出更多样化的训练样本,让模型学会忽略这些无关变化,专注于真正的类别特征。
与此同时,训练了若干轮之后,我们得到了一个不错的模型——但如果没有保存,下次就得从头再来。模型保存让训练成果得以持久化,随时可以加载使用。
本篇博客将围绕这两大主题展开,基于完整代码讲解数据增强、标准化、最优模型保存三大核心知识点。
二、数据增强
2.1 训练集 vs 验证集:两套不同的变换策略
代码中最醒目的设计,是定义了两套变换流程:
data_transforms = {
'trainda': transforms.Compose([
transforms.RandomRotation(45),
transforms.CenterCrop(256),
transforms.RandomHorizontalFlip(p=0.5),
transforms.RandomVerticalFlip(p=0.5),
transforms.ColorJitter(brightness=0.2, contrast=0.1, saturation=0.1, hue=0.1),
transforms.RandomGrayscale(p=0.1),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
]),
'valid': transforms.Compose([
transforms.Resize([256, 256]),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
]),
}
核心原则:
-
训练集:使用随机变换(数据增强),让每个epoch看到的图片都略有不同,提升泛化能力。
-
验证集/测试集:只做必要的尺寸统一和标准化,不能加入随机性,否则评估结果不稳定。
2.2 常用数据增强方法详解
(1)RandomRotation——随机旋转
transforms.RandomRotation(45)
在 -45°到45° 之间随机旋转图片。这模拟了拍摄角度不同的情况,让模型学会识别旋转后的物体。
(2)CenterCrop——中心裁剪
transforms.CenterCrop(256)
从图像中心裁剪出256×256的区域。配合RandomRotation使用,可以裁掉旋转后产生的黑边,保证输入尺寸一致。
(3)RandomHorizontalFlip / RandomVerticalFlip——随机翻转
transforms.RandomHorizontalFlip(p=0.5) # 水平翻转,50%概率
transforms.RandomVerticalFlip(p=0.5) # 垂直翻转,50%概率
以指定概率对图片进行翻转。水平翻转适合大多数场景(如动物、车辆),垂直翻转则要谨慎使用(对于人脸等有方向性的物体可能不合适)。
(4)ColorJitter——颜色抖动
transforms.ColorJitter(brightness=0.2, contrast=0.1, saturation=0.1, hue=0.1)
随机调整图像的亮度、对比度、饱和度、色相。这模拟了不同光照条件下的拍摄效果,提升模型对光照变化的鲁棒性。
| brightness | 亮度 | 0.2表示在[0.8, 1.2]倍之间随机调整 |
| contrast | 对比度 | 同上 |
| saturation | 饱和度 | 同上 |
| hue | 色相 | 0.1表示在[-0.1, 0.1]之间偏移 |
(5)RandomGrayscale——随机灰度化
transforms.RandomGrayscale(p=0.1)
以10%的概率将彩色图片转为灰度图(R=G=B)。这强制模型不依赖颜色信息,学习更本质的形状特征。
2.3 ToTensor 与 Normalize——标准化的两步
ToTensor:从PIL到张量
transforms.ToTensor()
作用:
-
将PIL图像或NumPy数组转为PyTorch张量
-
将像素值从 0-255 缩放到 0-1
-
将通道维度从 HWC 转为 CHW(PyTorch要求)
Normalize:标准化到标准正态分布
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
计算方式:

为什么用这组特定的均值和标准差? 它们是 ImageNet数据集上统计出来的RGB三通道均值和标准差。由于大多数预训练模型都是在ImageNet上训练的,使用相同的标准化参数可以保持数据分布一致。
标准化的意义:
-
让数据分布接近标准正态分布,加速梯度下降收敛
-
消除不同通道之间的量纲差异
-
是迁移学习中使用预训练模型时的必要步骤
三、自定义数据集回顾
数据集类与上一篇博客一致:
class food_dataset(Dataset):
def __init__(self, file_path, transform=None):
self.file_path = file_path
self.imgs = []
self.labels = []
self.transform = transform
with open(self.file_path) as f:
samples = [x.strip().split(' ') for x in f.readlines()]
for img_path, label in samples:
self.imgs.append(img_path)
self.labels.append(label)
def __len__(self):
return len(self.imgs)
def __getitem__(self, idx):
image = Image.open(self.imgs[idx])
if self.transform:
image = self.transform(image)
label = torch.from_numpy(np.array(self.labels[idx], dtype=np.int64))
return image, label
然后分别用训练变换和验证变换创建数据集:
training_data = food_dataset(file_path='./train.txt', transform=data_transforms['trainda'])
test_data = food_dataset(file_path='./test.txt', transform=data_transforms['valid'])
train_dataloader = DataLoader(training_data, batch_size=64, shuffle=True)
test_dataloader = DataLoader(test_data, batch_size=64, shuffle=True)
四、CNN模型结构
模型与上一篇相同,针对 3×256×256 彩色输入,输出20个类别:
class CNN(nn.Module):
def __init__(self):
super(CNN, self).__init__()
self.conv1 = nn.Sequential(
nn.Conv2d(3, 16, 5, 1, 2),
nn.ReLU(),
nn.MaxPool2d(2),
)
self.conv2 = nn.Sequential(
nn.Conv2d(16, 32, 5, 1, 2),
nn.ReLU(),
nn.Conv2d(32, 64, 5, 1, 2),
nn.ReLU(),
nn.MaxPool2d(2),
)
self.conv3 = nn.Sequential(
nn.Conv2d(64, 128, 5, 1, 2),
nn.ReLU(),
)
self.out = nn.Linear(128*64*64, 20)
def forward(self, x):
x = self.conv1(x)
x = self.conv2(x)
x = self.conv3(x)
x = x.view(x.size(0), -1)
output = self.out(x)
return output
尺寸变化:
-
输入:3×256×256
-
conv1后:16×128×128
-
conv2后:64×64×64
-
conv3后:128×64×64
-
展平:128×64×64 = 524288 维
-
输出:20类
五、训练函数
def train(dataloader, model, loss_fn, optimizer):
model.train()
batch_size_num = 1
for x, y in dataloader:
x, y = x.to(device), y.to(device)
pred = model.forward(x)
loss = loss_fn(pred, y)
optimizer.zero_grad()
loss.backward()
optimizer.step()
loss_value = loss.item()
if batch_size_num % 1 == 0:
print(f"loss:{loss_value:7f} [number:{batch_size_num}]")
batch_size_num += 1
训练过程与之前一致:前向传播→计算损失→梯度清零→反向传播→更新参数。
六、模型保存
这是本篇博客的重点。测试函数中,当模型准确率创新高时,会保存模型:
best_acc = 0
def test(dataloader, model, loss_fn):
global best_acc
size = len(dataloader.dataset)
num_batches = len(dataloader)
model.eval()
test_loss, correct = 0, 0
with torch.no_grad():
for X, y in dataloader:
X, y = X.to(device), y.to(device)
pred = model.forward(X)
test_loss += loss_fn(pred, y).item()
correct += (pred.argmax(1) == y).type(torch.float).sum().item()
test_loss /= num_batches
correct /= size
print(f"Test result: \\n Accuracy: {(100*correct)}%, Avg loss: {test_loss}")
# 保存最优模型
if correct > best_acc:
best_acc = correct
print(model.state_dict().keys())
torch.save(model.state_dict(), f"xxxxxx.pth")
script_model = torch.jit.script(model)
torch.jit.save(script_model, f"xxxxxxx.pth")
6.1 两种保存方式的对比
方式一:保存模型参数(state_dict)
torch.save(model.state_dict(), "xxxxxx.pth")
保存内容:仅保存模型的参数(权重w和偏置b),不包含模型结构。
加载方式:
model = CNN() # 先定义模型结构
model.load_state_dict(torch.load("xxxxxx.pth"))
model.eval()
优点:
-
文件小,只保存参数
-
灵活,可以加载到不同但结构相同的模型
-
是PyTorch推荐的方式
缺点:
-
加载时需要先定义模型结构
方式二:保存完整模型(TorchScript)
script_model = torch.jit.script(model)
torch.jit.save(script_model, "xxxxxxx.pth")
保存内容:模型结构 + 参数 + 计算图,是一个独立可执行的文件。
加载方式:
model = torch.jit.load("xxxxxxx.pth")
model.eval()
优点:
-
无需定义模型结构,直接加载即可用
-
可以跨平台部署(C++、移动端等)
-
适合生产环境
缺点:
-
文件较大
-
某些复杂动态结构可能无法脚本化
6.2 模型文件扩展名
| .pt / .pth | PyTorch通用模型文件 |
| .t7 | Torch7格式(旧版) |
| .onnx | 开放神经网络交换格式 |
6.3 best_acc 的作用
if correct > best_acc:
best_acc = correct
# 保存模型
通过维护一个全局的 best_acc,只有当当前epoch的准确率超过历史最优时才保存。这样可以避免保存效果较差的模型,确保最终保存的是训练过程中表现最好的版本。
七、完整训练流程
loss_fn = nn.CrossEntropyLoss()
optimizer = torch.optim.Adam(model.parameters(), lr=0.001)
epochs = 10
for t in range(epochs):
print(f"Epoch {t+1}\\n———————————–")
train(train_dataloader, model, loss_fn, optimizer)
print("Done!")
test(test_dataloader, model, loss_fn)
注意:test() 只在训练结束后调用一次。如果希望在每个epoch后都评估并保存最优模型,可以在训练循环内调用 test()。
八、数据增强的效果分析
| RandomRotation | 拍摄角度不同 | 提升旋转不变性 |
| RandomFlip | 镜像拍摄 | 提升翻转不变性 |
| ColorJitter | 光照条件不同 | 提升光照鲁棒性 |
| RandomGrayscale | 黑白照片 | 减少对颜色的依赖 |
| Normalize | 数据分布统一 | 加速收敛,提升稳定性 |
实践建议:
-
数据增强不是越多越好,要根据任务特点选择
-
对于人脸识别,垂直翻转通常不合适(人脸有方向性)
-
对于食物分类,旋转、翻转、颜色抖动都很合适
-
验证集必须使用与测试集相同的变换,不能加入随机性
九、总结
本篇博客围绕数据增强和模型保存两大主题,系统讲解了:
| 数据增强 | RandomRotation、RandomFlip、ColorJitter、RandomGrayscale |
| 标准化 | ToTensor + Normalize,使用ImageNet统计参数 |
| 训练/验证变换 | 训练集用增强,验证集只用必要变换 |
| 模型保存方式一 | torch.save(model.state_dict()),保存参数 |
| 模型保存方式二 | torch.jit.script() + torch.jit.save(),保存完整模型 |
| 最优模型保存 | 用 best_acc 追踪,只保存最好的版本 |
核心收获:
数据增强是提升模型泛化能力的关键手段,相当于“免费”扩充数据集。
标准化是深度学习训练的标准步骤,不可省略。
模型保存让训练成果可复用,是工程落地的必要环节。
两种保存方式各有优劣,根据部署需求选择。
网硕互联帮助中心




评论前必须登录!
注册