如何迭代PyTorch DataLoader直至累计读取到指定数量的样本
PyTorch按总样本数控制训练的DataLoader实现方案
PyTorch原生默认的DataLoader没有内置total参数来直接实现自动循环加载数据集、累计读取指定样本数后停止的功能,但通过简单的封装即可实现你想要的调用效果,适配progressive growing of GANs这类按总样本数控制训练进度的场景。
方案1:自定义迭代器包装类(推荐)
将原生DataLoader封装为支持总样本数限制的迭代器,复用性高,完全符合你预期的调用写法:
from torch.utils.data import DataLoader class LimitedSampleDataLoader: def __init__(self, dataloader: DataLoader, total_samples: int): self.dataloader = dataloader self.total_samples = total_samples self.current_samples = 0 self.dataloader_iter = iter(dataloader) def __iter__(self): self.current_samples = 0 self.dataloader_iter = iter(self.dataloader) return self def __next__(self): if self.current_samples >= self.total_samples: raise StopIteration try: batch = next(self.dataloader_iter) except StopIteration: # 数据集遍历完一轮,重置迭代器继续加载 self.dataloader_iter = iter(self.dataloader) batch = next(self.dataloader_iter) # 计算当前batch的样本数,适配batch返回多输入的场景 batch_size = len(batch[0]) if isinstance(batch, (list, tuple)) else len(batch) self.current_samples += batch_size return batch
使用方式和你预期的逻辑完全一致:
# 先正常定义你的数据集和原生DataLoader dataset = YourCustomDataset() base_loader = DataLoader(dataset, batch_size=32, shuffle=True, num_workers=4) # 封装后指定总样本数即可 loader = LimitedSampleDataLoader(base_loader, total_samples=800000) # 直接遍历,累计读取到80万样本后会自动停止 for batch in loader: # 训练逻辑 pass
如果需要严格控制总样本数不超出指定值,可以在返回batch前增加截断逻辑,仅保留剩余需要的样本即可。
方案2:临时快速实现
如果仅单次使用,不需要复用逻辑,可以直接用itertools.cycle配合计数实现,代码更简短:
import itertools total_samples = 800000 count = 0 base_loader = DataLoader(...) for batch in itertools.cycle(base_loader): if count >= total_samples: break batch_size = len(batch[0]) if isinstance(batch, (list, tuple)) else len(batch) count += batch_size # 训练逻辑
内容的提问来源于stack exchange,提问作者aurelia
相关产品推荐
相关产品推荐

