PyTorch 读取 MNIST 数据管道与 DataLoader 参数解析
2026/9/18 6:08:00 网站建设 项目流程

第一次把 PyTorch 装好,兴致勃勃敲下两行代码准备跑通 MNIST 数据集,结果终端里那个进度条要么停在 0% 一动不动,要么直接甩出一个 404,这种场面我见得太多。很多人对"读取 MNIST"这件事的预期是十分钟,最后却花了一整个下午在排查网络、路径和参数。问题不在于 MNIST 有多难,而在于大多数人只记住了datasets.MNIST(root, train=True, download=True)这一行,却完全不知道这行代码背后发生了什么。这篇内容想做的事情很简单:把这十分钟真正该花的时间,花在刀刃上——先搞清楚数据从哪来、长什么样、经过哪些处理,再去写模型。

适合读这篇的人有三类:刚配完 Anaconda 和 PyTorch 环境、准备跑第一个训练脚本的新手;能跑通代码但说不清transformDataLoader每个参数含义的中间层;以及准备从 MNIST 迁移到自己的图像数据集、想先摸清标准套路的开发者。下面所有的代码都在 CPU 环境下验证过,GPU 环境除了设备号之外没有差别。

1. MNIST 在 PyTorch 入门路径里的真实位置

1.1 为什么第一课永远是它

MNIST 是 6 万张 28×28 的手写数字灰度图,加 1 万张测试图,共 10 个类别。它的价值不在于"能训练出多好的模型",而在于它把深度学习的完整链路压缩到了最小的规模:数据读取、张量变换、批处理、前向传播、损失计算、反向传播、验证,每一步都有,但每一步的代价都很低。你在一台没有独立显卡的笔记本上,几十秒就能跑完一个 epoch,这意味着你可以用极低的试错成本去理解流程。

从工程角度看,这个数据集还有一个被低估的优点:格式极其规整。所有图像尺寸统一、通道数统一、标签是 0 到 9 的整数,没有任何脏数据需要清洗。你在真实项目里遇到的那些麻烦——图像大小不一、标注缺失、类别极度不平衡、文件夹结构混乱——它一个都没有。所以它是一个纯粹的"管道练习器",你在这里练的是数据流通的路径,而不是数据清洗的技巧。

我自己的经验是,把 MNIST 当作"读取管道"来练,比当作"分类任务"来练收益大得多。因为模型结构网上到处都有现成的,但数据管道出问题时的排查思路,很少有人系统讲过。

1.2 四个压缩文件里到底装了什么

当你执行下载时,实际拿到的是四个.gz压缩包,分别对应训练图像、训练标签、测试图像、测试标签。它们是经典的 IDX 格式,一种非常古老的二进制格式:头部有一段魔数用来标识维度信息,接着是数据本体。图像数据是连续的uint8,标签数据是连续的uint8

格式细节对使用者来说不重要,但有一个数字必须记住:训练集 60000 张,测试集 10000 张,每张 784 个像素点。因为后面你写DataLoader的时候,len(train_dataset)应该等于 60000,len(test_dataset)应该等于 10000。这两个数字如果不匹配,说明文件下载不完整或者被截断了,这是第一时间就能发现问题的自检点。

torchvision在首次读取时会把这些压缩包解压成idx格式的裸文件,之后每次运行都直接读裸文件,不再重复解压。所以你的root目录下通常会看到raw子目录,里面同时存在.gz和已解压的文件。这个细节很重要:如果你想迁移或者备份数据集,把整个raw目录打包带走,放到新机器的相同相对位置,就能实现零下载运行。

1.3 十分钟的时间预算应该怎么分配

很多人以为"十分钟搞懂"的意思是十分钟能从零跑到训练完一个 epoch。实际上,如果环境已经配好,下载速度正常,这两件事确实能在十分钟内完成。但如果把环境搭建算进来,那十分钟是远远不够的——装 PyTorch 本身就要考虑版本匹配、源的速度、CPU 还是 GPU 版本这些事。

所以我建议把这十分钟拆成三块:第一块大约是两分钟,用来把数据落到本地并且验证文件完整性;第二块是四分钟,用来理解transformDataLoader的每个参数在干什么;第三块是四分钟,用来做数据体检和一次前向的冒烟测试。这三块做完,读取这件事就算彻底通了,剩下的是模型的事,和读取无关。

