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

如何将MNIST的Numpy数组转换为PyTorch Dataset及DataLoader?

解决方案

首先明确PyTorch官方MNIST数据集的结构:

  • 训练集数据形状:(60000, 1, 28, 28),单通道灰度图(通道维度在前),数据类型为torch.float32,数值范围归一化到[0,1]
  • 训练集标签形状:(60000,),一维张量,数据类型为torch.long

你的数据需要做以下调整,再封装成Dataset加载到DataLoader:

1. 数据预处理与形状转换

先把numpy数组调整为符合CNN输入的格式:

import numpy as np
import torch
from torch.utils.data import Dataset, DataLoader

# 假设你的数据已加载为train_data (60000,784)和train_labels (60000,1)
# 1. 调整数据形状:从(60000,784)转为(60000,1,28,28)
train_data = train_data.reshape(-1, 1, 28, 28)
# 2. 归一化到[0,1](和官方MNIST一致,官方数据是0-255整数,转float后除以255)
train_data = train_data.astype(np.float32) / 255.0
# 3. 调整标签形状:从(60000,1)转为(60000,),并转为long类型
train_labels = train_labels.squeeze(axis=1).astype(np.int64)

2. 自定义Dataset类

继承PyTorch的Dataset类,实现必要方法:

class CustomMNISTDataset(Dataset):
    def __init__(self, data, labels):
        self.data = torch.from_numpy(data)
        self.labels = torch.from_numpy(labels)
    
    def __len__(self):
        return len(self.data)
    
    def __getitem__(self, idx):
        return self.data[idx], self.labels[idx]

3. 加载到DataLoader

创建Dataset实例后,传入DataLoader即可:

# 创建数据集实例
train_dataset = CustomMNISTDataset(train_data, train_labels)
# 创建DataLoader,参数可根据需求调整
train_loader = DataLoader(
    train_dataset,
    batch_size=64,  # 批次大小
    shuffle=True,   # 训练时打乱数据
    num_workers=2   # 多进程加载,根据CPU核心数调整
)

验证数据结构

你可以通过以下代码验证加载后的数据是否符合要求:

# 取一个批次的数据
batch_data, batch_labels = next(iter(train_loader))
print(f"批次数据形状: {batch_data.shape}")  # 应为(64,1,28,28)
print(f"批次标签形状: {batch_labels.shape}")  # 应为(64,)
print(f"数据类型: {batch_data.dtype}, 标签类型: {batch_labels.dtype}")  # 应为float32和long

这样处理后,你的数据结构就和PyTorch官方MNIST数据集的输出完全一致,后续可直接用于CNN模型训练。

内容的提问来源于stack exchange,提问作者Firestar-Reimu

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.02 03:23:10