Python迭代器从指定索引启动的高效方法及PyTorch适配
你提到当前循环调用sample_real_video_batch函数n次来跳过前面的批次效率极低,确实这种做法会做大量无用的加载/迭代操作,完全没必要。下面是几个高效的解决方案,而且完全适用于PyTorch ImageDataset这类标准数据集:
方案1:通过自定义Sampler/BatchSampler直接定位起始批次
PyTorch的DataLoader本身就支持通过sampler或batch_sampler来精准控制数据采样的起始位置,这是最规范的做法。
假设你的video_sampler是一个BatchSampler(比如基于SequentialSampler构建的),你可以直接构造一个从指定批次开始的采样器:
from torch.utils.data import SequentialSampler, BatchSampler # 假设你要从第start_batch个批次开始(批次索引从0开始) batch_size = 32 # 替换成你的实际批次大小 start_batch = 10 # 替换成你想要的起始批次号 # 获取数据集的完整索引列表 full_indices = list(SequentialSampler(your_dataset)) # 计算起始索引:跳过前start_batch个批次的所有样本 start_sample_idx = start_batch * batch_size # 截取从起始索引到末尾的样本索引 truncated_indices = full_indices[start_sample_idx:] # 构建自定义的批次采样器 custom_batch_sampler = BatchSampler( SequentialSampler(truncated_indices), batch_size=batch_size, drop_last=False # 根据你的需求设置是否丢弃最后一个不完整批次 )
之后把这个custom_batch_sampler传给你的DataLoader,新的迭代器就会直接从指定批次开始遍历了。
方案2:修改现有函数,直接初始化迭代器到指定位置
如果你不想重构整个DataLoader逻辑,可以直接修改sample_real_video_batch函数,添加起始批次的参数,初始化时就跳过不需要的批次:
def sample_real_video_batch(self, start_batch=None): if self.video_enumerator is None: # 先把所有批次转换成列表,方便截取 all_batches = list(self.video_sampler) if start_batch is not None: # 从指定批次开始截取列表 all_batches = all_batches[start_batch:] # 基于截取后的批次列表创建迭代器 self.video_enumerator = enumerate(all_batches) try: batch_idx, batch = next(self.video_enumerator) except StopIteration: # 迭代完后重置,如果需要循环遍历的话 all_batches = list(self.video_sampler) if start_batch is not None: all_batches = all_batches[start_batch:] self.video_enumerator = enumerate(all_batches) batch_idx, batch = next(self.video_enumerator) b = batch if self.use_cuda: # 注意Python3中用items()替代iteritems() for k, v in batch.items(): b[k] = v.cuda() return b
调用时只需传入start_batch参数,比如sample_real_video_batch(start_batch=142),就能直接从第142个批次开始获取数据。
关于PyTorch ImageDataset的适用性
上面的两种方案完全适用于PyTorch的ImageDataset(比如ImageFolder),因为PyTorch的Dataset和Sampler体系是通用的,不管是图像数据集还是你自定义的视频数据集,只要遵循PyTorch的标准接口,就能用这些方法来控制迭代起始位置。
为什么循环调用n次低效?
循环调用n次next()来跳过前面的批次,本质上是加载并丢弃了前n个批次的数据,这会带来不必要的IO和计算开销(尤其是数据集很大时)。而上面的方案是直接从索引层面跳过不需要的部分,没有多余的数据加载操作,效率提升非常明显。
内容的提问来源于stack exchange,提问作者Bahman Rouhani