需要提前说一句:如果你的环境还没搭好,先别急着上 GPU 版本。CPU 版本安装体积小、依赖少、出错概率低,跑 MNIST 完全够用。等读取链路跑通、模型验证过一遍之后再换 GPU 环境,排查成本会低很多。

2. torchvision.datasets.MNIST 的加载链路拆解

2.1 三个必填参数和两个最容易漏的参数

这个类的构造参数看起来不多,但每一个都有讲究。root是数据根目录,它不是"文件路径"而是"目录路径",传错成文件名会直接报错。train是布尔值,True取 6 万张训练集,False取 1 万张测试集,注意它控制的是两个完全不同的文件,不是同一份数据的切片。download控制是否在本地找不到文件时发起网络请求。

真正容易被忽略的是transformtarget_transform。前者处理图像,后者处理标签。绝大多数教程只写transform,因为标签通常不需要变换,但如果你想做标签平滑、one-hot 编码或者标签偏移校正,就需要动target_transform。这两个参数的默认值都是None,此时返回的是原始 PIL 图像对象和整数标签。

还有一个隐藏参数是transform到底应该在哪一层生效。有一种说法是"把它放在DataLoader里性能更好",这在早期的某些版本里确实成立,但现在的常规做法是放在Dataset构造时传入。理由很简单:transform是数据源的一部分定义,和数据集绑定,放在DataLoader里会让DataLoader承担它不该承担的职责,代码可读性变差。

import torch from torchvision import datasets, transforms train_transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) train_set = datasets.MNIST( root="./data", train=True, transform=train_transform, download=False ) test_set = datasets.MNIST( root="./data", train=False, transform=train_transform, download=False )

注意:root="./data"是相对当前工作目录的路径。如果你在 PyCharm 里跑和在命令行里跑,工作目录可能不一样,导致数据被下载到两个地方。要么统一用绝对路径,要么确认清楚工作目录。

2.2 download=True 背后的下载逻辑与 404 的成因

download=True时,torchvision会先检查root/MNIST/raw下是否存在对应的文件。如果不存在,就发起请求下载。下载完成后会做一次校验,校验不通过会重新下载或者直接抛异常。

常见的 404 有几类成因,我按出现频率排一下。第一类是数据源地址本身发生了变化。早期的一些下载地址在特定网络环境下返回 404 是真实存在的情况,后来官方把默认地址切到了更稳定的对象存储镜像上。如果你用的是很老的torchvision版本,可能仍然指向旧的地址。解决办法是升级torchvision,或者手动把文件下载好放进raw目录。

第二类是路径拼接导致的问题。比如root指向了一个不存在的深层目录,或者路径里带了奇怪的字符,某些系统下会出现资源定位失败。第三类是代理或企业网络环境对请求做了拦截,返回了一个 HTML 错误页而不是二进制文件,torchvision拿到这个 HTML 后校验失败,报错信息看起来像是格式错误,实际是网络问题。这类问题最迷惑人,因为错误信息不会直接告诉你"你拿到的是网页"。

判断方法很简单:去raw目录看文件大小。正常的图像压缩包在 9MB 上下,标签包在几十 KB。如果某个文件只有几 KB 或者干脆是 0 字节,基本可以确定是下载环节出了问题。

2.3 本地已有数据时如何把网络这一环彻底绕开

最省事的做法是准备一份别人的raw目录压缩包,解压到对应位置,然后永远设download=False。这样脚本在任何网络环境下都能跑,也省掉了每次运行的检查开销。

另一种做法是提前把四个文件准备好,通过download=False让它直接读。这里有个坑:如果你只放了.gz而没放解压后的文件,某些版本下也能正常工作,因为它会自己解压;但如果你放的是解压后的文件而缺少.gz,在download=True的情况下它可能会认为文件不全,试图重新下载。所以最稳的状态是:raw目录里既有.gz也有解压后的文件,且download=False

再补充一个经验:团队协作或者做实验记录的时候,把数据集路径写成配置项,而不是硬编码在代码里。因为不同机器上的盘符、挂载点都可能不一样,硬编码路径会让你的脚本只能在你自己的电脑上跑。

3. transform 不是可有可无的装饰

3.1 ToTensor 到底做了哪三件事

很多人把ToTensor()当成一个"格式转换器",其实它一次性做了三件事,每一件都直接影响后续计算。

