初识深度学习——DataLoader
一、引言:当数据不再是现成的MNIST
在前两篇博客中,我们使用PyTorch内置的MNIST数据集完成了手写数字识别。MNIST的好处是开箱即用——datasets.MNIST一行代码就帮我们下载、解析、转换好了数据。
但在实际项目中,我们面对的数据往往是自己的图片文件夹,比如一个食物分类数据集,结构可能长这样:
food_dataset/
├── train/
│ ├── pizza/
│ │ ├── 001.jpg
│ │ └── 002.jpg
│ ├── sushi/
│ │ ├── 003.jpg
│ │ └── 004.jpg
│ └── ...
└── test/
├── pizza/
└── sushi/
这时候,我们就需要自定义数据集——告诉PyTorch如何读取这些图片、如何对应标签、如何做预处理。本篇博客将基于一份完整的代码,讲解如何从零构建自定义数据集,并用CNN完成食物分类任务。
二、自动生成数据索引文件
在自定义数据集之前,我们首先需要一份“清单”,告诉程序每张图片的路径和对应的标签。
2.1 遍历目录生成索引
代码中的 train_test_file 函数完成了这个任务:
import os
def train_test_file(root, dir):
file_txt = open(dir + '.txt', 'w')
path = os.path.join(root, dir)
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()
逻辑解析:
-
os.walk(path)递归遍历目录,返回(当前路径, 子目录列表, 文件列表)。 -
当
directories非空时,说明当前是类别文件夹的上一级(如train/),此时dirs保存所有类别名称(如['pizza', 'sushi', ...])。 -
当
directories为空时,说明当前是具体的类别文件夹(如train/pizza/),此时遍历其中的图片文件,写入一行:图片路径 标签。 -
标签通过
dirs.index(now_dir[-1])获得,即类别在列表中的索引(0, 1, 2, ...)。
运行后,会在当前目录生成 train.txt 和 test.txt,内容示例:
.\data\food_dataset\train\pizza\001.jpg 0
.\data\food_dataset\train\pizza\002.jpg 0
.\data\food_dataset\train\sushi\003.jpg 1
...
2.2 为什么需要索引文件?
-
解耦:数据集的读取逻辑与文件系统分离,方便后续修改。
-
灵活:索引文件可以是 TXT、CSV、JSON 等格式,适应不同场景。
-
可复现:固定索引文件后,每次训练使用相同的数据划分。
三、Python魔术方法:__getitem__ 与 __len__
在自定义数据集类之前,我们需要理解两个重要的魔术方法。
代码中有一个小示例:
class USE_getitem:
def __init__(self, text):
self.text = text
def __getitem__(self, index):
return self.text[index].upper()
def __len__(self):
return len(self.text)
p = USE_getitem("pytorch")
print(p[1]) # 输出 'Y',因为调用了 __getitem__
print(len(p)) # 输出 7,因为调用了 __len__
核心结论:
-
当对象实现了
__getitem__,就可以用obj[index]的形式访问。 -
当对象实现了
__len__,就可以用len(obj)获取长度。 -
PyTorch的
Dataset类正是依赖这两个方法来实现数据的索引和总数统计。
四、自定义数据集类:food_dataset
现在,我们基于 Dataset 构建自己的数据集类。
import torch
from torch.utils.data import Dataset
from PIL import Image
from torchvision import transforms
import numpy as np
class food_dataset(Dataset):
def __init__(self, file_path, transform=None):
self.file_path = file_path
self.imgs = []
self.labels = []
self.transform = transform
with open(self.file_path) as f:
samples = [x.strip().split(' ') for x in f.readlines()]
for img_path, label in samples:
self.imgs.append(img_path)
self.labels.append(label)
def __len__(self):
return len(self.imgs)
def __getitem__(self, idx):
image = Image.open(self.imgs[idx])
if self.transform:
image = self.transform(image)
label = torch.from_numpy(np.array(self.labels[idx], dtype=np.int64))
return image, label
三个关键方法:
| 方法 | 作用 | 说明 |
|---|---|---|
__init__ | 初始化 | 读取索引文件,将图片路径和标签分别存入 self.imgs 和 self.labels |
__len__ | 返回样本总数 | len(dataset) 时调用 |
__getitem__ | 返回第 idx 个样本 | dataset[idx] 时调用,返回 (image_tensor, label_tensor) |
注意:
-
Image.open()读取的是PIL图像,需要经过transform转为张量。 -
标签必须转为PyTorch张量(这里用
torch.from_numpy将整数转为int64张量),因为后续损失函数需要张量输入。
五、数据预处理与增强
data_transforms = {
'trainda': transforms.Compose([
transforms.Resize([256, 256]),
transforms.ToTensor(),
]),
'valid': transforms.Compose([
transforms.Resize([256, 256]),
transforms.ToTensor(),
]),
}
transforms.Compose 将多个变换组合在一起,按顺序执行。
| 变换 | 作用 |
|---|---|
Resize([256, 256]) | 将图像统一缩放到 256×256,保证输入尺寸一致 |
ToTensor() | 将PIL图像转为张量,并将像素值从 0-255 缩放到 0-1,同时把通道维度放到最前面(C×H×W) |
数据增强:虽然这里只用了缩放和转张量,但实际项目中可以加入随机裁剪、翻转、颜色抖动等操作,提升模型泛化能力。
六、DataLoader:批量加载数据
from torch.utils.data import DataLoader
train_dataloader = DataLoader(training_data, batch_size=64, shuffle=True)
test_dataloader = DataLoader(test_data, batch_size=64, shuffle=True)
DataLoader的作用:
-
批量读取:每次返回
batch_size个样本,减少内存占用。 -
打乱顺序:
shuffle=True每个epoch重新打乱,避免模型学到顺序规律。 -
并行加速:可通过
num_workers开启多进程加载。
七、CNN模型设计
针对 3×256×256 的彩色图像,模型定义如下:
class CNN(nn.Module):
def __init__(self):
super(CNN, self).__init__()
self.conv1 = nn.Sequential(
nn.Conv2d(3, 16, 5, 1, 2), # 16×256×256
nn.ReLU(),
nn.MaxPool2d(2), # 16×128×128
)
self.conv2 = nn.Sequential(
nn.Conv2d(16, 32, 5, 1, 2), # 32×128×128
nn.ReLU(),
nn.Conv2d(32, 64, 5, 1, 2), # 64×128×128
nn.ReLU(),
nn.MaxPool2d(2), # 64×64×64
)
self.conv3 = nn.Sequential(
nn.Conv2d(64, 128, 5, 1, 2), # 128×64×64
nn.ReLU(),
)
self.out = 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) # 展平
output = self.out(x)
return output
尺寸变化总结:
| 阶段 | 操作 | 输出尺寸 |
|---|---|---|
| 输入 | - | 3×256×256 |
| conv1 | Conv+ReLU+Pool | 16×128×128 |
| conv2 | 双层Conv+ReLU+Pool | 64×64×64 |
| conv3 | Conv+ReLU | 128×64×64 |
| 展平 | view | (batch, 128×64×64) |
| 输出 | Linear | (batch, 20) |
参数量估算:卷积层参数约 10 万,全连接层参数约 128×64×64×20 ≈ 1048 万,参数量较大,但仍在可接受范围。
八、训练与测试
训练和测试函数与之前类似,核心步骤:
def train(dataloader, model, loss_fn, optimizer):
model.train()
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()
def test(dataloader, model, loss_fn):
model.eval()
size = len(dataloader.dataset)
correct = 0
with torch.no_grad():
for x, y in dataloader:
x, y = x.to(device), y.to(device)
pred = model(x)
correct += (pred.argmax(1) == y).type(torch.float).sum().item()
print(f"Accuracy: {100*correct/size}%")
配置:
-
损失函数:
nn.CrossEntropyLoss() -
优化器:
torch.optim.Adam(model.parameters(), lr=0.001) -
训练轮数:10
九、总结
本篇博客通过一个完整的食物分类项目,讲解了深度学习中自定义数据集的完整流程:
| 知识点 | 核心内容 |
|---|---|
| 数据索引 | 遍历目录生成 train.txt / test.txt,每行“路径 标签” |
| 魔法方法 | __getitem__ 支持索引,__len__ 支持 len() |
| 自定义Dataset | 继承 Dataset,实现 __init__、__len__、__getitem__ |
| 数据变换 | transforms.Compose 组合 Resize 和 ToTensor |
| DataLoader | 批量加载、打乱、并行 |
| CNN模型 | 针对 3×256×256 输入,输出 20 类 |
| 训练测试 | 标准训练循环与评估 |
关键收获:
-
自定义数据集让PyTorch能够处理任意格式的数据。
-
DataLoader负责高效的批量数据供给。 -
数据预处理和增强是提升模型性能的重要手段。
-
CNN的通道数递增、空间尺寸递减是经典设计模式。
更多推荐


所有评论(0)