PyTorch DataLoader使用自定义采样器时多轮迭代失效问题
问题分析与解决
问题根源
你的VaribleBatchSampler把迭代状态(batch_idx、start_idx、end_idx)存在了实例属性中,且__iter__直接返回自身。第一个epoch迭代完成后,这些状态变量停留在遍历结束的位置,下一次调用迭代器时,start_idx已经大于等于dataset_len,直接触发StopIteration,导致后续epoch无法生成新批次。
修复方案
正确的做法是让__iter__方法在每个epoch开始时重置状态,返回全新的迭代逻辑。以下提供两种实现方式:
方式一:生成器实现(简洁推荐)
用生成器替代实例状态维护,每次调用__iter__都会生成新的迭代器,自动重置状态:
class VariableBatchSampler(Sampler): def __init__(self, dataset_len: int, batch_sizes: list): self.dataset_len = dataset_len self.batch_sizes = batch_sizes def __iter__(self): start_idx = 0 for batch_size in self.batch_sizes: if start_idx >= self.dataset_len: break end_idx = min(start_idx + batch_size, self.dataset_len) yield torch.arange(start_idx, end_idx, dtype=torch.long) start_idx = end_idx # 处理batch_sizes总和小于数据集长度时的剩余数据 while start_idx < self.dataset_len: end_idx = self.dataset_len yield torch.arange(start_idx, end_idx, dtype=torch.long) start_idx = end_idx
方式二:重置实例状态的迭代器
如果坚持用__next__模式,需在__iter__中重置状态:
class VariableBatchSampler(Sampler): def __init__(self, dataset_len: int, batch_sizes: list): self.dataset_len = dataset_len self.batch_sizes = batch_sizes self._reset_state() def _reset_state(self): self.batch_idx = 0 self.start_idx = 0 self.end_idx = self.batch_sizes[self.batch_idx] if self.batch_sizes else self.dataset_len def __iter__(self): self._reset_state() return self def __next__(self): if self.start_idx >= self.dataset_len: raise StopIteration() batch_indices = torch.arange(self.start_idx, self.end_idx, dtype=torch.long) self.start_idx = self.end_idx self.batch_idx += 1 if self.batch_idx < len(self.batch_sizes): self.end_idx = min(self.start_idx + self.batch_sizes[self.batch_idx], self.dataset_len) else: self.end_idx = self.dataset_len return batch_indices
额外提示
- 类名拼写建议修正为
VariableBatchSampler(原类名少写了一个字母i) - 原代码中
self.start_idx += (self.end_idx - self.start_idx)可简化为self.start_idx = self.end_idx - 需确保
batch_sizes总和不足时,剩余数据能被正常采样
内容的提问来源于stack exchange,提问作者pyaj
相关产品推荐
相关产品推荐

