自定义COCO数据集配BatchSampler时PyTorch DataLoader索引报错
解决方法
1. 检查DataLoader的参数传递
这是最容易踩的坑:别把BatchSampler传给DataLoader的sampler参数,要传给batch_sampler参数。
错误写法:
from torch.utils.data import RandomSampler, BatchSampler, DataLoader sampler = RandomSampler(my_coco_dataset) batch_sampler = BatchSampler(sampler, batch_size=8, drop_last=False) # 错在这:把batch_sampler塞给了sampler参数 dataloader = DataLoader(my_coco_dataset, sampler=batch_sampler)
正确写法:
dataloader = DataLoader(my_coco_dataset, batch_sampler=batch_sampler)
要是传错参数,DataLoader会把BatchSampler返回的整组索引(比如[0,1,2,3])当成单个索引丢给__getitem__,但你的数据集只认整数索引,自然就报类型错了。
2. 检查自定义数据集MyCoco的__getitem__实现
确保__getitem__只接受单个整数索引,别写处理列表索引的逻辑。比如:
正确的__getitem__示例:
class MyCoco(CocoDetection): def __getitem__(self, idx): # idx是单个整数,直接取对应数据 img, target = super().__getitem__(idx) # 你的自定义处理逻辑... return img, target
PyTorch的DataLoader会自动根据batch_sampler的索引列表,挨个调用__getitem__再打包成batch,不用你自己在__getitem__里处理批量索引。
3. 确认Sampler的输出格式
如果是你自己写的RandomSampler或BatchSampler,检查__iter__方法返回的索引结构:
- RandomSampler的迭代器要返回单个整数(比如
0,1,2,...) - BatchSampler的迭代器要返回每组batch的索引列表(比如
[0,1], [2,3], ...)
要是自定义Sampler返回了嵌套列表(比如把[0,1]当成单个元素传给__getitem__),就得修正迭代逻辑。
4. 简化调试
可以先不用自定义Sampler,用默认参数测试DataLoader:
dataloader = DataLoader(my_coco_dataset, batch_size=8, shuffle=True)
如果这个能正常跑,说明问题肯定出在自定义Sampler的组合或参数传递上,再回头排查前面的步骤。
内容的提问来源于stack exchange,提问作者Dr. Prof. Patrick
相关产品推荐
相关产品推荐

