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

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

在前两篇博客中,我们完成了从自定义数据集到CNN模型训练的完整流程。但如果你仔细回顾,会发现一个潜在的问题:训练数据太单一

模型只见过正着摆放的、亮度固定的、同一角度的物品。一旦测试图片稍有旋转、翻转或颜色变化,模型就可能傻眼------这就是所谓的过拟合

解决这个问题的利器,就是数据增强(Data Augmentation)。它通过对训练图片进行随机变换(旋转、翻转、调色等),人为地制造出更多样化的训练样本,让模型学会忽略这些无关变化,专注于真正的类别特征。

与此同时,训练了若干轮之后,我们得到了一个不错的模型------但如果没有保存,下次就得从头再来。模型保存让训练成果得以持久化,随时可以加载使用。

本篇博客将围绕这两大主题展开,基于完整代码讲解数据增强标准化最优模型保存三大核心知识点。

二、数据增强

2.1 训练集 vs 验证集:两套不同的变换策略

代码中最醒目的设计,是定义了两套变换流程:

python 复制代码
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------随机旋转
python 复制代码
transforms.RandomRotation(45)

-45°到45° 之间随机旋转图片。这模拟了拍摄角度不同的情况,让模型学会识别旋转后的物体。

(2)CenterCrop------中心裁剪
python 复制代码
transforms.CenterCrop(256)

从图像中心裁剪出256×256的区域。配合RandomRotation使用,可以裁掉旋转后产生的黑边,保证输入尺寸一致。

(3)RandomHorizontalFlip / RandomVerticalFlip------随机翻转
python 复制代码
transforms.RandomHorizontalFlip(p=0.5)  # 水平翻转,50%概率
transforms.RandomVerticalFlip(p=0.5)    # 垂直翻转,50%概率

以指定概率对图片进行翻转。水平翻转适合大多数场景(如动物、车辆),垂直翻转则要谨慎使用(对于人脸等有方向性的物体可能不合适)。

(4)ColorJitter------颜色抖动
python 复制代码
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------随机灰度化
python 复制代码
transforms.RandomGrayscale(p=0.1)

以10%的概率将彩色图片转为灰度图(R=G=B)。这强制模型不依赖颜色信息,学习更本质的形状特征。

2.3 ToTensor 与 Normalize------标准化的两步

ToTensor:从PIL到张量
python 复制代码
transforms.ToTensor()

作用:

  • 将PIL图像或NumPy数组转为PyTorch张量

  • 将像素值从 0-255 缩放到 0-1

  • 将通道维度从 HWC 转为 CHW(PyTorch要求)

Normalize:标准化到标准正态分布
python 复制代码
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])

计算方式:

为什么用这组特定的均值和标准差? 它们是 ImageNet数据集上统计出来的RGB三通道均值和标准差。由于大多数预训练模型都是在ImageNet上训练的,使用相同的标准化参数可以保持数据分布一致。

标准化的意义

  • 让数据分布接近标准正态分布,加速梯度下降收敛

  • 消除不同通道之间的量纲差异

  • 是迁移学习中使用预训练模型时的必要步骤

三、自定义数据集回顾

数据集类与上一篇博客一致:

python 复制代码
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

然后分别用训练变换和验证变换创建数据集:

python 复制代码
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个类别:

python 复制代码
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类

五、训练函数

python 复制代码
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

训练过程与之前一致:前向传播→计算损失→梯度清零→反向传播→更新参数。

六、模型保存

这是本篇博客的重点。测试函数中,当模型准确率创新高时,会保存模型:

python 复制代码
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)
python 复制代码
torch.save(model.state_dict(), "xxxxxx.pth")

保存内容 :仅保存模型的参数(权重w和偏置b),不包含模型结构。

加载方式

python 复制代码
model = CNN()  # 先定义模型结构
model.load_state_dict(torch.load("xxxxxx.pth"))
model.eval()

优点

  • 文件小,只保存参数

  • 灵活,可以加载到不同但结构相同的模型

  • 是PyTorch推荐的方式

缺点

  • 加载时需要先定义模型结构
方式二:保存完整模型(TorchScript)
python 复制代码
script_model = torch.jit.script(model)
torch.jit.save(script_model, "xxxxxxx.pth")

保存内容 :模型结构 + 参数 + 计算图,是一个独立可执行的文件。

加载方式

python 复制代码
model = torch.jit.load("xxxxxxx.pth")
model.eval()

优点

  • 无需定义模型结构,直接加载即可用

  • 可以跨平台部署(C++、移动端等)

  • 适合生产环境

缺点

  • 文件较大

  • 某些复杂动态结构可能无法脚本化

6.2 模型文件扩展名

扩展名 说明
.pt / .pth PyTorch通用模型文件
.t7 Torch7格式(旧版)
.onnx 开放神经网络交换格式

6.3 best_acc 的作用

python 复制代码
if correct > best_acc:
    best_acc = correct
    # 保存模型

通过维护一个全局的 best_acc,只有当当前epoch的准确率超过历史最优时才保存。这样可以避免保存效果较差的模型,确保最终保存的是训练过程中表现最好的版本。

七、完整训练流程

python 复制代码
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 追踪,只保存最好的版本

核心收获

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

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

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

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

相关推荐
找方案1 小时前
AI+气象预报:华为盘古大模型如何让天气预报精准到街区
人工智能·算法·机器学习
IT_陈寒1 小时前
Vite的HMR怎么突然罢工了?原来是我漏了这个配置
前端·人工智能·后端
狂师1 小时前
最近火爆出圈的,FDE 到底是个什么岗位?
人工智能·程序员·全栈
不要生病了1 小时前
Through Their Eyes:用简单对齐实现跨被试与跨数据集视觉脑解码
人工智能·深度学习
2601_962304251 小时前
AI照片上色新手好上手:怎么调出自然肤色?
人工智能
临床数据科学和人工智能兴趣组1 小时前
399元现在超值!学R语言,订阅我们专栏就够了,包括了所有的内容,不断更新!
人工智能·数据挖掘·r语言·r语言-4.2.1·临床研究
sjh7524229691 小时前
Deepseek Harness的四种模式
人工智能
一比七品牌咨询1 小时前
科技企业品牌定位:如何把技术优势转化为品牌优势?
大数据·人工智能·品牌策划·品牌全案策划·深圳品牌策划公司·品牌定位
一木 之林1 小时前
深度学习-计算优化与分布式训练-111-128
人工智能·深度学习