第一件是维度置换。PIL 图像在内存里的布局是高度、宽度、通道,也就是 HWC。而 PyTorch 的卷积操作期望的布局是通道、高度、宽度,即 CHW。ToTensor()会把这个顺序换过来,所以 MNIST 的一张图从(28, 28)变成(1, 28, 28)。这就是为什么你在调试时打印图像形状会看到那个 1,它不是多余的维度。

第二件是数值类型转换。原始像素是 0 到 255 的整数,ToTensor()会把它变成浮点数并且除以 255,缩放到 0 到 1 之间。这一步如果漏掉,输入就是 0 到 255 的大数值,配合默认学习率训练会非常不稳定,损失曲线要么爆炸要么纹丝不动。

第三件是封装成张量对象。转换之后你拿到的是torch.Tensor,可以参与自动求导、可以.to(device)、可以和模型输出直接做运算。如果跳过这一步,你拿到的是 PIL 对象,根本没法和模型交互。

提示:ToTensor()内部做的是除以 255,如果你的数据本身已经是张量并且已经归一化过,就不要再套一次,否则数值会被压缩到几乎为零。

3.2 Normalize 的 0.1307 和 0.3081 是怎么算出来的

这两个数字几乎在所有 MNIST 教程里都能看到,但很少有人解释它们的来源。它们是在整个训练集上统计出来的:先把所有像素缩放到 0 到 1,然后计算全体像素的均值和标准差,得到大约 0.1307 和 0.3081。

Normalize做的事情是对每个通道执行"减去均值、除以标准差"。也就是(x - mean) / std。这样处理后,数据分布被拉到均值接近 0、标准差接近 1 的状态。为什么这么做?因为神经网络的参数初始化通常假设输入是近似标准正态分布的,如果输入整体偏大或者偏小,梯度传播的效率会明显下降,收敛会变慢。

你可以自己验证这两个数字,代码只有三行:

import torch from torchvision import datasets raw = datasets.MNIST(root="./data", train=True, download=False) imgs = raw.data.float() / 255.0 # 形状 [60000, 28, 28] print(imgs.mean().item()) # 约 0.1307 print(imgs.std().item()) # 约 0.3081

这里要注意一个细节:严格来说应该只统计训练集,而不是把测试集也算进去,否则会引入轻微的数据泄漏。上面演示用的是训练集,这是正确做法。另外,这两个值只对 MNIST 成立。换成 FashionMNIST,均值约 0.2860、标准差约 0.3530,两者完全不一样,不能混用。同理,换成彩色数据集就是三个通道各自一组值。

3.3 训练集和测试集的处理必须区别对待

有一个非常常见但很隐蔽的错误:训练集和测试集用了同一个带数据增强的transform。比如训练时加了随机旋转、随机裁剪,测试时也套了同一个管道。这会导致测试结果不稳定——同一张图每次评估都可能得到不同的预测结果,你根本无法判断模型到底是好还是坏。

正确的处理方式是分成两条管道。训练管道可以包含数据增强,测试管道只保留确定性的操作,也就是ToTensorNormalize,一个随机的都不要有。

train_tf = transforms.Compose([ transforms.RandomRotation(10), transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) eval_tf = transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ])

对于 MNIST 这种字迹本来就很规整的数据集,我个人的建议是增强可以非常轻,甚至不加。加太多旋转和裁剪反而会让数字变形到无法辨认,比如把 6 转成 9 的边缘状态,模型学到的东西就被污染了。如果你想验证增强是否有效,就做对照实验,别凭感觉加。

4. DataLoader 的几个参数决定了你的训练效率

4.1 batch_size、shuffle 和 drop_last 的实际影响

batch_size决定每次送入模型的样本数。MNIST 这种小图上,32 到 128 都是常用区间。batch 太小会让梯度噪声大、训练震荡;batch 太大会让内存占用上升,而且每个 epoch 的更新次数变少,收敛可能需要更多轮。选一个中间值先跑通,再根据显存和收敛曲线微调,这是比较务实的顺序。

shuffle只在训练集上设为True,测试集一律False。原因有两层。第一层是训练时打乱能避免模型学到"样本顺序"这种无关规律,特别是在数据按类别排序的情况下。第二层是评估时的顺序固定,能让你复现出完全一致的结果,方便对比不同模型。

drop_last处理的是最后一个不完整批次。60000 除以 128 是 468 余 96,最后一批只有 96 个样本。如果模型里用了批归一化层,一个 96 样本的批次和一个 128 样本的批次统计特性不同,可能造成轻微波动。设drop_last=True会丢掉这 96 个样本,代价是每个 epoch 少看一点数据。对 MNIST 来说这点损失可以忽略,所以我一般开着。

