为何遍历PyTorch Dataset实例会无限循环?其长度明明为50000
问题分析与解决
问题原因
PyTorch的Dataset遍历逻辑是:从idx=0开始,不断调用__getitem__(idx)并递增idx,直到触发IndexError才会终止遍历。
你的MyDataLoader类里,__getitem__方法完全没用到传入的idx参数,也没有做任何索引越界的校验——不管idx多大,都能通过my_function()返回数据,永远不会抛出异常,所以遍历会无限进行下去。
虽然你实现了__len__返回50000,但这个方法主要是给PyTorch自带的DataLoader用的(比如计算批次数量),不是用来终止遍历的依据。另外注意:你的类名MyDataLoader容易和PyTorch的torch.utils.data.DataLoader混淆,建议改成MyDataset更合理。
修复方案
方案1:添加索引校验(基础修复)
在__getitem__里先判断idx是否超出self.data_size,超出就抛出IndexError,这样遍历到idx=50000时就会停止:
import torch import typing as t class MyDataset(torch.utils.data.Dataset): def __init__(self, data_size=50000): self.data_size = data_size def __len__(self) -> int: return self.data_size def __getitem__(self, idx) -> t.Tuple[torch.Tensor, torch.Tensor]: # 索引越界校验 if idx >= self.data_size: raise IndexError("index out of range") image, label = my_function() return image[None], label dl = MyDataset() print(len(dl)) # 输出50000 # 现在遍历会正常终止 for j, i in enumerate(dl): if j % 10000 == 0: print(j) # 输出:0、10000、20000、30000、40000 后停止
方案2:结合idx生成对应数据(更符合Dataset设计)
Dataset的核心是按索引返回对应数据,建议让my_function接收idx参数,生成对应的数据(而不是每次返回随机值),这样才符合PyTorch Dataset的设计初衷:
import torch import typing as t # 假设my_function可以接收idx参数,生成对应的数据 def my_function(idx): # 示例:生成和idx相关的tensor image = torch.randn(32,32) label = torch.tensor(idx % 10) return image, label class MyDataset(torch.utils.data.Dataset): def __init__(self, data_size=50000): self.data_size = data_size def __len__(self) -> int: return self.data_size def __getitem__(self, idx) -> t.Tuple[torch.Tensor, torch.Tensor]: if idx >= self.data_size: raise IndexError("index out of range") image, label = my_function(idx) return image[None], label
内容的提问来源于stack exchange,提问作者Saeed
相关产品推荐
相关产品推荐

