news 2026/9/14 3:36:52

初识深度学习——DataLoader

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
初识深度学习——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.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的通道数递增、空间尺寸递减是经典设计模式。

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/14 3:36:36

信创环境下SNMP协议栈选型:从Net-SNMP到国产自研SDK的实践思考

1. 信创改造现场,SNMP采集模块是怎么"崩溃"的1.1 一个真实迁移场景:从x86CentOS到ARM国产OS前段时间帮客户做网管系统的信创适配,其中一块工作就是SNMP采集模块的迁移。客户原来的架构很简单:网管服务器跑在x86 CentOS…

作者头像 李华
网站建设 2026/9/14 3:36:07

VS Code实现Typora级Markdown可视化编辑

简介:这是一款专为 VS Code 用户打造的 Markdown 增强插件,面向前端开发者、技术文档撰写者及轻量级内容创作者,旨在解决原生编辑器在可视化编辑、实时预览与富媒体支持方面的短板。插件提供 Typora 级别的流畅体验:支持表格可视化…

作者头像 李华
网站建设 2026/9/14 3:36:04

BP神经网络PID控制在Simulink中的实现与优化

1. 项目概述:BP神经网络PID控制在Simulink中的实现价值在工业控制领域,PID控制器因其结构简单、鲁棒性好等特点被广泛应用,但面对非线性、时变系统时,传统PID参数整定往往显得力不从心。我在某型海洋压力模拟设备的开发中就遇到过…

作者头像 李华
网站建设 2026/9/14 3:33:19

SpringBoot环保网站开发:毕业设计实战指南

1. 项目概述与核心价值这个基于SpringBoot的环境保护宣传网站项目,本质上是一个典型的计算机专业毕业设计解决方案包。它包含了从技术实现到论文撰写的完整闭环,特别适合需要快速搭建环保类Web应用的学生开发者。我经手过二十多个类似项目,这…

作者头像 李华