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

PyTorch文件夹类自定义数据集构建及DataLoader类别返回错误排查

解决PyTorch自定义数据集返回错误类别的问题

问题回顾

你想要基于文件夹名称作为类别构建PyTorch分类数据集,但遇到了DataLoader返回类别错误的问题:当batch size设为1时,你期望每个batch返回单个类别索引,但现在每个图像对应的返回结果是全部15个类别的字典,完全不符合分类任务的需求。

问题根源分析

看了你的代码,主要有两个核心问题导致了这个错误:

  1. 返回了整个类别映射字典而非单个类别索引:你的__getitem__方法里直接返回了class_to_idx(整个类别到索引的映射字典),而没有根据当前图像的路径提取它所属的类别,再返回对应的单个索引值。这就是为什么每个样本都返回全部15个类别的原因。
  2. 初始化方法的参数混乱与冗余:
    • 你在CustomDataset的__init__里重新调用了find_classes,还硬编码了固定路径,直接覆盖了传入的classes和class_to_id参数,完全浪费了手动定义的类别映射。
    • 创建数据集时传入的参数顺序(train_image_paths, class_to_idx, classes)和__init__定义的参数顺序(image_paths, classes, class_to_id)完全颠倒,导致参数赋值混乱,进一步加剧了问题。

修复方案

针对这些问题,我给你整理了修正后的代码,核心改动如下:

  • 移除__init__里冗余的find_classes调用,自动从数据根目录获取类别(推荐,避免手动维护出错)。
  • 在__getitem__中,从图像路径提取所属的文件夹名称(也就是类别名),再通过类别映射得到对应的单个索引,返回这个索引而非整个字典。
  • 统一参数传递的顺序和命名,避免混乱。

完整修正代码

import os
import glob
import numpy as np
from PIL import Image
import torch
from torch.utils.data import Dataset, DataLoader
from torchvision import transforms

def find_classes(dir):
    # Finds the class folders in a dataset, dir (string): Root directory path.
    classes = [d.name for d in os.scandir(dir) if d.is_dir()]
    classes.sort()
    class_to_idx = {classes[i]: i for i in range(len(classes))}
    return classes, class_to_idx

def main():
    class CustomDataset(Dataset):
        def __init__(self, image_paths, data_root):
            self.image_paths = image_paths
            self.transforms = transforms.ToTensor()
            # 从数据根目录自动获取类别和映射,避免手动维护出错
            self.classes, self.class_to_idx = find_classes(data_root)

        def __getitem__(self, index):
            img_path = self.image_paths[index]
            # 从图像路径提取所属的类别文件夹名称
            class_name = os.path.basename(os.path.dirname(img_path))
            # 获取对应的类别索引
            label = self.class_to_idx[class_name]
            
            # 图像加载与预处理
            image = Image.open(img_path)
            t_image = image.convert('L')
            t_image = self.transforms(t_image)
            
            # 返回图像、单个类别索引、图像路径
            return t_image, label, img_path

        def __len__(self):
            return len(self.image_paths)

    # 数据路径配置
    data_root = "D:\\Neda\\Echo_View_Classification\\avi_images\\"
    folder_data = glob.glob(os.path.join(data_root, "*", "*.png"))
    len_data = len(folder_data)
    print("count of dataset: ", len_data)

    # 数据集划分
    split_1 = int(0.6 * len_data)
    split_2 = int(0.8 * len_data)
    folder_data.sort()
    train_image_paths = folder_data[:split_1]
    print("count of train images is: ", len(train_image_paths))
    np.savetxt('im_training_path_1.csv', np.c_[train_image_paths], fmt=['%s'], comments='', delimiter = ",")
    
    valid_image_paths = folder_data[split_1:split_2]
    print("count of validation image is: ", len(valid_image_paths))
    np.savetxt('im_valid_path_1.csv', np.c_[valid_image_paths], fmt=['%s'], comments='', delimiter = ",")
    
    test_image_paths = folder_data[split_2:]
    print("count of test images is: ", len(test_image_paths))
    np.savetxt('im_testing_path_1.csv', np.c_[test_image_paths], fmt=['%s'], comments='', delimiter = ",")

    # 创建数据集和DataLoader
    train_dataset = CustomDataset(train_image_paths, data_root)
    train_loader = DataLoader(train_dataset, batch_size=1, shuffle=False, num_workers=0)
    
    valid_dataset = CustomDataset(valid_image_paths, data_root)
    valid_loader = DataLoader(valid_dataset, batch_size=1, shuffle=False, num_workers=0)
    
    test_dataset = CustomDataset(test_image_paths, data_root)
    test_loader = DataLoader(test_dataset, batch_size=1, shuffle=False, num_workers=0)

    dataLoaders = {
        'train': train_loader,
        'valid': valid_loader,
        'test': test_loader,
    }

    # 测试一下返回结果是否正确
    for img, label, path in train_loader:
        print(f"图像路径: {path[0]}")
        print(f"对应的类别索引: {label[0]}")
        break

if __name__ == "__main__":
    main()

关键改动说明

  1. 自动获取类别映射:通过传入数据根目录,让CustomDataset自动调用find_classes获取类别和映射,避免手动维护类别字典时可能出现的拼写错误或不一致。
  2. 提取单个类别索引:在__getitem__中,利用os.path.dirname和os.path.basename从图像路径中提取所属的类别文件夹名称,再通过self.class_to_idx得到对应的整数索引,这才是分类任务需要的标签格式。
  3. 参数传递更清晰:创建数据集时只需要传入图像路径列表和数据根目录,参数更少更清晰,避免了之前的顺序混乱问题。

现在你运行代码后,每个batch(batch size=1)会返回单个图像、对应的单个类别索引和图像路径,完全符合你的预期需求。

内容的提问来源于stack exchange,提问作者AI_NA

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.13 08:43:00