4.2 num_workers 与 persistent_workers 的取舍

num_workers是最容易出问题的一个参数。它指定用几个子进程来加载数据,0 表示在主进程里加载。设成大于 0 可以并行读取,理论上更快。

但实际效果取决于任务复杂度。如果在transform里只做了ToTensorNormalize,计算量极小,多进程的开销反而可能比收益大,num_workers设成 2 或者 4 就足够,设成 16 只会让启动变慢。真正需要多进程的是那些在transform里做复杂解码的图像任务,或者磁盘 IO 很慢的场景。

persistent_workers=True的含义是每个 epoch 结束后不销毁工作进程,下一个 epoch 直接复用。这能省掉反复创建进程的开销,但必须配合num_workers > 0使用,否则会直接报错。

注意:在 Windows 上使用多进程加载,你的训练代码必须包在if __name__ == "__main__":保护块里,否则子进程会重新导入主模块,导致无限递归创建进程。这个坑在 Linux 上不存在,所以从 Linux 迁到 Windows 的人几乎都会踩一次。

4.3 pin_memory、collate_fn 与迭代器的语义

pin_memory=True的作用是把批数据放进锁页内存,这样从内存拷贝到 GPU 显存时能走更快的通道。只有在使用 GPU 训练时才有意义,纯 CPU 环境下设了不会报错但也没有收益。

collate_fn决定了怎么把一批单独的样本拼成一个批次。默认实现会做两件事:把一批(C, H, W)的图像沿第 0 维堆叠成(B, C, H, W),把一批标量标签堆叠成(B,)。所以从DataLoader里取出一个批次,你会得到两个张量,形状分别是(B, 1, 28, 28)(B,)。如果你需要序列任务那种不等长填充,就必须自己写collate_fn,但在 MNIST 上完全用不到。

还有一个语义要搞清楚:DataLoader是可迭代对象,不是列表。每次for循环都会重新从数据集开头迭代一遍,所以你可以放心地在每个 epoch 用同一个DataLoader。但如果你用iter()手动拿了一个迭代器然后想复用,那是行不通的,迭代器耗尽就没了。

from torch.utils.data import DataLoader train_loader = DataLoader( train_set, batch_size=128, shuffle=True, num_workers=2, drop_last=True, pin_memory=False ) imgs, labels = next(iter(train_loader)) print(imgs.shape, labels.shape, imgs.dtype)

5. 一条完整的读取链路与数据体检

5.1 从文件到批次的完整代码

把前面所有内容串起来,一份能直接跑的脚本大概长这样。写得比我平时用的稍微啰嗦一点,是为了让每一步都看得见。

import torch from torchvision import datasets, transforms from torch.utils.data import DataLoader def build_loaders(data_root="./data", batch_size=128): train_tf = transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) eval_tf = transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) train_set = datasets.MNIST(root=data_root, train=True, transform=train_tf, download=False) test_set = datasets.MNIST(root=data_root, train=False, transform=eval_tf, download=False) train_loader = DataLoader(train_set, batch_size=batch_size, shuffle=True, num_workers=2, drop_last=True) test_loader = DataLoader(test_set, batch_size=batch_size, shuffle=False, num_workers=2) return train_set, test_set, train_loader, test_loader if __name__ == "__main__": tr_set, te_set, tr_loader, te_loader = build_loaders() print("train size:", len(tr_set), "test size:", len(te_set)) for imgs, labels in tr_loader: print("batch:", imgs.shape, labels.shape) break

这段代码里有两个刻意的设计。一个是我把构建逻辑封装成函数,这样换数据集、换 batch size 都只改参数不改逻辑。另一个是if __name__ == "__main__":保护块,即使现在num_workers=2也能在任何系统上安全运行,不会因为多进程产生问题。

5.2 上手先做的四项数据体检

跑通之后别急着写模型,先花两分钟做体检。这四件事能帮你提前发现九成以上的数据问题。

第一项,检查样本数量。len(train_set)应该是 60000,len(test_set)应该是 10000。数字不对说明文件有问题。

第二项,检查单张样本的形状和数值范围。取train_set[0],你会拿到一个元组,第一个元素形状是(1, 28, 28),第二个元素是 0 到 9 的整数。第一个元素的值应该大致落在 -0.4242 到 2.8215 之间。这个区间是算出来的:(0 - 0.1307) / 0.3081(1 - 0.1307) / 0.3081。如果你的值域超出这个范围,说明归一化参数用错了。

