PyTorch Dataset.__getitem__被传入超出__len__的索引问题排查
PyTorch Dataset迭代时出现超出范围索引的问题
问题代码
from torch.utils.data import Dataset import torch class TempDataset(Dataset): def __init__(self, window_size=200): self.window = window_size self.x = torch.randn(4340, 10, dtype=torch.float32) self.y = torch.randn(4340, 3, dtype=torch.float32) self.len = len(self.x) - self.window + 1 # 计算得4141 def __len__(self): return self.len def __getitem__(self, index): # 预期该条件永远不会触发 if index == self.len: print('self.__len__(): ', self.__len__()) print('Tried to access element @ index: ', index) return self.x[index: index + self.window], self.y[index + self.window - 1] ds = TempDataset(window_size=200) print('len: ', len(ds)) counter = 0 for x, y in ds: counter += 1 print('counter: ', counter)
运行输出
len: 4141 self.__len__(): 4141 Tried to access element @ index: 4141 counter: 4141
疑问
按预期__getitem__()应仅接收0到__len__()-1范围内的索引,为何会调用index=4141?且该索引被传入后,循环计数仍为4141,未计入这次调用,原因是什么?用DataLoader包装Dataset后现象依旧。
原因分析与解决
1. 超出范围索引的触发原因
Python遍历Dataset对象时,底层迭代器遵循**"尝试调用直到抛出IndexError"**的逻辑:从index=0开始,每次调用__getitem__后index自增,直到捕获到IndexError才终止迭代。
你的__len__返回4141,合法索引范围是0~4140,但迭代器会尝试index=4141:
- 访问
self.x[4141:4341]时,因为self.x长度为4340,切片会返回self.x[4141:4340](一个长度为199的张量),不会报错; - 当执行
self.y[4141+200-1] = self.y[4340]时,self.y的最大索引是4339,此时才会抛出IndexError; - 在抛出错误前,你的
if index == self.len条件已经触发,所以打印了index=4141的信息,随后迭代器捕获错误并停止循环。
2. 循环计数为4141的原因
循环里的counter +=1仅在__getitem__成功返回数据时执行。index=4141的调用最终抛出了IndexError,没有生成有效的(x,y)对,因此这次调用不会被计入循环次数,counter最终等于合法索引的数量(4141次,对应0~4140)。
3. 解决方法
要避免这种不必要的越界调用,建议在__getitem__开头添加索引合法性检查,主动抛出IndexError:
def __getitem__(self, index): if index < 0 or index >= self.len: raise IndexError(f"Index {index} out of bounds for dataset of size {self.len}") # 原有的切片与返回逻辑 return self.x[index: index + self.window], self.y[index + self.window - 1]
添加检查后,迭代器会在index达到self.len时立即捕获错误并停止,不会执行后续的切片和打印逻辑。
内容的提问来源于stack exchange,提问作者Mahesha999
相关产品推荐
相关产品推荐

