PyTorch搭建CNN食品图像分类器:验证集(val)文件夹创建求助
食品图像分类器验证集创建与PyTorch数据加载方案
一、验证集(val)的正确创建逻辑
- 禁止直接使用测试集(test)充当验证集,必须从**训练集(train)**中拆分10%-20%的数据作为验证集,保留测试集用于最终模型泛化能力评估
- 目录结构需与train、test保持完全一致,确保每个类别都有对应的子目录:
content/drive/MyDrive/TDI 2023/Deteccion_auto_comidas/ ├── train/ │ ├── class_1/ │ ├── class_2/ │ ... │ └── class_11/ ├── val/ │ ├── class_1/ │ ├── class_2/ │ ... │ └── class_11/ └── test/ ├── class_1/ ├── class_2/ ... └── class_11/ - 拆分时保证每个类别按比例分配,避免类别不平衡问题,比如每个类别从train中抽取10%数据到val
二、训练集拆分代码实现(文件移动版)
使用shutil和sklearn实现按类别比例拆分:
import os import shutil from sklearn.model_selection import train_test_split # 主目录路径 base_dir = "content/drive/MyDrive/TDI 2023/Deteccion_auto_comidas/" train_dir = os.path.join(base_dir, "train") val_dir = os.path.join(base_dir, "val") # 创建验证集根目录及子类别目录 os.makedirs(val_dir, exist_ok=True) for class_name in os.listdir(train_dir): class_train_path = os.path.join(train_dir, class_name) if not os.path.isdir(class_train_path): continue class_val_path = os.path.join(val_dir, class_name) os.makedirs(class_val_path, exist_ok=True) # 获取当前类别下所有图像路径 img_paths = [os.path.join(class_train_path, f) for f in os.listdir(class_train_path) if f.lower().endswith(('.png', '.jpg', '.jpeg'))] # 9:1拆分训练/验证数据,random_state保证拆分结果可复现 train_imgs, val_imgs = train_test_split(img_paths, test_size=0.1, random_state=42) # 移动验证集图像到对应目录 for img in val_imgs: shutil.move(img, class_val_path)
三、PyTorch数据集加载代码
用ImageFolder和DataLoader加载训练、验证、测试集:
import torch import torchvision.transforms as transforms from torchvision.datasets import ImageFolder from torch.utils.data import DataLoader # 数据预处理流程 transform = transforms.Compose([ transforms.Resize((224, 224)), # 统一图像尺寸 transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) # 通用图像归一化参数 ]) # 加载数据集 train_dataset = ImageFolder(os.path.join(base_dir, "train"), transform=transform) val_dataset = ImageFolder(os.path.join(base_dir, "val"), transform=transform) test_dataset = ImageFolder(os.path.join(base_dir, "test"), transform=transform) # 创建数据加载器 train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True) val_loader = DataLoader(val_dataset, batch_size=32, shuffle=False) test_loader = DataLoader(test_dataset, batch_size=32, shuffle=False) # 验证类别匹配情况 print("训练集类别:", train_dataset.classes) print("验证集类别:", val_dataset.classes)
四、无文件移动的拆分方案(可选)
如果不想修改原train目录结构,用Subset实现索引拆分:
from torch.utils.data import Subset import numpy as np # 按类别拆分索引 indices = np.arange(len(train_dataset)) targets = np.array(train_dataset.targets) train_indices = [] val_indices = [] for class_idx in range(11): class_mask = targets == class_idx class_indices = indices[class_mask] split_point = int(0.9 * len(class_indices)) train_indices.extend(class_indices[:split_point]) val_indices.extend(class_indices[split_point:]) # 创建子集数据集 train_subset = Subset(train_dataset, train_indices) val_subset = Subset(train_dataset, val_indices) # 生成加载器 train_loader = DataLoader(train_subset, batch_size=32, shuffle=True) val_loader = DataLoader(val_subset, batch_size=32, shuffle=False)
内容的提问来源于stack exchange,提问作者SANDRA SANTOS GALVEZ
相关产品推荐
相关产品推荐

