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

如何将MNIST图像加载到PyTorch DataLoader?及通用加载方法咨询

针对JPG格式MNIST的PyTorch简易图像加载方案

嘿,我来帮你搞定这两个问题!既然你不想用现成的数据集类,也不需要CSV这类复杂的东西,那我们直接从最基础的自定义Dataset入手,一步步实现:

1. 把JPG格式的MNIST加载到PyTorch DataLoader

首先假设你的MNIST文件夹结构是这样的(如果你的结构略有不同,后面的代码也很容易调整):

mnist_jpg/
├── train/
│   ├── 0/          # 所有数字0的jpg图片
│   ├── 1/          # 所有数字1的jpg图片
│   └── ...         # 剩下的数字2-9文件夹
└── test/
    ├── 0/
    ├── 1/
    └── ...

我们先写一个极简的自定义Dataset类,直接从文件夹读取图片和标签:

步骤1:导入必要的库

import os
from PIL import Image
import torch
from torch.utils.data import Dataset, DataLoader
from torchvision import transforms

步骤2:自定义MNIST专用Dataset

class MNISTJPGDataset(Dataset):
    def __init__(self, root_dir, transform=None):
        self.root_dir = root_dir
        self.transform = transform
        self.image_paths = []
        self.labels = []
        
        # 遍历所有数字类别文件夹(0-9)
        digit_folders = sorted([f for f in os.listdir(root_dir) if os.path.isdir(os.path.join(root_dir, f))], key=int)
        
        for digit_label, digit_folder in enumerate(digit_folders):
            folder_path = os.path.join(root_dir, digit_folder)
            # 收集该文件夹下所有jpg图片
            for img_name in os.listdir(folder_path):
                if img_name.endswith(".jpg"):
                    self.image_paths.append(os.path.join(folder_path, img_name))
                    self.labels.append(digit_label)
    
    def __len__(self):
        return len(self.image_paths)
    
    def __getitem__(self, idx):
        # 加载灰度图(MNIST是单通道)
        img = Image.open(self.image_paths[idx]).convert("L")
        label = self.labels[idx]
        
        # 应用预处理变换
        if self.transform:
            img = self.transform(img)
        
        return img, label

步骤3:创建DataLoader

# 定义预处理变换(MNIST标准尺寸是28x28,加上标准化效果更好)
mnist_transform = transforms.Compose([
    transforms.Resize((28, 28)),  # 如果你的jpg已经是28x28可以省略
    transforms.ToTensor(),
    transforms.Normalize((0.1307,), (0.3081,))  # MNIST官方的均值和标准差
])

# 初始化数据集
train_dataset = MNISTJPGDataset(root_dir="mnist_jpg/train", transform=mnist_transform)
test_dataset = MNISTJPGDataset(root_dir="mnist_jpg/test", transform=mnist_transform)

# 创建DataLoader
train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True)
test_loader = DataLoader(test_dataset, batch_size=64, shuffle=False)

这样你就可以像使用官方MNIST数据集一样,通过train_loader迭代获取批量的图片和标签了!

2. 通用的简易图像加载类

上面的类是针对MNIST的,我们可以把它改成更通用的版本,适配任何按类别分文件夹存储的图像数据集(比如猫狗分类、花卉分类等),不需要依赖CSV或复杂配置:

class GenericImageDataset(Dataset):
    def __init__(self, root_dir, transform=None, img_extensions=(".jpg", ".jpeg", ".png"), label_mapping=None):
        self.root_dir = root_dir
        self.transform = transform
        self.img_extensions = img_extensions
        # 可选:手动指定类别到标签的映射,比如{"cat":0, "dog":1}
        self.label_mapping = label_mapping
        
        self.image_paths = []
        self.labels = []
        
        # 获取所有类别文件夹
        class_folders = [f for f in os.listdir(root_dir) if os.path.isdir(os.path.join(root_dir, f))]
        # 如果有标签映射,按映射排序;否则默认排序
        if self.label_mapping:
            class_folders.sort(key=lambda x: self.label_mapping[x])
        else:
            class_folders.sort()
        
        for class_name in class_folders:
            class_path = os.path.join(root_dir, class_name)
            # 收集所有符合格式的图片
            for img_name in os.listdir(class_path):
                if img_name.lower().endswith(self.img_extensions):
                    self.image_paths.append(os.path.join(class_path, img_name))
                    # 确定标签
                    if self.label_mapping:
                        label = self.label_mapping[class_name]
                    else:
                        label = class_folders.index(class_name)
                    self.labels.append(label)
    
    def __len__(self):
        return len(self.image_paths)
    
    def __getitem__(self, idx):
        # 加载图片(自动适配彩色/灰度)
        img = Image.open(self.image_paths[idx])
        label = self.labels[idx]
        
        if self.transform:
            img = self.transform(img)
        
        return img, label

通用类的用法示例

比如你有一个猫狗分类数据集,文件夹结构是cat_dog/train/cat和cat_dog/train/dog,可以这样用:

# 定义预处理变换
transform = transforms.Compose([
    transforms.Resize((224, 224)),
    transforms.ToTensor(),
])

# 手动指定标签映射(可选)
label_map = {"cat": 0, "dog": 1}
dataset = GenericImageDataset(root_dir="cat_dog/train", transform=transform, label_mapping=label_map)
dataloader = DataLoader(dataset, batch_size=32, shuffle=True)

这个通用类完全满足你的需求:不需要CSV,只依赖文件夹结构,适配绝大多数图像分类场景,而且代码简洁易懂,方便你根据自己的数据集调整细节。

内容的提问来源于stack exchange,提问作者Terry

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.26 09:00:36