PyTorch中重创建MNIST数据集及DataLoader相关疑问
PyTorch数据集相关问题解答
一、内置MNIST数据集的标签与DataLoader逻辑问题
你使用的内置MNIST加载代码:
import torch import torchvision data = torch.utils.data.DataLoader( torchvision.datasets.MNIST("/Users/Myself/PyTorch_tutorials", transform=torchvision.transforms.ToTensor(), download=True), batch_size=128, shuffle=True)
核心逻辑说明
DataLoader本身不负责区分数据和标签,区分逻辑完全来自传入的Dataset类:
- torchvision的MNIST数据集类中,
__getitem__方法会返回一个(图像张量, 标签张量)的元组。 - DataLoader在遍历的时候,会把每个样本的元组收集起来,分别将图像和标签按batch维度堆叠,最终返回
(批量图像张量, 批量标签张量)的元组,所以你可以用for (x,y) in data:直接解包。
多维标签的处理方式
如果需要使用多维标签,只要让你的Dataset类的__getitem__方法返回(数据张量, 多维标签张量)即可。DataLoader会自动将多个样本的多维标签堆叠成批量的多维标签张量,遍历的时候直接用x, y = batch就能拿到批量的多维标签,y的形状就是[batch_size, *标签维度]。
二、自定义CSV版MNIST数据集的加载问题
你尝试的加载代码:
import pandas as pd import numpy as np data = pd.read_csv("mnist_train.csv") labels = data["5"].values datapoints = data.iloc[:,1:]
batch_size = 128 dataset_pytor = TensorDataset(torch.from_numpy(datapoints.values.reshape(-1,28,28)).unsqueeze(1)) my_loader = DataLoader(dataset_pytor, shuffle=True, batch_size=batch_size)
报错原因
当TensorDataset只传入一个张量时,它会把每个样本包装成一个长度为1的元组,所以遍历DataLoader时,每个batch是一个包含单个张量的列表(或元组),而不是单独的张量。这就是调用x.to(device)会报错的原因——x是列表类型,没有to方法。
而内置MNIST的DataLoader能返回(x,y)张量,是因为它的Dataset返回的是二元组,DataLoader会分别堆叠成两个独立的张量,最终返回二元组,所以可以直接解包。
正确的加载方式
方式1:修改遍历逻辑(快速解决)
直接取列表中的第一个元素:
for x in my_loader: x = x[0].to(device) # 提取列表内的张量
方式2:自定义Dataset类(规范方案,适合扩展)
自编码器只需要图像数据,所以自定义Dataset时只返回图像张量即可:
from torch.utils.data import Dataset, DataLoader import torch import pandas as pd class CustomMNISTDataset(Dataset): def __init__(self, csv_path): data = pd.read_csv(csv_path) # 预处理:转成float32类型,归一化到0-1区间,reshape为(通道数, 高, 宽) self.images = torch.tensor(data.iloc[:, 1:].values, dtype=torch.float32) / 255.0 self.images = self.images.reshape(-1, 1, 28, 28) def __len__(self): # 返回数据集总长度 return len(self.images) def __getitem__(self, idx): # 返回单个样本的图像张量 return self.images[idx] # 加载数据集并创建DataLoader dataset = CustomMNISTDataset("mnist_train.csv") my_loader = DataLoader(dataset, shuffle=True, batch_size=128) # 遍历时直接拿到张量 for x in my_loader: x = x.to(device) # 后续自编码器训练逻辑
内容的提问来源于stack exchange,提问作者user37292
相关产品推荐
相关产品推荐

