【深度学习】卷积神经网络 数据增强、保存最优模型实现,详细解读

文章目录

  • 一、数据增强详解
    • [1. 什么是数据增强](#1. 什么是数据增强)
    • [2. 核心目标](#2. 核心目标)
    • [3. 常用数据增强方法](#3. 常用数据增强方法)
    • [4. 数据预处理与增强的代码实现](#4. 数据预处理与增强的代码实现)
    • [5. 自定义数据集类与增强集成](#5. 自定义数据集类与增强集成)
  • 二、训练过程中保存最优模型
    • [1. 为什么要保存最优模型](#1. 为什么要保存最优模型)
    • [2. 定义 CNN 模型(用于图像分类)](#2. 定义 CNN 模型(用于图像分类))
    • [3. 训练函数](#3. 训练函数)
    • [4. 测试函数与最优模型保存(两种方式)](#4. 测试函数与最优模型保存(两种方式))
    • [5. 训练主循环](#5. 训练主循环)
    • [6. 模型文件说明](#6. 模型文件说明)

一、数据增强详解

1. 什么是数据增强

数据增强(Data Augmentation)是指在不改变原始数据语义的前提下,通过一系列随机变换和组合操作,对已有训练样本进行扩展,生成大量"新"样本的过程。其本质是人为增加训练集的规模和多样性,从而使深度学习模型在面对实际场景中的各种变化(如光照、角度、遮挡)时具备更强的适应能力和稳定性。

2. 核心目标

核心目标是模拟现实世界的复杂多变环境,迫使模型学习到更抽象、更鲁棒的特征表示,而非仅仅记住训练集的具体样本,从而有效降低过拟合风险,提升模型的泛化性能。

3. 常用数据增强方法

方法 描述
随机旋转 将图像绕中心旋转一定角度(如 -45°~45°)
水平/垂直翻转 沿水平或垂直轴镜像翻转图像
随机缩放 按比例放大或缩小图像尺寸
随机平移 沿水平或垂直方向移动若干像素
随机裁剪 从原图中截取部分区域
亮度/对比度/饱和度调整 改变颜色空间的数值
添加噪声 叠加高斯、椒盐等噪声
几何扭曲 仿射变换、弹性变形等

4. 数据预处理与增强的代码实现

在 PyTorch 中,通常使用 torchvision.transforms 组合多种操作,并分别对训练集和验证集设置不同的处理流水线。

c 复制代码
import torch
from torch.utils.data import DataLoader, Dataset
import numpy as np
from PIL import Image
from torchvision import transforms

# 定义训练和验证阶段的不同预处理
data_transforms = {
    'train': transforms.Compose([
        transforms.Resize([300, 300]),                # 统一尺寸
        transforms.RandomRotation(degrees=45),        # 随机旋转 [-45°, 45°]
        transforms.CenterCrop(256),                   # 中心裁剪为 256x256
        transforms.RandomHorizontalFlip(p=0.5),       # 水平翻转概率 50%
        transforms.RandomVerticalFlip(p=0.5),         # 垂直翻转概率 50%
        transforms.ColorJitter(brightness=0.2, contrast=0.1, saturation=0.1, hue=0.1),  # 颜色抖动
        transforms.RandomGrayscale(p=0.1),            # 10% 概率转为灰度图
        transforms.ToTensor(),                        # 转为 Tensor 并归一化到 [0,1]
        transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])  # ImageNet 标准化
    ]),
    'valid': transforms.Compose([
        transforms.Resize([256, 256]),
        transforms.ToTensor(),
        transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
    ])
}

提醒:数据增强并非总能提升效果,需根据具体任务和数据集进行调优,但通常情况下会带来正向收益。

5. 自定义数据集类与增强集成

我们通过继承 Dataset 类,在 getitem 中应用上述变换,从而在每次读取样本时动态生成增强后的图像。

c 复制代码
class FoodDataset(Dataset):
    """自定义食物图片数据集"""
    def __init__(self, file_path, transform=None):
        self.file_path = file_path
        self.transform = transform
        self.image_paths = []
        self.labels = []
        # 解析文件,每行格式:图片路径 类别标签
        with open(file_path) as f:
            lines = [line.strip().split() for line in f.readlines()]
            for img_path, label in lines:
                self.image_paths.append(img_path)
                self.labels.append(label)

    def __len__(self):
        return len(self.image_paths)

    def __getitem__(self, idx):
        img = Image.open(self.image_paths[idx])
        if self.transform:
            img = self.transform(img)
        # 标签转换为 Tensor
        label = torch.from_numpy(np.array(int(self.labels[idx]), dtype=np.int64))
        return img, label

# 实例化训练集和验证集
train_dataset = FoodDataset(file_path='./trainda.txt', transform=data_transforms['train'])
valid_dataset = FoodDataset(file_path='./testda.txt', transform=data_transforms['valid'])

# 设备选择
device = "cuda" if torch.cuda.is_available() else "mps" if torch.backends.mps.is_available() else "cpu"
print(f"当前使用的设备: {device}")

其中 trainda.txt 和 testda.txt 的内容格式为:

...

(每行一个样本,路径与标签用空格分隔)

二、训练过程中保存最优模型

1. 为什么要保存最优模型

在深度学习的迭代训练中,模型参数会随着优化步骤不断更新。通常,我们会在每个 Epoch 结束时在验证集上评估性能,并将当前验证集上表现最好的模型参数持久化到磁盘(常见扩展名为 .pt、.pth 或 .t7)。这样做可以避免因过拟合或训练后期震荡而丢失最佳状态,也便于后续部署或继续微调。

2. 定义 CNN 模型(用于图像分类)

这里构建一个三层卷积 + 全连接的简单 CNN,输入为 3×256×256 的 RGB 图像。

c 复制代码
from torch import nn

class CNN(nn.Module):
    def __init__(self):
        super(CNN, self).__init__()
        self.conv1 = nn.Sequential(
            nn.Conv2d(in_channels=3, out_channels=16, kernel_size=5, stride=1, padding=2),  # -> (16, 256, 256)
            nn.ReLU(),
            nn.MaxPool2d(kernel_size=2)  # -> (16, 128, 128)
        )
        self.conv2 = nn.Sequential(
            nn.Conv2d(16, 32, kernel_size=5, stride=1, padding=2),  # -> (32, 128, 128)
            nn.ReLU(),
            nn.MaxPool2d(2)  # -> (32, 64, 64)
        )
        self.conv3 = nn.Sequential(
            nn.Conv2d(32, 128, kernel_size=5, stride=1, padding=2),  # -> (128, 64, 64)
            nn.ReLU()
        )
        self.fc = nn.Linear(128 * 64 * 64, 20)  # 假设类别数为 20

    def forward(self, x):
        x = self.conv1(x)
        x = self.conv2(x)
        x = self.conv3(x)
        x = x.view(x.size(0), -1)
        out = self.fc(x)
        return out

model = CNN().to(device)
print(model)

3. 训练函数

c 复制代码
def train_one_epoch(dataloader, model, loss_fn, optimizer):
    model.train()
    batch_idx = 1
    for X, y in dataloader:
        X, y = X.to(device), y.to(device)
        pred = model(X)                 # 前向传播
        loss = loss_fn(pred, y)         # 计算损失

        optimizer.zero_grad()           # 清零梯度
        loss.backward()                 # 反向传播
        optimizer.step()                # 更新参数

        if batch_idx % 100 == 0:
            print(f"  批次 {batch_idx} 损失: {loss.item():.4f}")
        batch_idx += 1

4. 测试函数与最优模型保存(两种方式)

定义全局变量 best_acc 跟踪最高准确率,若当前验证准确率更高则保存模型。

c 复制代码
best_acc = 0.0

def evaluate_and_save(dataloader, model, loss_fn):
    global best_acc
    size = len(dataloader.dataset)
    num_batches = len(dataloader)
    model.eval()
    test_loss, correct = 0.0, 0

    with torch.no_grad():
        for X, y in dataloader:
            X, y = X.to(device), y.to(device)
            pred = model(X)
            test_loss += loss_fn(pred, y).item()
            correct += (pred.argmax(1) == y).type(torch.float).sum().item()

    test_loss /= num_batches
    accuracy = correct / size
    print(f"验证结果: 准确率 = {accuracy:.2%}, 平均损失 = {test_loss:.4f}")

    # 保存最优模型
    if accuracy > best_acc:
        best_acc = accuracy
        # 方式一:仅保存模型参数(推荐,占用空间小)
        # torch.save(model.state_dict(), 'best_params.pth')
        # 方式二:保存完整模型(包含架构和参数)
        torch.save(model, 'best_model.pt')
        print(f"模型已更新保存,当前最佳准确率: {best_acc:.2%}")

5. 训练主循环

c 复制代码
loss_fn = nn.CrossEntropyLoss()
optimizer = torch.optim.Adam(model.parameters(), lr=0.001)

train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True)
valid_loader = DataLoader(valid_dataset, batch_size=64, shuffle=False)

epochs = 150
for epoch in range(epochs):
    print(f"\nEpoch {epoch+1}/{epochs}")
    train_one_epoch(train_loader, model, loss_fn, optimizer)
    evaluate_and_save(valid_loader, model, loss_fn)

运行过程会逐轮输出训练损失和验证准确率,当验证准确率超过历史最佳时,自动保存新模型。

6. 模型文件说明

训练结束后,best_model.pt(或 best_params.pth)即为最优模型文件。加载方法:

若保存的是完整模型:

c 复制代码
model = torch.load('best_model.pt')

若保存的只是状态字典:先实例化模型结构,再

c 复制代码
model.load_state_dict(torch.load('best_params.pth'))
相关推荐
rain_sxr1 小时前
把多步串起来:Agent 前端编排的状态机与进度可视化
人工智能
天青色等烟雨..1 小时前
全流程ArcGISPro空间分析、三维建模、可视化及Python融合应用技术
开发语言·python
极光代码工作室1 小时前
基于Spark的日志监控与分析平台
大数据·hadoop·python·spark·数据可视化
硅谷秋水1 小时前
PhyGround:生成式世界模型中的物理推理基准测试
人工智能·深度学习·机器学习·计算机视觉·语言模型
AOwhisky1 小时前
Python 学习笔记(第十三期)——运维自动化(下·前篇):远程命令执行——paramiko基础篇
运维·python·学习·云原生·自动化·运维开发·paramiko
深海鱼肝油ya1 小时前
基于FastAPI的AI智能体Web系统构建(二)
人工智能·fastapi·python开发·异步框架·agent开发
浊酒南街1 小时前
subprocess.check_output函数介绍
python
TheBestRucy1 小时前
RAG知识库问答系统落地:从向量检索到上下文增强的全链路实践
人工智能·python·langchain·aigc·交互