使用PyTorch处理ShanghaiTech人群计数数据集时遭遇TypeError: 'DataLoader' object is not subscriptable错误求助
使用PyTorch处理ShanghaiTech人群计数数据集时遭遇TypeError: 'DataLoader' object is not subscriptable错误求助
嗨,看起来你是把PyTorch里的Dataset和DataLoader的核心概念搞混了,这才导致了这个棘手的错误。咱们一步步理清楚问题出在哪,再给你对应的修复方案。
核心问题:混淆Dataset与DataLoader的职责,且命名冲突
在PyTorch的数据流管道里,这两个组件的分工是完全明确的:
- Dataset:负责定义「如何读取、预处理原始数据」(比如从目录加载图片和对应的密度图),它是可下标、可迭代的,是整个数据管道的「数据源」。
- DataLoader:负责把Dataset包装起来,实现批量加载、数据打乱、多进程加速等功能,它必须接收Dataset对象作为输入,而不能直接接收目录路径这类原始数据。
你的代码里犯了两个关键错误:
- 你自定义的数据源类(本该是Dataset)被命名成了
DataLoader,和PyTorch内置的torch.utils.data.DataLoader完全重名,直接导致代码逻辑混乱。 - 你试图用PyTorch的
Subset去包装一个非Dataset对象(可能是你命名错误的DataLoader),而Subset只支持可下标的Dataset类型,这就触发了「不可下标」的错误。
修复方案:重构你的数据管道
咱们按照PyTorch的标准数据流流程来修正代码:
1. 先写一个规范的自定义Dataset类
首先把你的数据源类改个不冲突的名字,比如ShanghaiTechDataset,继承PyTorch的Dataset并实现必要的方法:
import os import torch import h5py from PIL import Image from torch.utils.data import Dataset class ShanghaiTechDataset(Dataset): def __init__(self, root_dir, shuffle=False): self.root_dir = root_dir # 遍历目录,获取图片和对应密度图的路径(根据ShanghaiTech的目录结构调整) image_dir = os.path.join(root_dir, "images") gt_dir = os.path.join(root_dir, "ground_truth") self.image_paths = [os.path.join(image_dir, f) for f in os.listdir(image_dir) if f.endswith(".jpg")] self.density_paths = [os.path.join(gt_dir, f.replace(".jpg", ".h5")) for f in os.listdir(image_dir) if f.endswith(".jpg")] # 可选:打乱数据顺序 if shuffle: combined = list(zip(self.image_paths, self.density_paths)) torch.random.shuffle(combined) self.image_paths, self.density_paths = zip(*combined) def __len__(self): # 返回数据集总样本数 return len(self.image_paths) def __getitem__(self, idx): # 实现单样本的读取和预处理 # 加载并预处理图片 img = Image.open(self.image_paths[idx]).convert("RGB") img = torch.tensor(img).permute(2, 0, 1).float() / 255.0 # 转成CHW格式并归一化 # 加载密度图并计算人数 with h5py.File(self.density_paths[idx], "r") as hf: density_map = torch.tensor(hf["density"][:]) person_count = torch.sum(density_map).item() return img, density_map, person_count
2. 用标准流程拆分并加载数据集
现在用这个规范的Dataset来重构你的主逻辑,彻底区分Dataset和DataLoader的职责:
import torch from torch.utils.data import DataLoader, Subset batch_size = 8 device = 'cuda:0' if torch.cuda.is_available() else 'cpu' # 1. 用自定义Dataset加载原始训练数据 train_root_dir = "data/part_A/train_data/" init_training_dataset = ShanghaiTechDataset(train_root_dir, shuffle=True) # 2. 拆分训练集和验证集:用Subset包装Dataset(而非DataLoader) train_size = int(0.9 * len(init_training_dataset)) val_size = len(init_training_dataset) - train_size # 生成随机拆分的索引(固定种子可保证拆分结果一致) indices = list(range(len(init_training_dataset))) torch.random.shuffle(indices) train_indices = indices[:train_size] val_indices = indices[train_size:] train_dataset = Subset(init_training_dataset, train_indices) val_dataset = Subset(init_training_dataset, val_indices) # 3. 用PyTorch的DataLoader包装Dataset,实现批量加载 train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=2) val_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False, num_workers=2) # 测试集同理 test_root_dir = "data/part_A/test_data/" test_dataset = ShanghaiTechDataset(test_root_dir, shuffle=False) test_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False, num_workers=2) # 现在测试迭代就不会报错了 dataiter = iter(train_loader) ex_images, ex_dmaps, ex_n_people = next(dataiter) # 显示图片和密度图 plot_corresponding_pairs(ex_images, ex_dmaps)
为什么之前会触发那个错误?
你之前把自定义Dataset命名成了DataLoader,和PyTorch的内置DataLoader重名,导致后续用Subset包装的时候,实际上是在包装一个DataLoader对象(而非Dataset)。PyTorch的DataLoader是不可下标的(不能用[]访问单个样本),而Subset要求输入必须是可下标的Dataset,所以在迭代最终的train_loader时,就抛出了'DataLoader' object is not subscriptable的错误。
额外小提醒
- 永远不要让自定义类和PyTorch的核心API重名,这会带来很多难以排查的隐性问题。
- 记住PyTorch数据流的标准链路:
原始数据 → Dataset(读取+预处理) → Subset(可选拆分) → DataLoader(批量加载)
备注:内容来源于stack exchange,提问作者Tan
相关产品推荐
相关产品推荐

