一、引言:当数据不再是现成的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的通道数递增、空间尺寸递减是经典设计模式。