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

|--------------------------------------------------------------------------------------------------------------------|
| Plain Text food_dataset ├─train │ ├─八宝粥 │ ├─哈密瓜 │ ├─圣女果 │ └─......各类食物文件夹 └─test ├─八宝粥 ├─哈密瓜 ├─圣女果 └─......各类食物文件夹 |
每个类别单独放在一个文件夹内,文件夹名称就是类别名称,文件夹内存放该类别的图片。
我们需要实现:
- 遍历文件夹,自动读取图片路径并生成对应标签;
- 自定义 Dataset 类,加载图片与标签;
- 使用数据增强扩充训练样本;
- 通过 DataLoader 实现批量加载、送入 GPU 训练;
- 训练 CNN 模型,保存最优模型,最后单独拿一张图片做推理预测。
二、什么是 Dataset
torch.utils.data.Dataset是 PyTorch 提供的数据集抽象基类 ,用来定义数据集读取逻辑。
想要自定义数据集,必须重写 3 个核心方法:
- init:初始化,读取图片路径、标签、预处理变换;
- len:返回数据集总样本数量;
- getitem:按下标索引,返回单条样本(图像 + 标签)。
Dataset 的作用:解耦数据读取和模型训练,只负责定义如何获取单条样本,不负责分批、打乱。
三、数据增强
深度学习模型的训练需要充足的样本支撑,若数据集样本数量有限,模型极易死记硬背训练图像的像素细节,无法学习到物体的通用特征,最终导致在测试集上表现不佳,出现典型的过拟合问题。
数据增强是解决小样本训练、抑制过拟合的核心手段。在模型训练过程中,通过对原始图像进行随机几何变换、像素调整等操作,生成多样化的虚拟样本。该方式无需额外采集真实数据,即可极大丰富数据集的样本多样性,引导模型聚焦学习物体的核心轮廓、纹理等固有特征,而非图像的位置、角度、光影等无关干扰信息,有效提升模型的泛化能力与鲁棒性。
|---------------------------------------------------------------------|
| 注意:训练集使用随机增强,验证 / 测试集只用固定缩放,不能加随机操作 。测试阶段我们需要真实评估模型性能,不能随机修改图片。 |
常见数据增强种类
- Resize:统一缩放图片到固定尺寸,保证输入网络的图片大小一致;
- RandomRotation:随机旋转图片;
- RandomHorizontalFlip:随机水平翻转;
- RandomCrop / CenterCrop:随机裁剪、中心裁剪;
- ColorJitter:随机调整亮度、对比度、饱和度、色相;
- ToTensor:将 PIL 图片像素值从0,255转为0,1张量,同时调换通道顺序。
四、DataLoader
Dataset 只定义了单个样本怎么读取,而DataLoader 是加载器,在 Dataset 基础上实现批量打包。 DataLoader 的核心好处:
-
batch 批量打包:一次取出 batch_size 个样本,组成一个批次送入网络,充分利用 GPU 并行计算;
-
shuffle 打乱:训练集开启打乱,防止模型记住样本顺序,提升收敛效果;
-
多线程读取:可通过 num_workers 开启多进程,磁盘读取与 GPU 计算并行,加快数据加载速度;
-
自动堆叠张量,直接输出可以送入模型的批量图像张量和标签张量。
拓展:模型训练时,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 主程序入口:启动训练 + 单图推理测试
实验小结
本节课我们完成自定义图片数据集的完整工程流程:
- 通过 os.walk 遍历本地图片,生成图片路径与标签的索引文件;
- 继承 Dataset 自定义数据集类,实现图片读取;
- 使用 transforms 做数据增强,缓解小数据集过拟合;
- DataLoader 实现批量加载数据,配合 GPU 加速训练;
- 训练 CNN 并保存模型,最后实现单张图片推理。
这套流程是深度学习项目最常用的本地数据集开发范式,后续不管是图像分类、医学影像等任务,都可以复用这套 Dataset+DataLoader 的框架。