使用Python实现PyTorch自定义Dataset时__getitem__报错求助
PyTorch自定义Dataset报错解决:list indices must be integers or slices, not list
错误原因分析
原代码的核心问题有三个:
- __getitem__重定义索引变量:方法内部将
idx赋值为os.listdir()的返回结果(列表类型),最后执行return file_path[idx]时,用列表作为索引触发类型错误。 - 路径收集逻辑低效混乱:每次调用
__getitem__都重新遍历目录生成路径列表,既浪费资源,又导致索引逻辑失效。 - 初始化阶段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
相关产品推荐
相关产品推荐

