如何将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
相关产品推荐
相关产品推荐

