ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

深入解析load_data_fashion_mnist:PyTorch数据加载与Fashion-MNIST实战

深入解析load_data_fashion_mnist:PyTorch数据加载与Fashion-MNIST实战 学《动手学深度学习》这一系列内容的人没有一个能绕过load_data_fashion_mnist这个名字。第一次碰到它你大概率是照着书里敲一行from d2l import torch as d2l然后调用这个函数拿train_iter和test_iter接着就稀里糊涂开始训练了。但如果你只停留在能用的层面后面做模型改造、调参、上线部署时一定会回来补课。这篇文章我就把这个函数从头到尾拆开给你看。你会知道Fashion-MNIST到底是个什么数据集load_data_fashion_mnist在背后做了哪几件事源码里每一行为什么要这样写以及在实际使用中你会踩到哪些典型坑位。无论你是刚入门深度学习的新手还是已经能跑通几个模型但想补齐数据加载这块短板的同学这篇文章都值得你花十分钟读完。1. Fashion-MNIST到底是个什么数据集1.1 为什么入门都从它开始先说结论Fashion-MNIST是传统MNIST数据集的现代替代品。MNIST是手写数字识别曾经是深度学习界的hello world但今天它已经太简单了——随便一个线性模型加个隐藏层就能跑到95%以上的准确率卷积神经网络更是能刷到99%以上。问题在于数字识别的特征太单一背景干净、笔画规整、类别容易被区分你用MNIST训练出来的模型很难迁移到真实的图像场景里。Fashion-MNIST由Zalando公司基于电商平台的商品图片制作一共10个类别都是日常服饰T恤、裤子、套头衫、连衣裙、外套、凉鞋、衬衫、运动鞋、包包、靴子。每张图片是28x28像素的灰度图训练集60000张测试集10000张。整体难度比MNIST高了一个档次因为衣服类别的类内差异很大比如不同款式的衬衫长得完全不一样类间又有相似性比如套头衫和衬衫、运动鞋和凉鞋容易混淆这更接近真实图像识别任务的状态。我特别想强调一点Fashion-MNIST并没有用更复杂的图像尺寸或者彩色通道来增加难度它坚持28x28灰度图这个设计非常聪明。对初学者来说输入维度越小模型迭代越快调试成本越低而图像内容本身又有足够的判别难度不至于让模型躺赢。等你能在这个数据集上把卷积神经网络、批量归一化、数据增强这些技巧都跑通一遍再换到CIFAR-10或者ImageNet这种更复杂的数据集思路是完全一样的。1.2 数据下载与本地目录结构当你第一次调用load_data_fashion_mnist时程序会联网下载数据集文件。具体来说它内部走的是torchvision.datasets.FashionMNIST这个类下载完成后会解压并缓存到本地目录。不同版本的《动手学深度学习》代码里root参数的默认值略有不同常见的是../data或者用户目录下的.d2l文件夹目的是一样的把数据集中管理起来避免每次运行都重新下载。这里有一个值得留意的细节。torchvision的数据集接口都遵循同一套设计模式先检查本地是否已有对应文件有就直接加载没有才发起下载。这意味着你第一次运行时会感觉特别慢取决于网络状况但第二次运行就会快很多。如果你仔细观察下载过程会发现数据集文件其实是四个.gz压缩包训练图像、训练标签、测试图像、测试标签文件格式是IDX也就是一种简单的二进制数组格式需要用torchvision内置的reader来解析。有些同学会问能不能自己写代码读取IDX文件当然可以但你没必要重复造轮子。FashionMNIST这个类已经帮你把解压、解析、切片、按索引取样本这些琐事全部封装好了你要做的只是传入transform参数告诉它拿到原始图片后要怎么处理。1.3 torchvision数据集接口的通用逻辑FashionMNIST不是特例MNIST、CIFAR-10、CIFAR-100、ImageFolder、DatasetFolder都继承了torch.utils.data.Dataset这个基类。所以你只要吃透FashionMNIST的用法后面遇到任何torchvision内置数据集都能很快上手。Dataset的核心契约有两个一是能返回数据集的大小也就是实现__len__方法二是能根据索引返回一个样本对(image, label)也就是实现__getitem__方法。至于数据从哪里来、要不要做预处理Dataset内部不管这些都由具体的子类去完成。理解这个接口设计很重要。因为DataLoader在训练循环里实际上是不断调用dataset[i]来取样本的然后把取出来的样本打包成一个batch返回给你。如果你以后要加载自己的业务数据比如一堆病理切片、工业质检图片你只需要写一个自定义类继承Dataset在__getitem__里返回(图片张量, 标签张量)剩下的DataLoader调度逻辑完全不用改。2. 手把手拆解load_data_fashion_mnist源码2.1 函数签名与返回值设计我们先看《动手学深度学习》PyTorch版中这个函数的经典实现def load_data_fashion_mnist(batch_size, resizeNone): 下载Fashion-MNIST数据集然后将其加载到内存中 trans [transforms.ToTensor()] if resize: trans.insert(0, transforms.Resize(resize)) trans transforms.Compose(trans) mnist_train torchvision.datasets.FashionMNIST( root../data, trainTrue, transformtrans, downloadTrue) mnist_test torchvision.datasets.FashionMNIST( root../data, trainFalse, transformtrans, downloadTrue) return (torch.utils.data.DataLoader(mnist_train, batch_size, shuffleTrue, num_workersget_dataloader_workers()), torch.utils.data.DataLoader(mnist_test, batch_size, shuffleFalse, num_workersget_dataloader_workers()))函数签名只有两个参数batch_size和resize。batch_size决定每个batch装多少张图resize是可选的用来统一缩放图片尺寸。返回值是两个DataLoader对象一个用于训练打乱顺序一个用于测试不打乱顺序。为什么返回DataLoader而不直接返回Dataset因为训练时我们通常要小批量地取数据DataLoader把取数据-打包-多进程预取这套流程全部接管了。你的训练循环里直接for X, y in train_iter就行每次拿到的X是一个四维张量形状是(batch_size, 1, 28, 28)y是一维张量形状是(batch_size,)。2.2 ToTensor到底做了什么transforms.ToTensor()是这个函数里最重要的一行。它做了两件事缺一不可。第一件事是把图像数据从HWC布局转成CHW布局。原始图片在读取后通常是(高, 宽, 通道)的排列但PyTorch的卷积层、池化层约定输入是(通道, 高, 宽)所以必须在数据进入模型之前完成维度调换。如果你自己写训练代码时忽略了这一步模型运行时会直接报维度不匹配的错误。第二件事是把像素值从0到255的整数缩放到0到1之间的浮点数。神经网络在训练时对输入的数值范围非常敏感如果直接把255这种量级的数值喂给模型初始梯度很容易爆炸损失函数也难以收敛。ToTensor会把每个像素除以255得到范围在[0,1]的float32张量这样模型训练会更稳定。有些同学会想那要不要顺便做标准化让数据分布接近均值为0、方差为1load_data_fashion_mnist默认不做因为Fashion-MNIST灰度图的分布相对规整对简单模型影响不大。但如果你在做真实项目尤其是图像内容复杂、光照差异大的场景我强烈建议你加上transforms.Normalize((mean,), (std,))提前计算好数据集的均值和标准差能显著加速收敛。2.3 resize参数背后的模型输入适配为什么load_data_fashion_mnist要留一个resize参数因为不是所有模型都接受28x28的输入。比如你后面学到AlexNet时它设计的输入是224x224如果你直接用原始Fashion-MNIST的28x28图去喂模型结构里第一层卷积的尺寸就完全对不上。函数里用了transforms.Resize(resize)把它插入到ToTensor之前。这个顺序是有讲究的。Resize期望的输入是PIL Image或者numpy数组它负责像素插值缩放而ToTensor期望的输入是PIL Image或numpy数组输出Tensor。如果顺序反了先把图变成Tensor再去做Resize你会发现很多操作根本没法执行因为Resize不认Tensor输入除非你在代码里额外做了适配。关于Resize本身有两个细节值得注意。第一transforms.Resize((h, w))会强制把图片拉伸到指定宽高这可能改变原始长宽比transforms.Resize(224)则会把短边缩放到224同时保持长宽比。第二对于分类任务纯粹拉伸通常也能接受但如果你做目标检测长宽比改变会直接影响标注框的坐标那时候就需要结合CenterCrop或Pad来保持几何关系。2.4 DataLoader四个关键参数逐个说torch.utils.data.DataLoader(mnist_train, batch_size, shuffleTrue, num_workersget_dataloader_workers())这个调用里有几个参数每个都有讲究。shuffle训练集设为True每个epoch开始前都重新打乱数据顺序避免模型学到样本顺序带来的伪规律测试集设为False因为评估时你不关心顺序而且保持顺序方便调试和可视化。num_workers指定用几个子进程来预取数据。如果你的机器是多核CPU把这个值设大一些数据加载和模型训练可以重叠进行不会出现GPU等待CPU传数据的空闲期。《动手学深度学习》里的get_dataloader_workers()在Linux/macOS下通常返回CPU核心数在Windows下会返回0原因是Windows的多进程启动机制和Linux不同盲目设大反而容易报错这个我在第4章会展开讲。pin_memory这个参数在很多官方示例里没有明确写出来但实际训练中很有用。它会把数据放到锁页内存page-locked memory里CPU到GPU的拷贝速度会快很多。如果你的机器有NVIDIA显卡建议训练时加上pin_memoryTrue如果在CPU上训练这个参数没有意义保持默认即可。drop_last最后一个batch如果不足batch_size默认情况下它会保留这可能让梯度更新时batch大小不一致。对Fashion-MNIST来说60000能被常见batch size整除所以影响不大但如果你在大数据集上做分布式训练通常会设drop_lastTrue保证每个batch大小完全一致避免某些框架对batch size一致性敏感的问题。3. 实操跑通数据加载并可视化第一批样本3.1 环境准备与版本检查在跑load_data_fashion_mnist之前先确认你的环境没问题。我的建议是直接用conda建一个独立环境Python版本用3.9或3.10然后安装PyTorch、torchvision、d2l。conda create -n d2l python3.9 -y conda activate d2l pip install torch torchvision pip install d2l装好后跑一下版本检查python -c import torch, torchvision; print(torch.__version__, torchvision.__version__)能正常输出版本号就说明核心依赖没问题。需要提醒的是d2l这个包更新比较频繁不同版本里load_data_fashion_mnist的实现可能有细微差异但核心逻辑一致。如果你的版本比较新源码里可能已经改用d2l.Downloader来管理数据缓存路径这都不影响你按本文的思路去理解。3.2 首次加载会发生什么接着我们写一个最基础的调用from d2l import torch as d2l batch_size 256 train_iter, test_iter d2l.load_data_fashion_mnist(batch_size)运行到这一行如果本地没有数据程序会开始下载。下载完成后你可以先打印一下数据形状确认没问题for X, y in train_iter: print(X.shape, X.dtype, y.shape, y.dtype) break输出会是这样torch.Size([256, 1, 28, 28]) torch.float32 torch.Size([256]) torch.int64X的第一维是batch size 256第二维是通道数1灰度图后面两维是高和宽28x28。y是长度为256的整型张量每个值对应一个类别编号。我建议你养成拿到数据先看形状和dtype的习惯这能帮你快速判断数据流哪里出了问题。3.3 自己写一个样本可视化函数d2l包里提供了show_images函数但为了加深理解我建议你手动实现一个简易版本import matplotlib.pyplot as plt import torch def show_images(imgs, num_rows, num_cols, titlesNone, scale1.5): figsize (num_cols * scale, num_rows * scale) _, axes plt.subplots(num_rows, num_cols, figsizefigsize) axes axes.flatten() for i, (ax, img) in enumerate(zip(axes, imgs)): if isinstance(img, torch.Tensor): img img.detach().numpy() ax.imshow(img, cmapgray) ax.axes.get_xaxis().set_visible(False) ax.axes.get_yaxis().set_visible(False) if titles: ax.set_title(titles[i]) plt.show()注意一点train_iter里拿到的X是四维张量(256, 1, 28, 28)但imshow需要的是二维或者三维数据所以要先把多余的通道维去掉。你可以用X[i].reshape(28, 28)或X[i].squeeze()来处理batch next(iter(train_iter)) show_images(batch[0][:10].reshape(10, 28, 28), 2, 5, titles[str(label.item()) for label in batch[1][:10]])这样做能一次性看到10张图每张图上方标的是类别编号。第一次看时你会直观感受到不同衣服类别之间的差异确实存在但有些类别比如套头衫和衬衫连人眼都可能看走眼这对模型来说就不是一个简单任务。3.4 把编号转成可读的类别名如果你觉得只管看数字不够直观可以做一个标签映射表。Fashion-MNIST官网给出的类别顺序是固定的标签值类别英文名常见翻译0t-shirtT恤1trouser裤子2pullover套头衫3dress连衣裙4coat外套5sandal凉鞋6shirt衬衫7sneaker运动鞋8bag包9ankle boot短靴翻译成中文后你可能会立刻发现一个有趣的现象第0类和第6类一个叫T恤一个叫衬衫视觉上确实容易让人混淆这正是Fashion-MNIST比MNIST难度更高的原因之一。你可以把titles参数从数字标签改成真实类别名再可视化观察一下模型要面对的难度分布。4. 常见问题与排查技巧实录4.1 下载卡住、超时或者进度条一直不动这是初学者遇到最多的问题。Fashion-MNIST的原始数据托管在GitHub上有些网络环境访问GitHub下载很慢甚至直接卡在0%。解决办法有好几种。第一种是设置镜像。torchvision在下载时走的是URL你可以提前下载好四个.gz压缩包然后放到本地目录再修改root参数指向那个目录。要特别注意如果数据集已经被部分下载但没解压成功建议先把目录清空再手动放文件避免文件冲突。第二种是使用学术资源镜像站下载压缩包。下载后你需要把它放到root参数指定的路径下并且保证文件名和torchvision期望的文件名一致。常见的文件包括train-images-idx3-ubyte.gz train-labels-idx1-ubyte.gz t10k-images-idx3-ubyte.gz t10k-labels-idx1-ubyte.gz文件放好后把download参数改成False再运行程序会直接解压本地文件不再联网。另外提醒一下有些同学把root指向../data结果在不同目录下运行代码时各下载了一遍数据。我建议你在项目里固定一个绝对路径比如os.path.join(os.path.expanduser(~), .d2l, data)这样不管从哪个脚本启动都能复用同一份数据缓存。4.2 Windows系统下的num_workers报错在Windows下如果你的num_workers设置成大于0的数运行到迭代DataLoader时可能直接报错弹出类似这样的信息RuntimeError: DataLoader worker (pid(s) 12345) exited unexpectedly原因是Windows不像Linux那样使用fork方式创建子进程而是用spawn方式它会重新导入主模块。如果你的数据加载代码没有被if __name__ __main__:保护子进程导入时会递归执行最终崩溃。解决方案有三个按优先级推荐。第一最省事的办法是把num_workers设为0让数据加载在主进程里完成Fashion-MNIST这种百万张以下的小数据集0 workers的速度完全可以接受。第二把调用load_data_fashion_mnist和训练循环的代码统一放到if __name__ __main__:块里。第三如果你确实需要多进程加速可以在代码顶部设置torch.multiprocessing.freeze_support()但整体复杂度会高一些。4.3 resize之后图片变形严重有同学把resize设为(224, 224)后发现图像变得很扁或者很胖原因是Resize((h, w))会对图像做非等比拉伸。如果衣服的原始宽高比是1:1Fashion-MNIST本来就是方形(224, 224)其实不会变形但如果你换到其他数据集上比如宽高比是4:3的照片强拉伸就会失真。在深度学习中正方形输入是很多经典卷积网络的要求所以一个更稳妥的组合是transforms.Compose([ transforms.Resize(224), transforms.CenterCrop(224), transforms.ToTensor() ])先用Resize(224)保持长宽比把短边缩放到224再用CenterCrop(224)从中心裁出224x224区域。这样虽然会损失边缘信息但不会带来几何形变。在图像分类任务里这个组合是一个非常常见的预处理方案。4.4 维度对不上、图片显示异常有时候你明明调用成功了但训练时模型报维度错误常见原因是transform顺序写错了。前面提到过Resize要放在ToTensor之前因为Resize不认Tensor输入。如果你在Compose里先写了ToTensor再写Resize运行时会报错或者产生不可预期的行为。另一个容易踩的坑是可视化时维度乱掉。X是(batch, channel, height, width)但matplotlib的imshow默认接收(height, width)或(height, width, channel)。如果你直接把X[i]传进去会看到类似Invalid shape的报错或者显示出来的图像颜色完全错乱。记住用squeeze()去掉通道维灰度图才能正常显示。4.5 常见问题速查表为了让你排查时更顺手我把上面这些典型问题整理成一个表格现象可能原因解决方案下载进度条一直为0%网络访问源站不通手动下载压缩包放到root目录改downloadFalseWindows下worker进程崩溃spawn进程模型导致递归导入num_workers0或加if __name__ __main__保护图片变形、拉伸严重Resize((h, w))强制拉伸改用Resize(224) CenterCrop(224)维度不匹配报错没做CHW转换或transform顺序写反确认ToTensor在Resize之后检查数据形状训练速度慢GPU利用率低CPU数据加载成为瓶颈调大num_workers开启pin_memoryTrue每个epoch结果不一致数据没打乱或随机种子未固定确认训练集shuffleTrue必要时设torch.manual_seed4.6 验证数据链路是否正常遇到疑难杂症时我习惯做一个小步快跑验证。先不加载完整数据集而是只取前一小部分样本检查它的数值范围、类型、形状是否符合预期。比如你可以写一行代码打印X.min()和X.max()正常情况下经过ToTensor后会在0到1之间如果你发现最大值为255说明数据根本没有经过ToTensor问题大概率出在transform没有传进Dataset。另外一个非常实用的小技巧是把一个batch的样本和标签都打印出来肉眼核对一下标签和图像内容是否匹配。很多时候模型训练效果差不是模型结构的问题而是数据标签错位、数据增强过度、或者归一化参数不对。我自己在实际项目里会把load_data_fashion_mnist重构成更通用的工具函数把root、transform、dataset_type这些参数都抽出去让它能同时加载MNIST、CIFAR-10和我自己的业务数据。这样做的好处是整个项目的数据加载代码只有一份后面换数据集、改预处理只需要改配置不用动训练循环。最后再分享一个小技巧如果你用了Resize等数据增强操作调试时可以先把增强关掉只保留ToTensor看模型能不能正常收敛。如果模型在原始数据上表现正常加了增强后反而变差那问题大概率出在增强的强度上比如旋转角度太大、裁剪比例太小、亮度扰动太狠。这个排查方法是我踩过很多次坑之后总结出来的远比对着报错信息一行行猜效率高。
返回列表