PyTorch中如何将训练、测试、验证数据集分别保存至独立文件夹
解决方案:拆分数据集为训练/验证/测试独立文件夹并避免冗余下载
核心实现思路
- 优先检查目标训练集、验证集文件夹是否存在,若已存在则直接加载,跳过全量训练集的下载与拆分流程
- 仅当目标文件夹不存在时,才加载全量训练集,按比例拆分后将样本分别保存到对应独立文件夹
- 测试集保留原有逻辑直接加载
完整代码实现
import os import torch import shutil from torchvision import datasets, transforms from torch.utils.data import random_split from PIL import Image # 替换为你的实际配置参数 path_to_data = "./datasets" dataset_name = config['Pytorch_Dataset']['dataset'] val_split_ratio = 0.2 # 验证集占全量训练集的比例 transform = transforms.ToTensor() # 定义各数据集的存储路径 train_dir = os.path.join(path_to_data, f"data_train_{dataset_name}") val_dir = os.path.join(path_to_data, f"data_validation_{dataset_name}") test_dir = os.path.join(path_to_data, f"data_test_{dataset_name}") # 检查训练/验证集文件夹是否已存在且非空 train_ready = os.path.exists(train_dir) and len(os.listdir(train_dir)) > 0 val_ready = os.path.exists(val_dir) and len(os.listdir(val_dir)) > 0 if train_ready and val_ready: # 直接加载已保存的训练/验证集 train_data = datasets.ImageFolder(root=train_dir, transform=transform) val_data = datasets.ImageFolder(root=val_dir, transform=transform) else: # 临时加载全量训练集(仅在目标文件夹缺失时执行) temp_train_root = os.path.join(path_to_data, f"temp_all_train_{dataset_name}") all_training_data = getattr(datasets, dataset_name)( root=temp_train_root, train=True, download=True, transform=transforms.ToPILImage() # 转PIL格式方便保存图片 ) # 拆分训练集与验证集 train_size = int((1 - val_split_ratio) * len(all_training_data)) val_size = len(all_training_data) - train_size train_subset, val_subset = random_split(all_training_data, [train_size, val_size]) # 保存训练集:按类别创建子文件夹,逐个保存样本 for idx, (img, label) in enumerate(train_subset): class_folder = os.path.join(train_dir, str(label)) os.makedirs(class_folder, exist_ok=True) img.save(os.path.join(class_folder, f"{idx}.png")) # 保存验证集:同训练集的目录结构逻辑 for idx, (img, label) in enumerate(val_subset): class_folder = os.path.join(val_dir, str(label)) os.makedirs(class_folder, exist_ok=True) img.save(os.path.join(class_folder, f"{idx}.png")) # 加载刚保存的训练/验证集 train_data = datasets.ImageFolder(root=train_dir, transform=transform) val_data = datasets.ImageFolder(root=val_dir, transform=transform) # 删除临时全量训练集文件夹,清理冗余文件 shutil.rmtree(temp_train_root) # 加载测试集(保留原有逻辑) test_data = getattr(datasets, dataset_name)( root=test_dir, train=False, download=True, transform=transform )
关键细节说明
- 目录校验逻辑:通过检查文件夹是否存在且非空,确保只有在必要时才执行拆分与保存操作,避免重复劳动
- 临时文件夹处理:用临时目录存放全量训练集,拆分完成后直接删除,不会留下冗余的
data_all_train文件夹 - 样本存储结构:按类别创建子文件夹保存样本,完全适配
ImageFolder的加载规范,后续可直接用标准方式调用 - 比例可调:修改
val_split_ratio参数即可自定义验证集占比,适配不同场景需求
内容的提问来源于stack exchange,提问作者Noumeno
相关产品推荐
相关产品推荐

