词却头像
关注

深度学习入门:数据预处理与自定义数据集

深度学习入门:数据预处理与自定义数据集

前言:一篇我们学习了卷积神经网络(CNN),并用它实现了 MNIST 手写数字识别。但 MNIST 是 PyTorch 自带的数据集,直接加载就能用。在实际项目中,我们面对的是自己的图片数据,需要自己组织数据、编写数据集类、配合 DataLoader 使用。本篇我们将学习如何用 PyTorch 完成数据预处理,包括数据增强、transforms 操作,以及如何构建自定义数据集,为后续实战项目打下基础。

目录

  • 一、PyTorch 数据处理流程
  • 二、数据集文件结构
  • 三、生成数据索引文件
  • 四、自定义 Dataset 类
  • 五、DataLoader 批量加载
  • 六、搭建 CNN 模型
  • 七、总结

一、PyTorch 数据处理流程

1.1 三个核心模块

PyTorch 训练深度学习模型通常包含三个模块:

数据处理模块
(Dataset,DataLoader)

模型构建模块
(nn.Module)

训练控制模块
(train / test)

1.2 各模块作用

模块作用
数据处理模块负责读取图片、预处理、分批加载
模型构建模块定义神经网络结构
训练控制模块定义训练循环、损失函数、优化器

二、数据集文件结构

2.1 文件目录结构

假设我们有一个食物分类数据集,结构如下:

food_dataset/
    ├── train/
    │   ├── 0/          # 类别0(如:八宝粥)
    │   │   ├── img1.jpg
    │   │   └── img2.jpg
    │   ├── 1/          # 类别1(如:哈密瓜)
    │   ├── ...
    │   └── 19/         # 类别19
    └── test/
        ├── 0/
        ├── 1/
        ├── ...
        └── 19/
说明内容
训练集train 文件夹
测试集test 文件夹
类别数量20 类
每类数据存放在对应编号的子文件夹中

2.2 为什么用这种结构?

PyTorch 的 Dataset 需要我们自己提供图片路径和对应标签。用子文件夹区分类别是最直观的方式,读取时目录名即为标签。

三、生成数据索引文件

3.1 为什么需要索引文件?

每次训练都遍历文件夹读取图片路径,效率较低。更好的做法是提前生成一份索引文件,记录所有图片路径和标签,训练时直接读取。

3.2 生成索引文件的函数

import os


def train_test_file(root, dir):
    """生成图片路径和标签的索引文件"""
    file_txt = open(dir + ".txt", "w")  # 创建 train.txt / test.txt 用于写入索引
    path = os.path.join(root, dir)  # 拼接数据集根目录与子目录(train / test)

    for roots, directories, files in os.walk(path):  # 遍历目录树,roots为当前路径,directories为子文件夹,files为文件
        if len(directories) != 0:  # 如果当前层还有子文件夹,说明是类别目录,跳过
            dirs = directories  # 记录类别列表(如 [0,1,2,...,19])
        else:  # 否则说明已到最底层(图片所在目录)
            now_dir = roots.split("\\")  # 按反斜杠拆分路径,取出当前类别文件夹名
            for file in files:  # 遍历该类别下的所有图片
                path_1 = os.path.join(roots, file)  # 拼接图片完整路径
                print(path_1)  # 打印路径,方便查看进度
                file_txt.write(path_1 + ' ' + str(dirs.index(now_dir[-1])) + '\n')  # 写入 "图片路径 标签"

    file_txt.close()  # 关闭文件


root = r'..\data\food_classification\food_dataset'  # 数据集根目录
train_dir = 'train'  # 训练集子目录名
test_dir = 'test'  # 测试集子目录名

train_test_file(root, train_dir)  # 生成 train.txt
train_test_file(root, test_dir)  # 生成 test.txt

3.3 生成结果

运行后在当前目录生成 train.txttest.txt,格式如下:

..\data\food_classification\food_dataset\train\八宝粥\img_八宝粥罐_22.jpeg
..\data\food_classification\food_dataset\train\八宝粥\img_八宝粥罐_29.jpeg
..\data\food_classification\food_dataset\train\八宝粥\img_八宝粥罐_65.jpeg
..\data\food_classification\food_dataset\train\八宝粥\img_八宝粥罐_68.jpeg
......
..\data\food_classification\food_dataset\test\骨肉相连\img_骨肉相连_331.jpeg
..\data\food_classification\food_dataset\test\鸡翅\img_鸡翅_311.jpeg
..\data\food_classification\food_dataset\test\鸡翅\img_鸡翅_335.jpeg

四、自定义 Dataset 类

4.1 Dataset 的作用

PyTorch 的 Dataset 是一个抽象类,我们需要继承它并实现两个方法:

方法作用
__len__返回数据集的总样本数
__getitem__根据索引返回一张图片和对应标签

4.2 数据增强(Data Augmentation)

数据增强是缓解深度学习中数据不足的重要手段,在图像领域应用广泛。它的核心思想是通过变换增加训练数据的多样性,从而提高模型的泛化能力

常见的数据增强方式:

