一、引言:当数据不再是现成的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.txttest.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.imgsself.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
conv1Conv+ReLU+Pool16×128×128
conv2双层Conv+ReLU+Pool64×64×64
conv3Conv+ReLU128×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 类
训练测试标准训练循环与评估

关键收获

  1. 自定义数据集让PyTorch能够处理任意格式的数据。

  2. DataLoader 负责高效的批量数据供给。

  3. 数据预处理和增强是提升模型性能的重要手段。

  4. CNN的通道数递增、空间尺寸递减是经典设计模式。

更多推荐