第三项,检查标签分布。用torch.bincount(torch.tensor(train_set.targets))看一下每类有多少张。MNIST 大致均衡,每类在 5400 到 6800 之间。如果某类数量明显异常,说明数据被破坏过。

第四项,可视化一批。把批次里的图反归一化回来再显示:乘上标准差加上均值,然后压缩掉通道维。这一步是我最推荐的,因为很多预处理错误肉眼一看就发现,比如图像上下颠倒、数值全黑、标签和图像对不上。

import torch, matplotlib.pyplot as plt imgs, labels = next(iter(tr_loader)) mean, std = 0.1307, 0.3081 show = imgs[:8] * std + mean fig, axes = plt.subplots(1, 8, figsize=(12, 2)) for ax, img, lb in zip(axes, show, labels[:8]): ax.imshow(img.squeeze(0), cmap="gray") ax.set_title(str(lb.item())) ax.axis("off") plt.show()

5.3 划分训练集与验证集的两种写法

严格来说 MNIST 并没有官方验证集,只有训练集和测试集。把测试集当验证集调参是一种常见但有争议的做法,因为调参多了之后测试集就不再"干净"。规范一点的做法是从训练集里切出一部分当验证集。

第一种写法是random_split,简单直接。但要注意它每次运行的结果都不一样,做实验对比时必须固定随机种子,否则两次实验的验证集不同,指标没法比较。

from torch.utils.data import random_split g = torch.Generator().manual_seed(42) val_size = 5000 train_size = len(train_set) - val_size tr_sub, val_sub = random_split(train_set, [train_size, val_size], generator=g)

第二种写法是用Subset配合固定的索引数组。这样做的好处是索引可以保存下来,跨机器、跨实验完全可复现。

from torch.utils.data import Subset import numpy as np idx = np.random.RandomState(42).permutation(len(train_set)) val_idx, tr_idx = idx[:5000], idx[5000:] tr_sub = Subset(train_set, tr_idx) val_sub = Subset(train_set, val_idx)

两种写法在功能上等价,但第二个版本在需要和其它人复现结果的时候更靠谱。我的习惯是默认用第二种,因为它把"哪 5000 张是验证集"这件事变成了一个显式的、可记录的数据。

6. 踩过的坑与排查链路

6.1 下载中断与文件完整性

最典型的场景是:下载到一半网络波动,raw目录里留下了一个不完整的文件。下次运行时,download=True会认为文件已存在而跳过下载,结果读取时直接抛格式异常。这类问题的排查链路是这样的。

第一步,先看raw目录的文件清单和大小。四个文件应该齐全,图像包大约 9.5MB 上下,标签包几十 KB。第二步,如果发现某个文件明显偏小,直接删掉它,重新跑一次,让它重新下载。第三步,如果反复下载都失败,找一份别人打包好的raw目录,整体替换。

这里有个容易忽略的点:torchvision在下载后会做校验,校验失败会抛异常,但校验失败的原因是"拿到的东西不对",可能是网络返回错误页,也可能是文件被截断。错误信息本身通常不会明说,所以不要只盯着报错文字,要去看文件本身的大小。

还有一种更隐蔽的情况:路径里存在同名但内容不同的旧数据。比如你之前在另一个项目里下过一次 MNIST,路径是./data,现在换了项目但路径没变,脚本读到的可能是旧的那份,而你以为读的是新的。所以我一直建议路径里带上项目标识,或者干脆用同一个固定的数据根目录,避免到处散落副本。

6.2 Windows 下多进程与重复运行脚本

num_workers > 0在 Windows 上的行为差异是新手最容易卡住的地方。表现是脚本莫名其妙地反复启动、内存占用飙升、最后报一个和进程相关的错误,或者干脆卡死。

根因在于 Windows 用 spawn 方式创建子进程,子进程会重新导入你的主模块。如果你的数据加载代码写在模块顶层、没有被if __name__ == "__main__":保护,子进程在导入时又会创建新的DataLoader,又会创建新的子进程,无限递归。

修法就是加保护块,把所有实际执行的逻辑放进去。只留函数和类的定义在顶层。这个习惯养成之后,在 Linux 和 Windows 上都能跑,不用维护两个版本。