增强方式说明
垂直翻转上下翻转图片
随机旋转在指定角度范围内随机旋转
随机裁剪从图片中随机裁剪出部分区域
颜色变换调整亮度、对比度、饱和度、色调

4.3 训练集与测试集的预处理

数据增强只在训练集上使用,测试集只需要做基础的尺寸调整和标准化:

from torchvision import transforms

data_transforms = {
    'trainda':
        transforms.Compose([  # 对图片做预处理的组合
            transforms.Resize((256, 256)),  # 调整大小为 [256,256]
            transforms.RandomRotation(45),  # 随机旋转,-45 到 45 度之间随机选
            transforms.CenterCrop(256),  # 从中心开始裁剪 [256,256]
            transforms.RandomHorizontalFlip(p=0.5),  # 随机水平翻转,概率 0.5
            transforms.RandomVerticalFlip(p=0.5),  # 随机垂直翻转,概率 0.5
            transforms.ColorJitter(0.2, 0.1, 0.1, 0.1),  # 颜色变换:亮度、对比度、饱和度、色调
            transforms.RandomGrayscale(p=0.1),  # 概率转换为灰度图(3 通道 R=G=B)
            transforms.ToTensor(),  # 转换为 Tensor,通道维度放在前面
            transforms.Normalize([0.485, 0.456, 0.406],  # 标准化:ImageNet 均值
                                 [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])
        ]),
}

4.4 各操作参数说明

操作参数说明
RandomRotation(45)角度范围在 -45° ~ 45° 之间随机旋转
RandomHorizontalFlip(p=0.5)概率有 50% 的概率水平翻转
RandomVerticalFlip(p=0.5)概率有 50% 的概率垂直翻转
ColorJitter(0.2, 0.1, 0.1, 0.1)亮度、对比度、饱和度、色调各自的变化幅度
RandomGrayscale(p=0.1)概率有 10% 的概率转为灰度图
Normalize(mean, std)均值、标准差通常使用 ImageNet 的统计值

4.5 为什么训练集和测试集不同?

数据集处理方式原因
训练集数据增强 + 标准化增加多样性,提升泛化能力
测试集仅标准化保持数据真实分布,评估才准确

注意:数据增强只在训练阶段使用,验证和测试阶段必须使用与真实应用一致的预处理方式。

4.6 Dataset 类

from torch.utils.data import Dataset
from PIL import Image
import torch
import numpy as np


class FoodDataset(Dataset):  # 自定义数据集类,必须继承 Dataset
    def __init__(self, file_path, transform=None):
        """类的初始化,解析数据文件txt"""
        self.file_path = file_path  # 索引文件路径(如 train.txt)
        self.imgs = []  # 存放所有图片的路径
        self.labels = []  # 存放所有图片对应的标签
        self.transform = transform  # 图像预处理操作(缩放、翻转等)

        with open(self.file_path) as f:  # 打开索引文件
            samples = [x.strip().split(' ') for x in f.readlines() if x.strip()]  # 逐行读取,按空格拆分为 [路径, 标签]
            for img_path, label in samples:  # 遍历每条样本
                self.imgs.append(img_path)  # 保存图像路径
                self.labels.append(label)  # 保存标签

    def __len__(self):
        """返回数据集样本总数"""
        return len(self.imgs)  # DataLoader 会据此决定遍历范围

    def __getitem__(self, idx):
        """根据索引获取一张图片和对应标签"""
        image = Image.open(self.imgs[idx]).convert('RGB')  # 根据索引读取图片(PIL格式,RGB通道)
        if self.transform:  # 如果有预处理操作
            image = self.transform(image)  # 执行预处理(缩放、增强、转Tensor等)
        label = self.labels[idx]  # 取出对应标签
        label = torch.from_numpy(np.array(label, dtype=np.int64))  # 转为 int64 类型的 Tensor
        return image, label  # 返回图片和标签,供 DataLoader 打包成批次

4.7 关键方法说明

方法说明
__init__初始化时读取索引文件,保存所有图片路径和标签
__len__返回样本总数,供 DataLoader 使用
__getitem__根据索引返回一张图片和标签,DataLoader 会自动调用

五、DataLoader 批量加载

5.1 创建 Dataset 和 DataLoader

from torch.utils.data import DataLoader

# 创建数据集
training_data = FoodDataset(file_path='./train.txt', transform=data_transforms['trainda'])
test_data = FoodDataset(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=False)

5.2 DataLoader 的作用

作用说明
批量加载每次取 batch_size 张图片
打乱顺序shuffle=True 训练时打乱,测试时保持顺序
多进程加载自动使用多进程加速数据读取

5.3 DataLoader 参数说明

参数说明
batch_size每个批次包含多少张图片
shuffle是否打乱顺序,训练集设为 True,测试集设为 False

六、搭建 CNN 模型

数据处理完成后,模型构建和训练部分我们之前搭建过,可以直接沿用思路。

6.1 CNN 模型定义

from torch import nn

