PyTorch自定义数据集报错:list indices must be integers or slices, not list
PyTorch自定义Dataset调用异常:外部文件类接收列表idx,Notebook内类正常
问题场景
- 模型训练10、100轮epoch均正常,调整为500轮后GPU崩溃
- 重启GPU后Jupyter Notebook抛出500服务器错误,执行
pip install --upgrade nbconvert修复后,原代码无法运行 - 调试差异:
- 自定义数据集类放在外部Python文件(
src/custom_dataset.py)中调用时,__getitem__的参数idx为列表,触发TypeError - 同一类直接定义在Notebook单元格内调用时,
idx为整数,代码运行正常
- 自定义数据集类放在外部Python文件(
相关代码
外部文件src/custom_dataset.py中的数据集类
import os from natsort import natsorted from PIL import Image from datasets import Dataset class LoadPairedDataset(Dataset): def __init__(self, root_dir, transform=None): self.root_dir = root_dir self.transform = transform self.images = os.listdir(root_dir) def __len__(self): return len(self.images) def __getitem__(self, idx): print(idx) img_name = os.path.join(self.root_dir, self.images[idx]) image = Image.open(img_name) if self.transform: image = self.transform(image) return image
Jupyter Notebook调用代码
# Imports import os from PIL import Image import torch from torch.utils.data import Dataset from torchvision import transforms from src.custom_dataset import LoadPairedDataset # Notebook内定义的数据集类 class CustomImageDataset(Dataset): def __init__(self, root_dir, transform=None): self.root_dir = root_dir self.transform = transform self.images = os.listdir(root_dir) def __len__(self): return len(self.images) def __getitem__(self, idx): print(idx) img_name = os.path.join(self.root_dir, self.images[idx]) image = Image.open(img_name) if self.transform: image = self.transform(image) return image base_path = "../lol-custom" transform = transforms.Compose([ transforms.ToTensor() ]) # Notebook内类调用(正常运行) train_data = CustomImageDataset(root_dir=base_path + "/train/low", transform=transform) dataloader = torch.utils.data.DataLoader( train_data, batch_size=5, sampler=None, num_workers=0 ) print(next(iter(dataloader)).shape) # 输出0 1 2 3 4,以及torch.Size([5, 3, 400, 600]) print("######### 以下调用会报错 ##############") # 外部文件类调用(报错) train_data = LoadPairedDataset(root_dir=base_path + "/train/low", transform=transform) dataloader = torch.utils.data.DataLoader( train_data, batch_size=5, sampler=None, num_workers=0 ) print(next(iter(dataloader)).shape)
环境信息
- PyTorch 2.0.1
- Python 3.9.18
报错信息
0 1 2 3 4 torch.Size([5, 3, 400, 600]) ######### 以下调用会报错 ############## [0, 1, 2, 3, 4] --------------------------------------------------------------------------- TypeError Traceback (most recent call last) Cell In[1], line 61 52 dataloader = torch.utils.data.DataLoader( 53 train_data, 54 batch_size=5, 55 sampler=None, 56 num_workers=0 57 ) 58 # 输出[0, 1, 2, 3, 4]后触发错误 ---> 61 print(next(iter(dataloader)).shape) File ~\anaconda3\envs\mmie\lib\site-packages\torch\utils\data\dataloader.py:633, in _BaseDataLoaderIter.__next__(self) 630 if self._sampler_iter is None: 631 # TODO(https://github.com/pytorch/pytorch/issues/76750) 632 self._reset() # type: ignore[call-arg] --> 633 data = self._next_data() 634 self._num_yielded += 1 635 if self._dataset_kind == _DatasetKind.Iterable and \ 636 self._IterableDataset_len_called is not None and \ 637 self._num_yielded > self._IterableDataset_len_called: File ~\anaconda3\envs\mmie\lib\site-packages\torch\utils\data\dataloader.py:677, in _SingleProcessDataLoaderIter._next_data(self) 675 def _next_data(self): 676 index = self._next_index() # may raise StopIteration --> 677 data = self._dataset_fetcher.fetch(index) # may raise StopIteration 678 if self._pin_memory: 679 data = _utils.pin_memory.pin_memory(data, self._pin_memory_device) File ~\anaconda3\envs\mmie\lib\site-packages\torch\utils\data\_utils\fetch.py:49, in _MapDatasetFetcher.fetch(self, possibly_batched_index) 47 if self.auto_collation: 48 if hasattr(self.dataset, "__getitems__") and self.dataset.__getitems__: ---> 49 data = self.dataset.__getitems__(possibly_batched_index) 50 else: 51 data = [self.dataset[idx] for idx in possibly_batched_index] File ~\anaconda3\envs\mmie\lib\site-packages\datasets\arrow_dataset.py:2807, in Dataset.__getitems__(self, keys) 2805 def __getitems__(self, keys: List) -> List: 2806 """Can be used to get a batch using a list of integers indices.""" -> 2807 batch = self.__getitem__(keys) 2808 n_examples = len(batch[next(iter(batch))]) 2809 return [{col: array[i] for col, array in batch.items()} for i in range(n_examples)] File ~\Projects\mmie\src\custom_dataset.py:21, in LoadPairedDataset.__getitem__(self, idx) 19 def __getitem__(self, idx): 20 print(idx) ---> 21 img_name = os.path.join(self.root_dir, self.images[idx]) 22 image = Image.open(img_name) 24 if self.transform: TypeError: list indices must be integers or slices, not list
问题根源与解决方案
根源
外部文件中的LoadPairedDataset错误继承了Hugging Face datasets库的Dataset类,而非PyTorch原生的torch.utils.data.Dataset。
PyTorch的DataLoader会自动检测数据集是否实现了__getitems__方法:
- Hugging Face的
Dataset类内置了__getitems__,DataLoader会直接传入批量索引列表(如[0,1,2,3,4]) - 未实现
__getitems__的PyTorch原生Dataset子类,DataLoader会循环调用__getitem__并传入单个整数索引
解决方案
修改src/custom_dataset.py中的导入和继承,替换为PyTorch原生的Dataset:
import os from natsort import natsorted from PIL import Image # 替换导入为PyTorch的Dataset from torch.utils.data import Dataset class LoadPairedDataset(Dataset): def __init__(self, root_dir, transform=None): self.root_dir = root_dir self.transform = transform self.images = os.listdir(root_dir) def __len__(self): return len(self.images) def __getitem__(self, idx): print(idx) img_name = os.path.join(self.root_dir, self.images[idx]) image = Image.open(img_name) if self.transform: image = self.transform(image) return image
修改后重新运行代码,__getitem__将接收整数索引,与Notebook内的类行为一致,不再触发TypeError。
内容的提问来源于stack exchange,提问作者stic-lab
相关产品推荐
相关产品推荐

