You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.04 02:12:43