You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.04 22:50:34