顺带说一个相关的问题:有些人在transform里用了不可被 pickle 的对象,比如 lambda 函数。多进程加载时需要把Datasettransform序列化传给子进程,lambda 在很多情况下无法序列化,于是报错。解决办法是把 lambda 换成顶层函数或者用Compose组合已有的类。

6.3 常见报错速查

下面这张表是我自己攒的,覆盖了读取 MNIST 时最常遇到的几种报错。

报错关键词大概率原因处理方式
HTTP 错误 / 状态码异常下载地址不可达或返回错误页torchvision版本,或手动放置数据文件
文件解压或格式相关异常压缩包不完整或内容被替换删除raw下对应文件后重下
进程启动相关错误Windows 下缺少主模块保护if __name__ == "__main__":
序列化失败transform里含 lambda 或不可序列化对象改用顶层函数或类
张量维度不匹配漏了ToTensor或漏了squeeze检查单样本形状是否为(1, 28, 28)
损失一直不下降只做了ToTensor没做归一化补上Normalize
评估结果每次都不同测试集套了随机增强拆分训练与评估两条管道

这张表里最值钱的是最后两行。前面几行是环境问题,搜一下就有答案;后面两行是逻辑问题,报错不会告诉你,只能靠对流程的理解去发现。我见过太多人卡在"模型不收敛"上,查了半天模型结构,最后发现是数据管道里少了归一化。

排查这类问题的通用思路是:把数据管道和模型解耦,先单独验证数据。具体做法是在没有模型的情况下把一批数据打印出来,看形状、看值域、看可视化结果。数据对了再去怀疑模型。这个顺序如果反了,你会浪费大量时间在正确的地方找错误。

7. 从 MNIST 迁移到自己的数据集

7.1 自定义 Dataset 的三件套

等你把 MNIST 读通了,下一步通常是用自己的数据。自定义数据集需要实现三样东西:__init__负责记录文件列表和变换,__len__返回样本总数,__getitem__按下标返回一个样本。

from torch.utils.data import Dataset from PIL import Image class MyDataset(Dataset): def __init__(self, samples, transform=None): self.samples = samples # [(path, label), ...] self.transform = transform def __len__(self): return len(self.samples) def __getitem__(self, idx): path, label = self.samples[idx] img = Image.open(path).convert("RGB") if self.transform: img = self.transform(img) return img, label

这三个方法的职责边界很清楚。__init__里不要做重活,尤其不要在__init__里读图像,否则数据集一构建就要等很久,而且内存会被撑满。真正的读取放在__getitem__里,交给DataLoader的多进程并行处理。

7.2 迁移时要动哪几个地方

从 MNIST 迁到自己的数据集,需要改的东西其实只有四处。第一处是数据集类的替换,从datasets.MNIST换成你自己的类,或者换成ImageFolder这种按目录结构自动识别标签的通用类。第二处是归一化参数,MNIST 的 0.1307 和 0.3081 只对它自己成立,新数据集必须重新统计。

第三处是通道数相关的处理。MNIST 是单通道,你的数据大概率是三通道,模型第一层的输入通道数要跟着改,ToTensor之后张量形状从(1, H, W)变成(3, H, W)。第四处是图像尺寸,MNIST 是 28×28,你的数据可能大得多,通常需要加一步Resize,或者在模型里加自适应池化。

统计新数据集归一化参数的思路和 MNIST 一样,跑一遍数据集累加均值和方差。数据量大时可以用一个子集做抽样估计,精度损失很小但速度快很多。

# 抽样估计均值方差(三通道示例) import torch loader = DataLoader(my_dataset, batch_size=256, shuffle=True, num_workers=2) n, mean, m2 = 0, torch.zeros(3), torch.zeros(3) for imgs, _ in loader: b = imgs.size(0) flat = imgs.view(b, 3, -1) mean += flat.mean(dim=2).sum(dim=0) m2 += flat.var(dim=2, unbiased=False).sum(dim=0) n += b if n >= 5000: break print(mean / n) # 粗略均值

写到这里,我把整套读取流程从里到外过了一遍。自己这些年最深的体会是:数据管道这部分知识,看十篇教程不如自己动手把一批数据打印出来看一眼。形状、值域、可视化这三样东西,只要养成习惯每次都检查,绝大多数隐藏问题会在写模型之前就暴露出来。另外还有一个实用的小习惯,就是把数据体检写成一个独立的小函数,新项目接新数据时先跑一遍,比重头排查要省事得多。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询