PyTorch DataLoader未随机采样:每次epoch返回同图问题求助
解决PyTorch DataLoader始终返回同一张图片的问题
嘿,我来帮你搞定这个头疼的问题!从你描述的症状和给出的代码来看,大概率是你的Dataset类少了关键的__len__方法,这会导致DataLoader无法正确遍历整个数据集,只能一直取索引0的元素。咱们一步步来排查和解决:
核心问题:缺少__len__方法
你提供的MyDataset类只实现了__init__和__getitem__,但PyTorch的DataLoader必须通过__len__方法知道数据集的总长度,才能生成正确的索引序列。如果没有这个方法,DataLoader会默认认为数据集长度为0或者1,自然每次都只能取索引0的图片——哪怕你调整了batch size,也只是重复取同一张图凑够batch而已。
修复步骤:给Dataset添加__len__
修改你的MyDataset类,加上这个简单的方法:
class MyDataset(Dataset): def __init__(self, path, loader=pil_loader): self.path = path self.images = os.listdir(path) def __getitem__(self, index): image = self.images[index] # 你的图片加载、预处理逻辑... def __len__(self): # 返回数据集的总样本数 return len(self.images)
额外检查:确认DataLoader的shuffle设置
虽然这不是你当前问题的核心,但为了保证每个epoch的样本顺序随机,记得初始化DataLoader时设置shuffle=True:
train_loader = DataLoader(train_ds, batch_size=1, shuffle=True)
验证方法
你可以在__getitem__里加一行打印,看看每个batch对应的索引是否在变化:
def __getitem__(self, index): print(f"当前加载的索引:{index}") image = self.images[index] # 后续逻辑...
如果运行后能看到索引从0到len(self.images)-1循环变化,就说明问题解决了。
内容的提问来源于stack exchange,提问作者Monica Heddneck
相关产品推荐
相关产品推荐

