PyTorch多长度DataLoader列表的next()行为及多加载器取数疑问
问题解答与实现方案
一、现有实现疑问解答
next()的迭代行为
next()会依次取出迭代器的下一个batch,并非只取首个。但如果不对条件loader的迭代器做循环处理,当所有batch被取完后,再次调用next()会抛出StopIteration异常,不会自动重复遍历样本。比如A类loader有10个batch,第1次next()取第1个,第2次取第2个,第11次调用就会报错。样本耗尽的情况
按A(50%)、B(35%)、C(15%)的占比,C类样本量最少,对应的loader总batch数也最少。如果没做循环处理,遍历全量100个batch的过程中,C类loader的迭代器会最先耗尽,此时调用next(loader_dict['C'])会直接抛出StopIteration报错,不会自动循环C类样本。只有通过额外处理(比如用itertools.cycle包装迭代器),才会让条件loader循环取样本。
二、后续需求实现方案
1. 每次epoch从全数据集取不同随机样本
初始化全量数据集的train_loader时,设置shuffle=True即可。PyTorch的DataLoader在每个epoch开始时,会自动重新打乱数据集顺序,因此每次遍历train_loader拿到的都是不同的随机样本。
2. 遍历完所有条件样本,且以最小规模条件样本遍历完为epoch终止条件
可以按以下步骤实现:
- 步骤1:每个epoch初始化条件loader迭代器
每个epoch开始时,为每个条件子集的loader重新创建迭代器,确保每次epoch都能从头遍历条件样本:loader_iters = {key: iter(loader) for key, loader in loader_dict.items()} - 步骤2:确定epoch的终止边界
计算每个条件loader的总batch数,取最小值作为当前epoch的总batch数(以样本量最少的C类为基准):min_batch_num = min(len(loader) for loader in loader_dict.values()) - 步骤3:按批次混合样本并训练
循环min_batch_num次,每次从每个条件loader取一个batch,再从train_loader取一个随机batch,混合后执行训练:
这样就能保证所有条件样本都被遍历,且当样本量最少的C类遍历完成时,自动终止当前epoch。# 提前初始化全量loader的迭代器,避免每次重新创建 full_iter = iter(train_loader) for _ in range(min_batch_num): # 获取各条件样本batch a_batch = next(loader_iters['A']) b_batch = next(loader_iters['B']) c_batch = next(loader_iters['C']) # 获取全量随机样本batch full_batch = next(full_iter) # 自定义样本混合逻辑(比如拼接张量) mixed_batch = { 'data': torch.cat([a_batch['data'], b_batch['data'], c_batch['data'], full_batch['data']], dim=0), 'label': torch.cat([a_batch['label'], b_batch['label'], c_batch['label'], full_batch['label']], dim=0) } # 执行训练操作 train_model(mixed_batch)
内容的提问来源于stack exchange,提问作者Alex
相关产品推荐
相关产品推荐

