如何无需自定义脚本在PyTorch中为类别子文件夹数据集创建三类DataLoader
无需自定义脚本,用PyTorch构建脑肿瘤图像数据集的Train/Val/Test DataLoader
问题描述
我有一个包含脑肿瘤图像的数据集,想要构建CNN进行图像分类。此前接触的数据集是按“train”“test”文件夹划分的,但本次数据集的目录结构如下:
dataset_dir |_____tumor_type_1 |_____tumor_type_2 |_____tumor_type_3 |_____no_tumor现在我希望创建三个DataLoader,分别是train_dataloader、validation_dataloader和test_dataloader,请问如何在PyTorch中无需编写自定义脚本实现?
解决方案
PyTorch内置的torchvision.datasets.ImageFolder可以直接读取这种按类别划分目录的数据集,搭配拆分工具就能快速完成数据集拆分,全程无需自定义Dataset类。
1. 导入依赖库
import torch from torchvision import datasets, transforms from torch.utils.data import DataLoader, random_split
2. 定义图像预处理变换
根据任务需求设置图像增强和标准化操作,示例如下:
transform = transforms.Compose([ transforms.Resize((224, 224)), # 统一输入图像尺寸 transforms.ToTensor(), # 将PIL图像转换为PyTorch Tensor transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) # 用ImageNet均值方差标准化 ])
3. 加载完整数据集
ImageFolder会自动将子目录名称作为类别标签,标签顺序与子目录排序一致:
full_dataset = datasets.ImageFolder(root="dataset_dir", transform=transform)
4. 拆分数据集为Train/Val/Test
方法1:随机拆分(简单快速)
设定划分比例(例如70%训练、15%验证、15%测试),直接调用random_split:
# 计算各数据集大小 train_size = int(0.7 * len(full_dataset)) val_size = int(0.15 * len(full_dataset)) test_size = len(full_dataset) - train_size - val_size # 拆分数据集 train_dataset, val_dataset, test_dataset = random_split(full_dataset, [train_size, val_size, test_size])
方法2:分层拆分(保持类别分布)
如果需要确保拆分后每个数据集的类别比例与原数据集一致,可以结合sklearn的分层划分工具:
from sklearn.model_selection import train_test_split import numpy as np # 获取所有样本的标签 labels = np.array(full_dataset.targets) # 先拆分训练集和临时集(验证+测试) train_idx, temp_idx = train_test_split( np.arange(len(full_dataset)), test_size=0.3, stratify=labels, random_state=42 ) # 再从临时集拆分验证集和测试集 val_idx, test_idx = train_test_split( temp_idx, test_size=0.5, stratify=labels[temp_idx], random_state=42 ) # 用Subset构建各数据集 train_dataset = torch.utils.data.Subset(full_dataset, train_idx) val_dataset = torch.utils.data.Subset(full_dataset, val_idx) test_dataset = torch.utils.data.Subset(full_dataset, test_idx)
5. 创建DataLoader
为拆分后的数据集创建对应的DataLoader,设置批量大小、是否打乱等参数:
batch_size = 32 train_dataloader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True) validation_dataloader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False) test_dataloader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False)
内容的提问来源于stack exchange,提问作者SasikaA
相关产品推荐
相关产品推荐

