
一、引言当数据不再是现成的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, transformNone): 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], dtypenp.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_size64, shuffleTrue) test_dataloader DataLoader(test_data, batch_size64, shuffleTrue)DataLoader的作用批量读取每次返回batch_size个样本减少内存占用。打乱顺序shuffleTrue每个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×256conv1ConvReLUPool16×128×128conv2双层ConvReLUPool64×64×64conv3ConvReLU128×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(fAccuracy: {100*correct/size}%)配置损失函数nn.CrossEntropyLoss()优化器torch.optim.Adam(model.parameters(), lr0.001)训练轮数10九、总结本篇博客通过一个完整的食物分类项目讲解了深度学习中自定义数据集的完整流程知识点核心内容数据索引遍历目录生成train.txt/test.txt每行“路径 标签”魔法方法__getitem__支持索引__len__支持len()自定义Dataset继承Dataset实现__init__、__len__、__getitem__数据变换transforms.Compose组合 Resize 和 ToTensorDataLoader批量加载、打乱、并行CNN模型针对 3×256×256 输入输出 20 类训练测试标准训练循环与评估关键收获自定义数据集让PyTorch能够处理任意格式的数据。DataLoader负责高效的批量数据供给。数据预处理和增强是提升模型性能的重要手段。CNN的通道数递增、空间尺寸递减是经典设计模式。