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

使用Python实现PyTorch自定义Dataset时__getitem__报错求助

PyTorch自定义Dataset报错解决:list indices must be integers or slices, not list

错误原因分析

原代码的核心问题有三个:

  1. __getitem__重定义索引变量:方法内部将idx赋值为os.listdir()的返回结果(列表类型),最后执行return file_path[idx]时,用列表作为索引触发类型错误。
  2. 路径收集逻辑低效混乱:每次调用__getitem__都重新遍历目录生成路径列表,既浪费资源,又导致索引逻辑失效。
  3. 初始化阶段bug:使用未定义的root变量,路径硬编码拼接存在跨平台兼容问题。

修正后的实现代码

以下是重构后的自定义Dataset类,解决原问题并优化了整体逻辑:

import os
import torch
from PIL import Image
from torch.utils.data import Dataset

class My_custom(Dataset):
    def __init__(self, path: str, transform=None):
        self.root = path
        self.transform = transform
        self.file_paths = []  # 预存储所有图片的完整路径
        
        # 遍历根目录下的一级子目录(如train、test)
        for sub_dir in os.listdir(self.root):
            sub_dir_full = os.path.join(self.root, sub_dir)
            if not os.path.isdir(sub_dir_full):
                continue
            
            # 遍历一级子目录下的类别文件夹
            for class_dir in os.listdir(sub_dir_full):
                class_dir_full = os.path.join(sub_dir_full, class_dir)
                if not os.path.isdir(class_dir_full):
                    continue
                
                # 收集类别文件夹下的所有图片路径
                for img_name in os.listdir(class_dir_full):
                    if img_name.lower().endswith(('.png', '.jpg', '.jpeg')):
                        self.file_paths.append(os.path.join(class_dir_full, img_name))

    def __len__(self):
        # 直接返回预存路径的数量
        return len(self.file_paths)

    def __getitem__(self, idx):
        # 用整数索引直接获取预存的图片路径
        img_path = self.file_paths[idx]
        # 加载图片并转为RGB格式
        img = Image.open(img_path).convert('RGB')
        
        # 应用数据增强transform
        if self.transform:
            img = self.transform(img)
        
        # 从路径中提取标签(可根据需求调整逻辑)
        label = os.path.basename(os.path.dirname(img_path))
        # 若需将标签转为整数,可在__init__中建立类别映射字典
        # self.class_map = {cls:i for i, cls in enumerate(sorted(os.listdir(os.path.join(self.root, os.listdir(self.root)[0]))))}
        # label = self.class_map[label]
        
        return img, label

关键优化点

  • 预加载路径:在__init__阶段一次性遍历所有目录,将图片路径存入列表,避免重复遍历浪费资源。
  • 正确使用索引:__getitem__直接使用传入的整数索引访问预存路径,彻底解决类型错误。
  • 安全路径拼接:使用os.path.join拼接路径,适配Windows、Linux等不同系统的路径分隔符。
  • 增加容错判断:跳过非目录文件,只收集图片格式的文件,避免无效路径导致的报错。

使用示例

# 数据集根目录
root = "/your/dataset/root/path"

# 定义数据变换
from torchvision import transforms
transform = transforms.Compose([
    transforms.Resize((224, 224)),
    transforms.ToTensor()
])

# 创建Dataset实例
dataset = My_custom(path=root, transform=transform)
# 测试单个样本
sample_img, sample_label = dataset[2]
print(f"图片形状: {sample_img.shape}, 标签: {sample_label}")

# 创建DataLoader
from torch.utils.data import DataLoader
dataloader = DataLoader(dataset, batch_size=32, shuffle=True)

# 遍历DataLoader
for imgs, labels in dataloader:
    print(f"批次图片形状: {imgs.shape}, 批次标签: {labels}")
    break

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.23 10:50:00