class CNN(nn.Module):  # 定义卷积神经网络,继承 nn.Module
    def __init__(self):
        super(CNN, self).__init__()  # 初始化父类
        # 第1个卷积块:卷积 -> 激活 -> 池化
        self.conv1 = nn.Sequential(
            nn.Conv2d(in_channels=3, out_channels=16, kernel_size=5, stride=1, padding=2),  # 3×256×256 -> 16×256×256
            nn.ReLU(),  # 激活函数,不改变尺寸
            nn.MaxPool2d(kernel_size=2),  # 池化,尺寸减半 -> 16×128×128
        )
        # 第2个卷积块:两层卷积 -> 激活 -> 池化
        self.conv2 = nn.Sequential(
            nn.Conv2d(16, 32, 5, 1, 2),  # 16×128×128 -> 32×128×128
            nn.ReLU(),
            nn.Conv2d(32, 32, 5, 1, 2),  # 32×128×128 -> 32×128×128
            nn.ReLU(),
            nn.MaxPool2d(2),  # 尺寸减半 -> 32×64×64
        )
        # 第3个卷积块:卷积 -> 激活
        self.conv3 = nn.Sequential(
            nn.Conv2d(32, 128, 5, 1, 2),  # 32×64×64 -> 128×64×64
            nn.ReLU(),
        )
        # 全连接层:将 128×64×64 展平后映射到 20 个类别
        self.out = nn.Linear(128 * 64 * 64, 20)

    def forward(self, x):  # 前向传播,定义数据流向
        x = self.conv1(x)  # 经过第1个卷积块
        x = self.conv2(x)  # 经过第2个卷积块
        x = self.conv3(x)  # 经过第3个卷积块
        x = x.view(x.size(0), -1)  # 展平:(batch_size, 128*64*64)
        output = self.out(x)  # 全连接层输出 (batch_size, 20)
        return output

6.2 模型结构说明

层级操作输出尺寸
输入-3×256×256
Conv1Conv2d(3, 16, 5, 1, 2) + ReLU + MaxPool2d(2)16×128×128
Conv2Conv2d(16, 32, 5, 1, 2) + ReLU + Conv2d(32, 32, 5, 1, 2) + ReLU + MaxPool2d(2)32×64×64
Conv3Conv2d(32, 128, 5, 1, 2) + ReLU128×64×64
Flattenview128×64×64 = 524288
Linear524288 --> 2020

6.3 训练与测试函数

import torch

# 自动选择设备:优先CUDA,其次MPS,最后CPU
device = "cuda" if torch.cuda.is_available() else "mps" if torch.backends.mps.is_available() else "cpu"
print(f"Using {device} device")
model = CNN().to(device)  # 模型移到设备上


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()  # 取损失值
        print(f"loss: {loss_value:>7f}  [number:{batch_size_num}]")
        batch_size_num += 1


def test(dataloader, model, loss_fn):
    """测试集评估"""
    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}")

6.4 开始训练

loss_fn = nn.CrossEntropyLoss()  # 交叉熵损失

optimizer = torch.optim.Adam(model.parameters(), lr=0.001)  # Adam优化器

epochs = 10  # 训练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)  # 测试集评估
训练结果:
Epoch 1
====================================
loss: 3.001019  [number:1]
loss: 19.668638  [number:2]
loss: 6.641384  [number:3]
loss: 3.080421  [number:4]
loss: 2.910440  [number:5]
Epoch 2
......
Epoch 10
====================================
loss: 1.947709  [number:1]
loss: 2.374726  [number:2]
loss: 2.388289  [number:3]
loss: 2.342167  [number:4]
loss: 2.616031  [number:5]
Done!
Test result: 
 Accuracy: 15.384615384615385%, Avg loss: 2.68658185005188

结果说明:本案例中准确率较低(15.38%),主要原因是数据增强幅度较大、训练轮数较少,使用的数据集很小。实际应用中可以通过调整超参数、增加训练轮数、优化数据增强策略来提升准确率。

七、总结

核心知识点速查

知识点关键概念
DatasetPyTorch 数据集的基类,需实现 __len____getitem__
DataLoader批量加载数据,支持打乱和多进程
transforms图像预处理工具,如 ResizeToTensor
索引文件记录图片路径和标签的 txt 文件
__getitem__返回一张图片和标签,DataLoader 自动调用

核心 API 一览

用途对应模块 / 方法
数据集基类torch.utils.data.Dataset
数据加载器torch.utils.data.DataLoader
图像预处理torchvision.transforms.Compose
读取图片PIL.Image.open()
转为 Tensortorch.from_numpy()

注意事项

要点说明
索引文件提前生成,避免每次遍历文件夹
数据增强训练集和验证集可使用不同的预处理方式
shuffle训练集设为 True,测试集设为 False
标签类型需转为 torch.int64 类型才能用于 CrossEntropyLoss
图像格式Image.open() 读取的是 RGB 格式,通道维度顺序与 OpenCV 不同

系列直达

转载自 CSDN-专业IT技术社区

原文链接:https://blog.csdn.net/2301_79882046/article/details/165351607

文章来源转载

评论

赞0

评论列表

微信小程序
QQ小程序

关于作者

点赞数:0
关注数:0
粉丝:0
文章:0
关注标签:0
加入于:--