新手学PyTorch第一个卡壳的地方,往往不是模型怎么写,而是数据集怎么弄到手。MNIST作为最经典的入门数据集,很多人第一次跑教程就挂在数据加载这一步:torchvision下载半天报404、网络超时、文件损坏,心态直接裂开。其实MNIST的获取无非两条路——让torchvision自动下载,或者手动下载到本地再读取。这篇文章把两种方式从头到尾掰开揉碎讲清楚,附带常用的可视化代码和踩坑记录,适合刚装上PyTorch还没跑通第一个训练脚本的初学者。看完你就能理解数据集在PyTorch里的来龙去脉,以后换CIFAR、ImageNet这类数据集也能举一反三。
1. 先搞清楚MNIST到底是什么
1.1 数据长什么样
MNIST全称是Modified National Institute of Standards and Technology,最早从美国国家标准与技术研究院收集的手写数字样本里整理而来,是计算机视觉领域流传最广的入门数据集之一。它包含0到9共10个类别的灰度手写数字图片,训练集有60000张,测试集有10000张。
每张图片的尺寸是28×28像素,单通道灰度图。灰度图的每个像素只有一个数值,范围0到255,0表示纯黑,255表示纯白,中间值是各种深浅的灰。单张28×28的图片展开成一维就是784个数字,用普通文本格式存储的话,一张图不算标签只算像素就是784字节,六万张训练图裸数据就要接近47MB,再加上一些头部信息和标签,整个数据集压缩包大概几十MB,这也是它能如此广泛传播的原因——体积小、下载快、单张图简单,稍弱的CPU都能跑起来。
MNIST的图片是扁平的一维数组,没有长宽的概念,官方文件里存的时候按行排列。28×28的矩阵按行优先展平,第n个像素对应原图中的第n//28行、第n%28列。这个映射关系在可视化时必须搞清楚,画图前把784维的向量reshape成28×28,否则画出来就是一堆乱码。
1.2 为什么新手都从它开始
MNIST在深度学习领域的地位类似编程里的"Hello World"。它足够小,几百个样本就能完成一次训练迭代,几秒钟就能跑完一个epoch,方便快速验证模型的正确性。而且它不需要额外的预处理,下载下来就能直接用,不像ImageNet那样还要做各种尺寸裁剪和归一化才能喂进网络。
MNIST的官方文件格式也是很多数据集工具的通用模板。PyTorch的torchvision.datasets模块在实现MNIST时,使用了一套标准的解析流程——先检测本地有没有缓存文件,没有就尝试下载,下载完再解压解析成内存中的张量。你把它搞明白了,CIFAR-10、FashionMNIST、SVHN等数据集的加载逻辑也都是同一套路,测试集、训练集、标签文件、图像文件的结构大同小异。
2. 准备工作:环境配置与镜像源设置
2.1 确认PyTorch和torchvision都装好了
网上下载MNIST依赖torchvision这个视觉工具库,它是PyTorch官方出的扩展包,datasets、transforms、models这些常用的视觉组件都在里面。很多人只装了torch没装torchvision,跑到import torchvision就报ModuleNotFoundError。
确认是否装好,直接在Python环境里执行:
import torch import torchvision print(torch.__version__) print(torchvision.__version__)能打印出两个版本号就说明环境没问题。如果报错,先装torchvision:
pip install torchvision这里有个容易忽略的点:torchvision和torch有严格的版本对应关系,版本差太多可能会触发底层警告甚至报错。比如torch 2.0配torchvision 0.15是配套的,torch 2.1要配torchvision 0.16。装的时候最好用pip install torchvision --upgrade让pip自动挑选兼容版本,或者直接用pip install torch torchvision一起装,省得自己配。
如果用的conda环境,也可以:
conda install pytorch torchvision cpuonly -c pytorch没有NVIDIA显卡或者没装CUDA的同学,装CPU版就够了,跑MNIST根本不需要显卡,CPU几十秒就能完成一个epoch的训练。
2.2 给torchvision配国内镜像源
torchvision的MNIST默认下载地址在GitHub的release文件托管上,国内直连这个地址经常超时或者直接404,这也是很多教程跑不通的核心原因。解决办法是配置国内镜像源。
推荐清华大学的PyPI镜像,在终端里临时指定或者写进pip配置都可以。临时指定的方式:
pip install torchvision -i https://pypi.tuna.tsinghua.edu.cn/simple永久配置的话,在用户目录下找到pip.ini(Windows)或pip.conf(Linux/macOS),写入:
[global] index-url = https://pypi.tuna.tsinghua.edu.cn/simple但注意:pip的镜像源解决的是包安装问题,而torchvision下载MNIST数据时用的是它自己内置的URL。镜像源对torchvision内部的下载不一定生效,这个问题下面第3章详细讲。实际项目中,最稳妥的方式是手动下载数据文件再让代码读取,这也是标题里"本地读取"这种方式的主要价值。
3. 方式一:让torchvision自动下载
3.1 标准写法与参数详解
网上下载方式的核心是torchvision.datasets.MNIST这个类。最简单的代码如下:
import torch from torchvision import datasets, transforms # 定义预处理:转张量 + 标准化 transform = transforms.Compose([ transforms.ToTensor(), # 将PIL图像或numpy数组转为Tensor,并把像素值从0~255缩放到0~1 transforms.Normalize((0.1307,), (0.3081,)) # 用MNIST数据集的全局均值和标准差做标准化 ]) # 下载并加载训练集 train_dataset = datasets.MNIST( root='./data', # 数据保存到当前目录下的data文件夹 train=True, # True加载训练集,False加载测试集 transform=transform, # 对图像做的预处理 download=True # 本地没有数据时就下载 ) # 下载并加载测试集 test_dataset = datasets.MNIST( root='./data', train=False, transform=transform, download=True ) print(f"训练集样本数: {len(train_dataset)}") print(f"测试集样本数: {len(test_dataset)}")几个参数逐个说清楚:
root是数据保存的根目录,代码会在root下自动创建MNIST/raw目录,把下载的原始压缩包放到raw里,再生成一个MNIST/processed目录存放处理好的数据。第二次运行代码时,download=True也会先检查processed目录是否存在,存在就直接加载,不会再重新下载。
train参数决定加载的是训练集还是测试集。训练集和测试集的文件是分开的,下载时各下一份。
transform是一个"变换流水线",PyTorch里每次取一个样本时都会对图像依次执行这些变换。常见的变换有transforms.ToTensor()、transforms.Normalize()、transforms.Resize()、transforms.RandomCrop()等。这里用ToTensor把PIL图像转成Tensor,Normalize把像素分布标准化,让均值接近0、方差接近1,神经网络训练更容易收敛。
download参数最坑,设为True时如果本地没有processed文件,它就去网上拉原始压缩包。如果你的网络访问GitHub不畅,这里就会卡住或报错。
3.2 download=True时到底发生了什么
datasets.MNIST在下载过程中做这些事:先检查root/MNIST/processed/training.pt和test.pt是否存在;如果存在,直接从.pt文件加载,跳过下载;如果不存在,检查root/MNIST/raw目录下有没有四个.gz文件(训练图像、训练标签、测试图像、测试标签);缺哪个就下载哪个,使用torchvision内置的URL,全部下载完成后解析成张量,缓存为.pt文件。
404错误的根源就在这里——默认的下载地址是:
https://github.com/myleott/mnist_py/blob/master/train-images-idx3-ubyte.gz?raw=true这类GitHub地址在某些网络环境下非常不稳定,要么直接404,要么下载速度只有几KB/s,下到一半超时,留下一个残缺的文件占着坑,下次再试还是失败。
针对这个问题,可以通过为MNIST对象手动指定下载地址来绕过。torchvision的download_url支持镜像参数,新版torchvision里datasets.MNIST的download操作会调用torchvision.datasets.utils.download_url,而download_url内部会用urlopen去请求地址。你可以先手动下载这四个gz文件,把它们放到root/MNIST/raw目录下,再让download=True去检测——如果raw目录里已有完整文件,它就不会再请求网络。
实际上很多经历过这种折磨的PyTorch玩家,后来都会直接转用本地读取的方式,绕开网络这个不稳定因素。
3.3 网上下载失败的排查思路
遇到"下载失败"、"404 Not Found"这类的报错时,按顺序检查这几样:
第一,看raw目录里有没有残留的0字节或几十KB的小文件。下载中断会生成不完整的临时文件,程序检测到文件存在就不再下载,于是永远卡在损坏文件上。解决方法是把root/MNIST/raw目录整个删掉,重新跑。
第二,检查torchvision版本。部分老版本0.8、0.9的MNIST下载逻辑里写死了旧的URL,已经失效,建议升级到0.13以上的版本,或者直接新装torchvision 0.15+。
第三,试一下手动访问默认URL,看浏览器能不能打开。打不开就别折腾了,直接切到第4章的本地读取方式。
第四,如果你的环境有代理但代理配置不正确,也会下载失败。这类问题在新手期特别常见,先用curl -I <url>先测一下网络连通性再定位。
4. 方式二:手动下载到本地再读取
4.1 手动获取MNIST原始文件
本地读取方式的思路很简单:自己去网上下载官方原始文件,放到PyTorch能识别的位置,然后在代码里解析它。手动下载的好处是把"下载"和"解析"两个步骤分开,网络断了也能离线完成,而且下载一次以后可以永久复用。
官网MNIST原始数据有四个文件,都用gzip压缩:
| 文件名 | 内容 | 大小(约) |
|---|---|---|
| train-images-idx3-ubyte.gz | 训练集图像(60000×28×28) | 47MB |
| train-labels-idx1-ubyte.gz | 训练集标签(60000个) | 59KB |
| t10k-images-idx3-ubyte.gz | 测试集图像(10000×28×28) | 7.9MB |
| t10k-labels-idx1-ubyte.gz | 测试集标签(10000个) | 9.8KB |
可以到MNIST官网或各大镜像站下载。下载完后,在项目目录下创建data/MNIST/raw文件夹,把这四个gz文件放进去,目录结构长这样:
你的项目目录/ └── data/ └── MNIST/ └── raw/ ├── train-images-idx3-ubyte.gz ├── train-labels-idx1-ubyte.gz ├── t10k-images-idx3-ubyte.gz └── t10k-labels-idx1-ubyte.gz这个目录结构不是随便定的,它就是torchvision默认的MNIST存放布局。放在这个位置,后续用download=True时程序检测到raw里已经有文件,就不会再联网。
4.2 用python直接解析idx格式
MNIST的文件格式简单粗暴,就是一个二进制头部加一堆原始数据。
图像文件的头部是32字节(4个int32大端整数):前4字节是魔数magic number,固定为2051,表示整数图像文件;接着4字节是样本数量;再4字节是行数(28);最后4字节是列数(28)。头部之后就是像素数据,按uint8顺序排列。
标签文件的头部是8字节(2个int32):前4字节魔数2049,表示标签文件;接着4字节是标签数量,之后就是uint8的标签值。
用Python加标准库gzip就能解析。写一个函数load_mnist_images和load_mnist_labels:
import gzip import numpy as np import torch import os def load_mnist_images(image_path): """读取IDX格式的MNIST图像文件,返回float32类型的Tensor,值范围0~1""" with gzip.open(image_path, 'rb') as f: raw = f.read() # 前4字节是魔数,第5~8字节是样本数,第9~12字节是行数,第13~16是列数 num_images = int.from_bytes(raw[4:8], 'big') rows = int.from_bytes(raw[8:12], 'big') cols = int.from_bytes(raw[12:16], 'big') # 从第16字节开始是像素数据,每个像素1字节,共num*rows*cols个 data = np.frombuffer(raw[16:], dtype=np.uint8) data = data.reshape(num_images, rows * cols) # 转成torch tensor并缩放到0~1 tensor = torch.tensor(data, dtype=torch.float32) / 255.0 return tensor # 形状:[样本数, 784] def load_mnist_labels(label_path): """读取IDX格式的MNIST标签文件,返回int64类型的Tensor""" with gzip.open(label_path, 'rb') as f: raw = f.read() num_labels = int.from_bytes(raw[4:8], 'big') data = np.frombuffer(raw[8:], dtype=np.uint8) return torch.tensor(data, dtype=torch.long) # 形状:[样本数] base_dir = './data/MNIST/raw' train_images = load_mnist_images(os.path.join(base_dir, 'train-images-idx3-ubyte.gz')) train_labels = load_mnist_labels(os.path.join(base_dir, 'train-labels-idx1-ubyte.gz')) test_images = load_mnist_images(os.path.join(base_dir, 't10k-images-idx3-ubyte.gz')) test_labels = load_mnist_labels(os.path.join(base_dir, 't10k-labels-idx1-ubyte.gz')) print(train_images.shape, train_labels.shape) print(test_images.shape, test_labels.shape)运行后输出类似torch.Size([60000, 784]) torch.Size([60000]),说明解析成功。注意图像Tensor的值范围已经从0~255变成了0~1,这是因为除以了255.0。这个缩放很重要,不缩放直接喂给网络,数值范围太大,Loss下降会很慢,还会导致梯度爆炸。
我一般还会把raw目录的路径写成一个全局变量,这样不同项目里复用起来方便。如果你的项目目录改了,只需要改base_dir就行。
4.3 封装成PyTorch的Dataset类
上面的tensor加载进内存后,可以直接用下标访问,但要真正融入PyTorch的训练流程,最好封装成torch.utils.data.Dataset的子类。Dataset类是PyTorch数据管道的第一环,它决定了"怎么拿到一个样本"。封装之后可以配合DataLoader做批量读取、打乱顺序、多线程加载,训练代码才干净。
from torch.utils.data import Dataset, DataLoader class MyMNIST(Dataset): """自定义MNIST数据集类,支持可选的transform""" def __init__(self, images, labels, transform=None): self.images = images # shape: [N, 784] self.labels = labels # shape: [N] self.transform = transform def __len__(self): return len(self.labels) def __getitem__(self, idx): image = self.images[idx] # 形状:[784] label = self.labels[idx] # 如果需要,把一维向量变成28x28的二维图像,方便做数据增强 if self.transform: image = self.transform(image) return image, label # 使用示例 train_dataset = MyMNIST(train_images, train_labels) test_dataset = MyMNIST(test_images, test_labels) train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True, num_workers=2) print(f"训练集一个batch大小: {next(iter(train_loader))[0].shape}") # 输出: torch.Size([64, 784])封装后,本地读取方式和torchvision自带的MNIST在使用体验上就完全一致了,后面可以无缝切换到同样的训练代码。如果后续想对图像做随机旋转、平移、裁剪这些数据增强,只需要在MyMNIST的transform参数里传入transforms.Compose([...])即可。
4.4 两种方式怎么选
网上下载方式适合网络畅通、追求简单的人。一行download=True搞定所有下载和解析,代码量最少,只需处理三次,几乎所有教程都用这种方式演示。
本地读取方式适合网络环境不稳定、需要离线使用、或者想理解底层原理的人。它把数据获取和模型训练彻底解耦,下载一次可以永久使用,而且无所谓网络环境。很多生产环境跑模型时,数据文件都是提前放到NAS或对象存储上,代码里走的就是类似本地方案。
实际开发中我见过很多人两种都备着:团队里写统一的下载脚本,先尝试自动下载,失败则提示手动下载放指定目录。这种"自动为主、手动兜底"的做法最稳。
5. 数据可视化:看看你手里的数据
5.1 单张图展示
读取完数据后第一件事就是可视化,确认数据没读错。用matplotlib画一张图,注意把784维的向量reshape成28×28再显示:
import matplotlib.pyplot as plt # 取训练集第0张图 image = train_images[0].reshape(28, 28) label = train_labels[0].item() plt.imshow(image, cmap='gray') plt.title(f'Label: {label}') plt.axis('off') plt.show()如果图像显示出来是一个歪歪扭扭的"5"或者"7",说明数据解析正确。如果图案像雪花一样乱,大概率是reshape时尺寸或者顺序搞错了。
cmap='gray'参数必须指定,否则matplotlib默认用彩色映射显示灰度图,出来是花花绿绿的,容易误导自己。
5.2 网格图批量展示
一次看一张太慢,可以拼一个网格图,比如3行3列展示前9张图:
fig, axes = plt.subplots(3, 3, figsize=(6, 6)) for i in range(9): ax = axes[i // 3, i % 3] img = train_images[i].reshape(28, 28) ax.imshow(img, cmap='gray') ax.set_title(f'Label: {train_labels[i].item()}') ax.axis('off') plt.tight_layout() plt.show()也可以随机抽9张看看不同数字的长相。这种批量展示在调试阶段特别有用——模型训练完之后,你也可以用同样的方式展示模型的预测结果,对比预测标签和真实标签,一眼看出模型错在哪。
5.3 类别分布统计图
MNIST十类样本数量是否均衡,直接影响模型训练效果。画一个标签分布柱状图,确认每个数字的样本数大致都在六千张左右:
from collections import Counter train_label_counts = Counter(train_labels.numpy()) test_label_counts = Counter(test_labels.numpy()) fig, axes = plt.subplots(1, 2, figsize=(12, 4)) # 训练集分布 axes[0].bar(train_label_counts.keys(), train_label_counts.values()) axes[0].set_title('Train Label Distribution') axes[0].set_xlabel('Digit') axes[0].set_ylabel('Count') # 测试集分布 axes[1].bar(test_label_counts.keys(), test_label_counts.values()) axes[1].set_title('Test Label Distribution') axes[1].set_xlabel('Digit') axes[1].set_ylabel('Count') plt.tight_layout() plt.show()MNIST各类别样本数量基本均衡,所以不需要做类别加权。但如果你后面换成别的数据集,类别分布图就要重点关注,不均衡时需要在采样器或Loss函数里做处理。
5.4 像素值分布图
MNIST这类灰度图,像素值大部分集中在0(纯黑背景)和255(纯白笔画)两端,中间值比较少。画一个直方图验证一下:
plt.hist(train_images.numpy().flatten(), bins=50, color='gray', alpha=0.7) plt.xlabel('Pixel value') plt.ylabel('Frequency') plt.title('MNIST Pixel Value Distribution') plt.show()如果你看到分布集中在0附近、少部分在高位,说明数据正常。但注意,因为显示之前已经除以了255,像素值范围是0~1,直方图横轴看起来会集中在0和1两端。如果你像我一样做了归一化处理,横轴会在-0.42到2.82附近(0.1307和0.3081变换后的结果),那是另一番图景。
归一化对可视化的影响容易被忽视:Normalize之后图像不再是人眼能直观理解的0~255灰度范围,此时如果直接imshow,很多细节被压缩到不明显的区间,图会显得非常暗甚至发灰。验证时如果想像原始数据一样直观展示,需要反归一化:
def unnormalize(img_tensor, mean=0.1307, std=0.3081): # img_tensor: 已经归一化的图像 img = img_tensor * std + mean # 先还原到0~1 img = img.clamp(0, 1) # 保险起见,裁剪到合法区间 return img.numpy()这个细节在调试可视化结果时非常重要,不然你精心画出来的训练样本图全是一片灰度,搞不清是数据问题还是展示问题。
6. 实操进阶:数据管道里的几个实用技巧
6.1 用DataLoader做批量加载和打乱
拿到了Dataset,还不能直接循环训练。深度学习训练是"小批量梯度下降"打法,每次喂一小撮样本,而不是一次过全部数据。DataLoader就是干这个活的:
from torch.utils.data import DataLoader train_loader = DataLoader( dataset=train_dataset, batch_size=64, shuffle=True, # 每个epoch开始前打乱顺序 num_workers=2, # 用2个子进程预加载数据 drop_last=False # 最后一批样本不足64时也保留 ) for batch_images, batch_labels in train_loader: print(f"batch_images: {batch_images.shape}, batch_labels: {batch_labels.shape}") # 跑一次前向传播...shuffle这个参数很关键。如果训练时每个epoch都用相同的顺序喂数据,模型会记住顺序而不是学到真正的特征,导致泛化能力变差。一般训练集shuffle设为True,测试集通常设为False,因为测试时不需要打乱。
num_workers是在后台预先读取下一个batch数据的子进程数。Windows上如果num_workers设太高容易报错或者内存占用飙升,建议设0或2,跑起来再慢慢调。Jupyter Notebook里有的环境设num_workers>0会卡住,遇到这种直接设0就行。
6.2 把原始文件做成备份
MNIST数据下载过一次之后,强烈建议把四个gz文件拷到一个独立的backup目录里,不要只留在项目data文件夹中。因为data文件夹随时可能被你rm -rf清理掉,backup目录才是真正的"离线数据源"。
实践里我习惯在项目里放一个fetch_mnist_data.py脚本,专门负责从镜像站下载并解压到本地,之后所有模型代码都从本地固定目录读取。这个脚本会做三件事:检查目录是否存在;用gzip校验文件完整性;把文件移动到raw目录。这样即使哪天有别的项目要用MNIST,直接去backup目录拿就行,不用再经历一次下载的折磨。
6.3 手工构造一份"迷你MNIST"做调试
在实际训练之前,我强烈建议手工抽取一小部分数据来验证整套代码能不能跑通。别一上来就60k样本全部喂进去——万一数据处理逻辑错了,跑了一小时才发现,心态全崩。
最快的验证方式是用切片:
mini_train = MyMNIST(train_images[:1000], train_labels[:1000]) mini_loader = DataLoader(mini_train, batch_size=32, shuffle=True) # 先在上面跑一个极简的线性模型 import torch.nn as nn model = nn.Linear(784, 10) criterion = nn.CrossEntropyLoss() optimizer = torch.optim.SGD(model.parameters(), lr=0.01) for images, labels in mini_loader: outputs = model(images) loss = criterion(outputs, labels) optimizer.zero_grad() loss.backward() optimizer.step() print(loss.item()) break # 只跑一个batch就停如果这个batch能正常前向、反向、更新参数,再换全量数据跑就放心多了。
6.4 关于模型的输入形状
MNIST图像有两种常见的输入形态:
一种是展平后的784维向量,适合全连接网络。前两章的代码里,图像本来就是[60000, 784]的形状,这种形态最适合入门。
另一种是保留二维结构的[1, 28, 28]或[3, 28, 28],适合卷积神经网络。如果想要CNN输入,需要把图像reshape成[batch, 1, 28, 28]:
# 使用之前加载的train_images,形状[60000, 784] train_images_cnn = train_images.reshape(-1, 1, 28, 28)这里的-1表示自动推断,1是通道数(灰度图单通道),28和28是高和宽。
主流的LeNet、CNN模型在处理MNIST时,都是按后一种形态接收输入的。你在网上看到别人代码里x = x.view(x.size(0), 1, 28, 28),做的就是这件事。
7. 常见问题速查与避坑指南
7.1 报错清单
| 报错/问题 | 原因 | 解决方式 |
|---|---|---|
HTTP Error 404 | torchvision内置下载地址失效或网络受限 | 手动下载四个gz文件放入raw目录 |
timed out | 网络不稳定,下载超时 | 手动下载本地读取;或配置镜像源 |
RuntimeError: File not found | raw目录缺少文件或路径不对 | 检查目录结构是否为data/MNIST/raw/xxx.gz |
OSError: [Errno 22] Invalid argument | Windows上路径分隔符或文件名大小写问题 | 用正斜杠拼接路径,检查文件名是否少了下划线 |
| 图像显示乱码 | reshape尺寸或顺序错误 | 确保按28×28、行优先顺序reshape |
| 像素值出现负值 | 做了Normalize之后直接可视化 | 归一化后需要反变换或裁剪到0~1再显示 |
| 第二次运行还在重新下载 | processed缓存文件缺失 | 检查processed目录下training.pt是否存在 |
DataLoader worker (pid(s) xx) exited unexpectedly | num_workers设置过高 | 把num_workers设为0再试 |
7.2 404报错的系统解法
404是新手最常遇到的坑。根因是torchvision的MNIST默认下载地址在特定网络环境下无法访问。处理步骤:
- 先删掉raw目录下所有残留文件,确保目录是空的。
- 浏览器访问官方MNIST镜像地址,一个个下载四个gz文件。
- 在项目下创建
data/MNIST/raw目录,严格按这个结构放文件。 - 重新运行代码,不要设置download=False,让它去检测本地文件。
- 如果代码还是尝试从网络下载,检查torchvision源码,确认是否有URL写死——新版torchvision的部分接口对已有文件会优先复用本地文件。
这个解法的核心思想是永远保留一份"应急本地数据",不要依赖网络下载这条路走通。
7.3 一个容易忽略的坑:标签的dtype
load_mnist_labels里返回的Tensor用的dtype是torch.long(即int64),因为PyTorch的交叉熵损失CrossEntropyLoss要求label必须是LongTensor。如果你的label是float32,训练时会直接报错:
Expected dtype torch.long but got dtype torch.float如果你不确定自己的label是什么类型,可以打印一下train_labels.dtype确认。
同理,图像Tensor一般用float32,因为神经网络权重都是float32,输入也必须是float32,uint8不可直接与权重相乘。很多人从numpy读出uint8数据后忘了转float32,直接丢进模型,报类型错误又折腾半天。
7.4 torchvision缓存目录与磁盘清理
torchvision除了在root目录下写数据,还会在系统缓存目录存一些下载临时文件。如果你很久没跑过MNIST相关的代码,又突然报磁盘空间不足,可以去这几个位置看看有没有残留的traffic:
- Windows:
C:\Users\<用户名>\.cache\torch - Linux/macOS:
~/.cache/torch
这些缓存文件删掉不影响后续使用,下次用的时候重新下载就行。新版本的torchvision还会在root下生成.torchvision之类的隐藏文件,同样可以定期清理。
8. 写在最后的一点经验
回想我最初跑通MNIST的过程,最崩溃的不是模型调参,反而是"数据到底有没有读对"这件事。第一次用torchvision下载,404报错折腾了一下午;后来改为手动下载到本地,终于把数据握在自己手里,心里踏实了,后面所有实验都顺了。
我个人的体会是:初学阶段,不要过度依赖一键下载的便利性。手动走一遍下载、解压、解析、可视化、封装的完整流程,你对数据集的内部结构会有一个非常扎实的认知。后面无论换什么数据集,拿到文件先看格式、再写解析、再可视化确认,这套流程几乎能帮你避开八成以上的数据坑。
最后再分享一个小技巧:拿到MNIST数据后,先把可视化跑出来,眼睛确认数据没问题,再写模型。不要跳步——很多看似玄学的训练失败,追到根子上都是数据解析时少了一个字节、多了一个偏移。数据是地基,地基歪了,上面再漂亮的模型也白搭。祝你把MNIST稳稳拿下,接下来不管往CNN还是Transformer方向走,这篇笔记里的习惯都能一直有用。