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

PyTorch中如何将训练、测试、验证数据集分别保存至独立文件夹

解决方案:拆分数据集为训练/验证/测试独立文件夹并避免冗余下载

核心实现思路

  1. 优先检查目标训练集、验证集文件夹是否存在,若已存在则直接加载,跳过全量训练集的下载与拆分流程
  2. 仅当目标文件夹不存在时,才加载全量训练集,按比例拆分后将样本分别保存到对应独立文件夹
  3. 测试集保留原有逻辑直接加载

完整代码实现

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.17 05:21:15