深度学习入门:数据预处理与自定义数据集
前言:一篇我们学习了卷积神经网络(CNN),并用它实现了 MNIST 手写数字识别。但 MNIST 是 PyTorch 自带的数据集,直接加载就能用。在实际项目中,我们面对的是自己的图片数据,需要自己组织数据、编写数据集类、配合 DataLoader 使用。本篇我们将学习如何用 PyTorch 完成数据预处理,包括数据增强、transforms 操作,以及如何构建自定义数据集,为后续实战项目打下基础。
目录
- 一、PyTorch 数据处理流程
- 二、数据集文件结构
- 三、生成数据索引文件
- 四、自定义 Dataset 类
- 五、DataLoader 批量加载
- 六、搭建 CNN 模型
- 七、总结
一、PyTorch 数据处理流程
1.1 三个核心模块
PyTorch 训练深度学习模型通常包含三个模块:
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.txt 和 test.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 |
| Conv1 | Conv2d(3, 16, 5, 1, 2) + ReLU + MaxPool2d(2) | 16×128×128 |
| Conv2 | Conv2d(16, 32, 5, 1, 2) + ReLU + Conv2d(32, 32, 5, 1, 2) + ReLU + MaxPool2d(2) | 32×64×64 |
| Conv3 | Conv2d(32, 128, 5, 1, 2) + ReLU | 128×64×64 |
| Flatten | view | 128×64×64 = 524288 |
| Linear | 524288 --> 20 | 20 |
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%),主要原因是数据增强幅度较大、训练轮数较少,使用的数据集很小。实际应用中可以通过调整超参数、增加训练轮数、优化数据增强策略来提升准确率。
七、总结
核心知识点速查
| 知识点 | 关键概念 |
|---|---|
| Dataset | PyTorch 数据集的基类,需实现 __len__ 和 __getitem__ |
| DataLoader | 批量加载数据,支持打乱和多进程 |
| transforms | 图像预处理工具,如 Resize、ToTensor |
| 索引文件 | 记录图片路径和标签的 txt 文件 |
__getitem__ | 返回一张图片和标签,DataLoader 自动调用 |
核心 API 一览
| 用途 | 对应模块 / 方法 |
|---|---|
| 数据集基类 | torch.utils.data.Dataset |
| 数据加载器 | torch.utils.data.DataLoader |
| 图像预处理 | torchvision.transforms.Compose |
| 读取图片 | PIL.Image.open() |
| 转为 Tensor | torch.from_numpy() |
注意事项
| 要点 | 说明 |
|---|---|
| 索引文件 | 提前生成,避免每次遍历文件夹 |
| 数据增强 | 训练集和验证集可使用不同的预处理方式 |
| shuffle | 训练集设为 True,测试集设为 False |
| 标签类型 | 需转为 torch.int64 类型才能用于 CrossEntropyLoss |
| 图像格式 | Image.open() 读取的是 RGB 格式,通道维度顺序与 OpenCV 不同 |
系列直达
- 上篇:深度学习入门:卷积神经网络与 MNIST 手写数字识别
- 本篇:深度学习入门:数据预处理与自定义数据集(本文)
- 下篇:敬请期待
转载自 CSDN-专业IT技术社区
原文链接:https://blog.csdn.net/2301_79882046/article/details/165351607



