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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.23 09:35:25