云计算百科
云计算领域专业知识百科平台

初识深度学习——数据增强与模型保存

一、引言:让模型长见识,让成果留下来

在前两篇博客中,我们完成了从自定义数据集到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])

计算方式:

x_{normalize} = \\frac{x - mean}{std}

为什么用这组特定的均值和标准差? 它们是 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 追踪,只保存最好的版本

核心收获:

  • 数据增强是提升模型泛化能力的关键手段,相当于“免费”扩充数据集。

  • 标准化是深度学习训练的标准步骤,不可省略。

  • 模型保存让训练成果可复用,是工程落地的必要环节。

  • 两种保存方式各有优劣,根据部署需求选择。

  • 赞(0)
    未经允许不得转载:网硕互联帮助中心 » 初识深度学习——数据增强与模型保存
    分享到: 更多 (0)

    评论 抢沙发

    评论前必须登录!