PyTorch遍历DataLoader触发__getitem__()传参数量不符TypeError
问题背景
首次接触PyTorch的开发者编写自定义数据集类,计划通过DataLoader加载张量,构建DataLoader的代码如下:
train_loader = DataLoader(dataset_train, batch_size=6, drop_last=True)
执行遍历DataLoader的代码时:
for i,train_batch in enumerate(train_loader):
触发报错:TypeError: __getitem__() takes 1 positional argument but 2 were given。开发者怀疑报错和依赖库版本有关,所用依赖版本如下:
- matplotlib 3.5.2
- numpy 1.23.0
- opencv-python 4.6.0.66
- torch 1.12.0
- torch-tb-profiler 0.4.0
- torchaudio 0.12.0
- torchvision 0.13.0
问题根因
该报错和依赖版本没有关系,是自定义数据集类的__getitem__方法定义不符合PyTorch Dataset的接口规范。
PyTorch的DataLoader在取数据时,会自动给数据集的__getitem__方法传入两个位置参数:第一个是数据集实例自身(即self),第二个是当前要读取的样本索引值。如果自定义的__getitem__方法只定义了1个位置参数(通常是只写了self,漏写了索引参数),就会触发当前的参数数量不匹配报错。
修复方案
- 检查自定义数据集类的
__getitem__方法定义,必须显式传入索引参数,标准写法参考:
from torch.utils.data import Dataset class YourCustomDataset(Dataset): def __init__(self, # 你的初始化入参): # 自定义初始化逻辑,比如加载数据路径、预处理配置等 pass def __len__(self): # 返回数据集的总样本量 return 总样本数 # 必须保留index形参,接收DataLoader传入的索引值 def __getitem__(self, index): # 编写根据索引index取对应样本、标签的逻辑 # sample = 你的数据张量[index] # label = 你的标签张量[index] return sample, label
- 如果
__getitem__已经写了index参数仍然报同类错,检查是否错误给__getitem__加了@staticmethod静态方法装饰器,静态方法不会自动传入self参数,会导致参数计数错位,去掉错误装饰器即可。
内容的提问来源于stack exchange,提问作者eli
相关产品推荐
相关产品推荐

