动手学深度学习 04 | Dataset 和 DataLoader、数据增强

目录

摘要

一、需求

[二、什么是 Dataset](#二、什么是 Dataset)

三、数据增强

四、DataLoader

五、完整案例演示


摘要

前面几节我们使用了 TorchVision 内置好的 MNIST 数据集,这类数据集已经完成文件读取、标签封装等预处理。但在真实项目场景中,我们经常会使用自己收集整理的图片数据集。本节课我们就来解决一个工程问题:如何读取本地自定义图片,搭建属于自己的 Dataset,搭配 DataLoader 批量加载数据,并把图像送入 GPU 参与训练,同时学习数据增强技术,提升模型泛化能力。

一、需求

我们准备了一个食物分类数据集,目录结构如下:

|--------------------------------------------------------------------------------------------------------------------|
| Plain Text food_dataset ├─train │ ├─八宝粥 │ ├─哈密瓜 │ ├─圣女果 │ └─......各类食物文件夹 └─test ├─八宝粥 ├─哈密瓜 ├─圣女果 └─......各类食物文件夹 |

每个类别单独放在一个文件夹内,文件夹名称就是类别名称,文件夹内存放该类别的图片。

我们需要实现:

  1. 遍历文件夹,自动读取图片路径并生成对应标签;
  2. 自定义 Dataset 类,加载图片与标签;
  3. 使用数据增强扩充训练样本;
  4. 通过 DataLoader 实现批量加载、送入 GPU 训练;
  5. 训练 CNN 模型,保存最优模型,最后单独拿一张图片做推理预测。

二、什么是 Dataset

torch.utils.data.Dataset是 PyTorch 提供的数据集抽象基类 ,用来定义数据集读取逻辑。

想要自定义数据集,必须重写 3 个核心方法:

  1. init:初始化,读取图片路径、标签、预处理变换;
  2. len:返回数据集总样本数量;
  3. getitem:按下标索引,返回单条样本(图像 + 标签)。

Dataset 的作用:解耦数据读取和模型训练,只负责定义如何获取单条样本,不负责分批、打乱。

三、数据增强

深度学习模型的训练需要充足的样本支撑,若数据集样本数量有限,模型极易死记硬背训练图像的像素细节,无法学习到物体的通用特征,最终导致在测试集上表现不佳,出现典型的过拟合问题。

数据增强是解决小样本训练、抑制过拟合的核心手段。在模型训练过程中,通过对原始图像进行随机几何变换、像素调整等操作,生成多样化的虚拟样本。该方式无需额外采集真实数据,即可极大丰富数据集的样本多样性,引导模型聚焦学习物体的核心轮廓、纹理等固有特征,而非图像的位置、角度、光影等无关干扰信息,有效提升模型的泛化能力与鲁棒性。

|---------------------------------------------------------------------|
| 注意:训练集使用随机增强,验证 / 测试集只用固定缩放,不能加随机操作 。测试阶段我们需要真实评估模型性能,不能随机修改图片。 |

常见数据增强种类

  • Resize:统一缩放图片到固定尺寸,保证输入网络的图片大小一致;
  • RandomRotation:随机旋转图片;
  • RandomHorizontalFlip:随机水平翻转;
  • RandomCrop / CenterCrop:随机裁剪、中心裁剪;
  • ColorJitter:随机调整亮度、对比度、饱和度、色相;
  • ToTensor:将 PIL 图片像素值从0,255转为0,1张量,同时调换通道顺序。

四、DataLoader

Dataset 只定义了单个样本怎么读取,而DataLoader 是加载器,在 Dataset 基础上实现批量打包。 DataLoader 的核心好处:

  1. batch 批量打包:一次取出 batch_size 个样本,组成一个批次送入网络,充分利用 GPU 并行计算;

  2. shuffle 打乱:训练集开启打乱,防止模型记住样本顺序,提升收敛效果;

  3. 多线程读取:可通过 num_workers 开启多进程,磁盘读取与 GPU 计算并行,加快数据加载速度;

  4. 自动堆叠张量,直接输出可以送入模型的批量图像张量和标签张量。

拓展:模型训练时,CPU负责读取磁盘中的图像数据并完成预处理,再经由内存将数据传输 至GPU显存。每次会向GPU送入一个batch的样本,利用GPU的并行计算能力,同步完成批量样本的推理运算、损失计算、梯度求解与模型参数更新。当单批次训练迭代完成后,系统自动释放显存资源,再载入下一批样本继续训练,该高效迭代方式便是批量梯度下降算法,也是深度学习模型主流的训练方式

五、完整案例演示

5.1 第一步:遍历目录,生成图片路径 + 标签 txt 文件

先写脚本递归遍历数据集文件夹,把图片路径 类别编号写入 txt 文本,后续 Dataset 直接读取 txt 加载数据。

python 复制代码
import os

def train_test_file(root,dir):
    file_txt = open(dir+'.txt','w',encoding='utf-8')
    path = os.path.join(root,dir)
    # os.walk递归遍历目录:先遍历主文件夹,再逐层遍历子文件夹,深度优先遍历
    # roots:当前文件夹完整路径;directories:当前目录下子文件夹名列表;files:当前目录下所有文件名列表
    for roots,directories,files in os.walk(path):
        if len(directories) !=0:
            dirs = directories
        else:
            now_dir = roots.split('\\')
            for file in files:
                path_1 = os.path.join(roots,file)
                file_txt.write(path_1+' '+str(dirs.index(now_dir[-1]))+'\n')
    file_txt.close()

# 修改为你的数据集根目录
root = r'D:\Mystudy\data\food_dataset'
train_dir = 'train'
test_dir = 'test'

train_test_file(root,train_dir)
train_test_file(root,test_dir)

运行后会在当前目录生成train.txt和test.txt,每一行格式:图片完整路径 类别数字标签。

5.2 第二步:定义数据增强策略 & 自定义 Dataset

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

# 定义数据增强策略:训练集随机增强,验证集仅做固定缩放
data_transform ={
    'train':transforms.Compose([
        transforms.Resize([256,256]),
        transforms.RandomRotation(45), #随机旋转±45度
        transforms.ToTensor()
    ]),
    'val':transforms.Compose([
        transforms.Resize([256,256]),
        transforms.ToTensor()
    ])
}

# 自定义Dataset类,继承torch.utils.data.Dataset
class food_dataset(Dataset):
    def  __init__(self,file_path,transform = None):
        self.file_path = file_path
        self.img_paths = []
        self.label = []
        self.transform = transform
        #读取txt文件,解析图片路径与标签
        with open(file_path,'r',encoding='utf-8') as f :
            samples = [x.strip().split(' ') for x in f.readlines()]
            for img_path,label in samples:
                self.img_paths.append(img_path)
                self.label.append(int(label))

    # 返回数据集样本总数
    def __len__(self):
        return len(self.img_paths)

    # 根据索引读取单张图片和标签
    def __getitem__(self,idx):
        image = Image.open(self.img_paths[idx])
        # 如果定义了transform,则执行图像预处理/增强
        if self.transform:
            image = self.transform(image)
        label = torch.tensor(self.label[idx],dtype = torch.long)
        return image,label

# 实例化训练集、测试集
training_data = food_dataset(r"./train.txt",data_transform['train'])
test_data = food_dataset(r"./test.txt",data_transform['val'])

# DataLoader批量加载
train_loader = DataLoader(training_data,batch_size = 32,shuffle = True)
test_loader = DataLoader(test_data,batch_size = 32,shuffle = False)

|----------------------------------------------------------|
| 说明:训练集 shuffle=True 打乱样本;测试集一般建议 shuffle=False,方便观测预测结果。 |

5.3 第三步:搭建 CNN 分类网络

python 复制代码
class Cnn(nn.Module):
    def __init__(self):
        super().__init__()
        # 特征提取模块:卷积+激活+池化
        self.features = nn.Sequential(
            nn.Conv2d(in_channels=3, out_channels=32, kernel_size=3,stride=1, padding=1),
            nn.ReLU(),
            nn.MaxPool2d(kernel_size=2),

            nn.Conv2d(in_channels=32, out_channels=64, kernel_size=3,stride=1, padding=1),
            nn.ReLU(),
            nn.MaxPool2d(kernel_size=2),

            nn.Conv2d(in_channels=64, out_channels=32, kernel_size=3,stride=1, padding=1),
            nn.ReLU()
        )
        # 分类头,输出20个类别,对应20种食物
        self.classifier = nn.Sequential(
            nn.Flatten(),
            nn.Linear(32*64*64, 20)
        )

    def forward(self, x):
         x = self.features(x)
         x = self.classifier(x)
         return x

5.4 第四步:设备、损失函数、优化器配置

python 复制代码
class Cnn(nn.Module):
    def __init__(self):
        super().__init__()
        # 特征提取模块:卷积+激活+池化
        self.features = nn.Sequential(
            nn.Conv2d(in_channels=3, out_channels=32, kernel_size=3,stride=1, padding=1),
            nn.ReLU(),
            nn.MaxPool2d(kernel_size=2),

            nn.Conv2d(in_channels=32, out_channels=64, kernel_size=3,stride=1, padding=1),
            nn.ReLU(),
            nn.MaxPool2d(kernel_size=2),

            nn.Conv2d(in_channels=64, out_channels=32, kernel_size=3,stride=1, padding=1),
            nn.ReLU()
        )
        # 分类头,输出20个类别,对应20种食物
        self.classifier = nn.Sequential(
            nn.Flatten(),
            nn.Linear(32*64*64, 20)
        )

    def forward(self, x):
         x = self.features(x)
         x = self.classifier(x)
         return x

5.5 第五步:训练函数、测试函数 + 模型保存

python 复制代码
def train(model,device,train_dataloader,loss_fn,optimizer):
    length = len(train_dataloader.dataset)
    num_batches = len(train_dataloader)
    model.train() #开启训练模式
    sum_loss = 0.0
    correct = 0
    for X,y in train_dataloader:
        X=X.to(device)
        y=y.to(device)
        y_pred = model(X)
        loss = loss_fn(y_pred,y)
        optimizer.zero_grad() #清空梯度
        loss.backward() #反向传播求梯度
        optimizer.step() #更新参数
        sum_loss += loss.item()
        correct += (y_pred.argmax(1) == y).sum().item()
    print(f"平均准确率{correct / length:.4f},平均损失{sum_loss / num_batches:.4f} ")


def test(model, device, test_dataloader, loss_fn):
    best_c = 0.0
    length = len(test_dataloader.dataset)
    num_batches = len(test_dataloader)
    model.eval() #评估模式
    sum_loss =0.0
    correct = 0
    with torch.no_grad(): #关闭梯度计算,节省显存
        for X,y in test_dataloader:
            X=X.to(device)
            y=y.to(device)
            y_pred = model(X)
            loss = loss_fn(y_pred,y)
            sum_loss += loss.item()
            correct += (y_pred.argmax(1) == y).sum().item()
    acc = correct / length
    print(f"测试集平均准确率{acc:.4f},测试平均损失{sum_loss/num_batches:.4f}")

5.6 主程序入口:启动训练 + 单图推理测试

实验小结

本节课我们完成自定义图片数据集的完整工程流程:

  1. 通过 os.walk 遍历本地图片,生成图片路径与标签的索引文件;
  2. 继承 Dataset 自定义数据集类,实现图片读取;
  3. 使用 transforms 做数据增强,缓解小数据集过拟合;
  4. DataLoader 实现批量加载数据,配合 GPU 加速训练;
  5. 训练 CNN 并保存模型,最后实现单张图片推理。

这套流程是深度学习项目最常用的本地数据集开发范式,后续不管是图像分类、医学影像等任务,都可以复用这套 Dataset+DataLoader 的框架。

相关推荐
东方佑2 小时前
可微概率后缀超图:检索硬、聚合软 —— 与 ROSA 的对比及真实链路验证
人工智能
`流年づ2 小时前
人工智能学习笔记 - 自动微分
人工智能·笔记·学习
AIGC大时代2 小时前
防 AI bot 审稿:CARMA 闸门、失败含义与当天最小实验
人工智能·审稿·carma·人工闸门
Quor2 小时前
Zorv AI GenUI 技术架构深度解析:从双面设计到安全边界
人工智能·ui·架构
elseif1232 小时前
【双指针/二分】P1102 A-B 数对
c++·算法·二分·双指针
电梯界知识分子2 小时前
江西抚州临川九尊府五层别墅:受限楼梯间里,全黑铝合金观光井道配曳引龙门架的落地记录
大数据·前端·网络·算法·家用电梯
东离与糖宝3 小时前
SSE流式输出详解:大模型打字机效果底层原理
人工智能
李兆龙的博客3 小时前
从一到无穷大 #91:从 Habitat 看存储平台的整合与分工
数据库·人工智能·架构
阿文和她的Key3 小时前
OpenAI 关 Pro 入口事件复盘:企业 AI 架构的稳定性问题,不只是故障应急
人工智能·架构