在GCP Vertex AI的JupyterLab加载MNIST数据集遇RuntimeError
在GCP Vertex AI JupyterLab中加载MNIST数据集触发RuntimeError的解决方法
问题描述
运行以下PyTorch代码加载MNIST数据集时触发RuntimeError:
import torch from torchvision import transforms from torchvision import datasets train_data = datasets.MNIST(root='data', train=True, download=True, transform=None) print(train_data)
环境版本
- torch: 1.12.1+cu113
- torchvision: 0.13.1+cu113
报错信息
--------------------------------------------------------------------------- RuntimeError Traceback (most recent call last) /tmp/ipykernel_10081/229378695.py in <module> 11 from torchvision import datasets 12 ---> 13 train_data = datasets.MNIST(root='data', train=True, download=True, transform=None) 14 print(train_data) /opt/conda/lib/python3.7/site-packages/torchvision/datasets/mnist.py in __init__(self, root, train, transform, target_transform, download) 102 raise RuntimeError("Dataset not found. You can use download=True to download it") 103 ---> 104 self.data, self.targets = self._load_data() 105 106 def _check_legacy_exist(self): /opt/conda/lib/python3.7/site-packages/torchvision/datasets/mnist.py in _load_data(self) 121 def _load_data(self): 122 image_file = f"{'train' if self.train else 't10k'}-images-idx3-ubyte" ---> 123 data = read_image_file(os.path.join(self.raw_folder, image_file)) 124 125 label_file = f"{'train' if self.train else 't10k'}-labels-idx1-ubyte" /opt/conda/lib/python3.7/site-packages/torchvision/datasets/mnist.py in read_image_file(path) 542 543 def read_image_file(path: str) -> torch.Tensor: ---> 544 x = read_sn3_pascalvincent_tensor(path, strict=False) 545 if x.dtype != torch.uint8: 546 raise TypeError(f"x should be of dtype torch.uint8 instead of {x.dtype}") /opt/conda/lib/python3.7/site-packages/torchvision/datasets/mnist.py in read_sn3_pascalvincent_tensor(path, strict) 529 530 assert parsed.shape[0] == np.prod(s) or not strict ---> 531 return parsed.view(*s) 532 533 RuntimeError: shape '[60000, 28, 28]' is invalid for input of size 9437168
问题背景
该问题仅在GCP Vertex AI的JupyterLab环境中出现,本地及Colab环境无法复现,更换多个torch和torchvision版本均未解决。
问题分析
报错显示输入数据大小为9437168,而MNIST训练集正常应包含60000张28×28的图像,总大小应为60000×28×28=47040000。这说明自动下载的数据集文件损坏或未正确解压,大概率是GCP环境中下载过程出现网络中断、文件系统权限限制或解压异常导致。
解决步骤
1. 清理损坏的数据集文件
执行以下代码删除已下载的损坏文件:
import shutil import os if os.path.exists('data/MNIST'): shutil.rmtree('data/MNIST')
2. 手动下载并解压数据集
在JupyterLab终端中运行以下命令,手动下载并解压MNIST数据集到指定目录:
mkdir -p data/MNIST/raw cd data/MNIST/raw # 下载MNIST数据集文件 wget http://yann.lecun.com/exdb/mnist/train-images-idx3-ubyte.gz wget http://yann.lecun.com/exdb/mnist/train-labels-idx1-ubyte.gz wget http://yann.lecun.com/exdb/mnist/t10k-images-idx3-ubyte.gz wget http://yann.lecun.com/exdb/mnist/t10k-labels-idx1-ubyte.gz # 解压所有压缩文件 gunzip *.gz
3. 验证文件完整性
运行以下命令检查文件大小,确保与标准MNIST文件一致:
ls -l data/MNIST/raw
标准文件大小参考:
train-images-idx3-ubyte: 47040016 bytestrain-labels-idx1-ubyte: 60008 bytest10k-images-idx3-ubyte: 7840016 bytest10k-labels-idx1-ubyte: 10008 bytes
4. 重新加载数据集
再次执行最初的加载代码,此时应能正常加载MNIST数据集。
内容的提问来源于stack exchange,提问作者Adir Morgan
相关产品推荐
相关产品推荐

