PyTorch Lightning 任意可迭代对象与多 DataLoader 支持:CombinedLoader 模式详解
【免费下载链接】pytorch-lightningPretrain, finetune ANY AI model of ANY size on 1 or 10,000+ GPUs with zero code changes.项目地址: https://gitcode.com/gh_mirrors/py/pytorch-lightning
导读
本篇技术指南围绕 PyTorch Lightning Trainer 对「任意可迭代对象(arbitrary iterables)」与「多个可迭代对象集合」的内建支持展开,重点剖析多 DataLoader 场景下批次自动合并的核心机制——CombinedLoader及其四种采样模式(min_size、max_size_cycle、max_size、sequential)。读完本文,你将掌握在训练、验证、测试、预测各阶段传入单个或多个 DataLoader 的正确姿势,理解每种模式对批次长度与数据消耗顺序的影响,并能结合仓库源码(combined_loader.py)与测试用例(test_combined_loader.py)写出可正确运行、可应对复杂数据编排的实战代码。
Python 可迭代对象与 DataLoader 的关系
在 Python 中,可迭代对象(iterable)指任何可以被迭代或循环遍历的对象,典型的例子包括列表、字典等。在 PyTorch 中,torch.utils.data.DataLoader本身就是一个可迭代对象,它通常从一个torch.utils.data.Dataset或torch.utils.data.IterableDataset中取数。
PyTorch Lightning 的Trainer直接支持任意可迭代对象作为数据源,而不仅仅是DataLoader。这意味着:
- 你可以传入一个原生 Python 迭代器(如
list(range(1000))); - 也可以传入
DataLoader,这是绝大多数用户的选择; - 甚至可以把它们组合成字典或列表的集合(见下文),Lightning 会自动按模式合并批次。
Trainer的这一能力覆盖了数据流的四个入口:Trainer.fit、Trainer.validate、Trainer.test和Trainer.predict,它们在内部都会先通过_request_dataloader获取数据源,再交由对应循环的setup_data处理(见 fit_loop.py、evaluation_loop.py 与 prediction_loop.py)。
多个可迭代对象:字典、列表与嵌套组合
除了支持单个任意可迭代对象外,Trainer还支持「可迭代对象的集合」。典型写法如下:
# 单个 DataLoader return DataLoader(...) # 原生 Python 可迭代对象 return list(range(1000)) # 以字典传入多个 DataLoader,会生成这样的批次: # {'a': batch_from_loader_a, 'b': batch_from_loader_b} return {"a": DataLoader(...), "b": DataLoader(...)} # 以列表传入多个 DataLoader,会生成这样的批次: # [batch_from_dl_1, batch_from_dl_2] return [DataLoader(...), DataLoader(...)] # 嵌套组合:字典的值是列表,会生成这样的批次: # {'a': [batch_from_dl_1, batch_from_dl_2], 'b': [batch_from_dl_3, batch_from_dl_4]} return {"a": [dl1, dl2], "b": [dl3, dl4]}这些写法可以出现在LightningDataModule的train_dataloader/val_dataloader/test_dataloader/predict_dataloader钩子中,也可以直接作为Trainer.fit、Trainer.validate、Trainer.test、Trainer.predict的 dataloader 参数传入,支持范围完全一致。
从源码看,CombinedLoader在构造时会通过_tree_flatten将任意嵌套的集合(字典、列表、元组的组合)拍平成扁平列表,并保存原始结构描述_spec;产出批次时再通过tree_unflatten还原为原始嵌套结构(combined_loader.py)。因此无论你是传字典、列表还是二者的嵌套组合,最终training_step/validation_step收到的 batch 都与你定义的容器结构一致。
CombinedLoader:多可迭代对象的自动合并核心
Lightning 根据一个「模式(mode)」自动将来自多个可迭代对象的批次合并起来,这一工作由lightning.pytorch.utilities.combined_loader.CombinedLoader完成。它也是一个Iterable,可以直接交给Trainer。
四种采样模式
CombinedLoader的mode参数支持以下四种取值(见 combined_loader.py 中的_SUPPORTED_MODES注册表):
| 模式 | 行为 | 总批次数 |
|---|---|---|
min_size | 在最短的可迭代对象(批次数最少的那个)耗尽时停止 | min(lengths) |
max_size_cycle | 在最长可迭代对象耗尽时停止,期间对已耗尽的可迭代对象循环重置并继续取数 | max(lengths) |
max_size | 在最长可迭代对象耗尽时停止,已耗尽的迭代器返回None(不循环) | max(lengths) |
sequential | 逐个完整消费每个可迭代对象,返回三元组(data, idx, iterable_idx) | sum(lengths) |
内部实现上,每种模式对应一个迭代器类:_MinSize、_MaxSizeCycle、_MaxSize、_Sequential,它们都继承自_ModeIterator(combined_loader.py),统一返回(batch, batch_idx, dataloader_idx)三元组:
_MaxSizeCycle(L67-L106):某个迭代器抛出StopIteration时标记为已耗尽,若还有未耗尽的迭代器,则用iter(self.iterables[i])重新创建迭代器继续循环取数;_MaxSize(L184-L205):抑制StopIteration,耗尽后把该位置的输出置为None,直到所有迭代器都耗尽;_Sequential(L123-L181):只同时加载当前迭代器(_load_current_iterator每次只创建一个迭代器),避免多余 worker 进程启动,一个迭代器耗尽后切换到下一个,并返回真实的dataloader_idx。
默认模式与各阶段限制
- 训练阶段默认使用
max_size_cycle:最长 DataLoader 跑满,其余 DataLoader 循环复用,保证每个 epoch 的训练步数由最长数据决定。该默认值在 fit_loop.py 中设置。 - 验证、测试与预测阶段默认使用
sequential:多个 DataLoader 依次完整消费,互不交织。该默认值在 evaluation_loop.py 与 prediction_loop.py 中设置。 trainer.predict仅支持"sequential"模式:如果传入其他模式的CombinedLoader,会在prediction_loop.reset中抛出ValueError('trainer.predict() only supports the CombinedLoader(mode="sequential") mode.')(prediction_loop.py)。trainer.fit不支持"sequential"模式:fit_loop对传入的CombinedLoader模式有显式校验,其余模式的组合方式由_SUPPORTED_MODES决定。
手动选择模式
如果默认模式不满足需求,可以直接使用CombinedLoader并指定mode,再把它传给Trainer:
from lightning.pytorch.utilities import CombinedLoader iterables = {"a": DataLoader(), "b": DataLoader()} combined_loader = CombinedLoader(iterables, mode="min_size") model = ... trainer = Trainer() trainer.fit(model, combined_loader)CombinedLoader已在 utilities/init.py 中导出,因此可直接从lightning.pytorch.utilities导入。如果传入不支持的 mode 字符串,构造时会抛出ValueError并列出所有合法取值(combined_loader.py)。
各模式的行为示例
源码 docstring 中给出了一个非常直观的示例:iterables = {'a': DataLoader(range(6), batch_size=4), 'b': DataLoader(range(15), batch_size=5)},即a有 2 个 batch、b有 3 个 batch:
max_size_cycle:共 3 个 batch。a在第 2 个 batch 耗尽后循环重置,第 3 个 batch 输出{'a': tensor([0,1,2,3]), 'b': tensor([10,...,14])};max_size:共 3 个 batch。第 3 个 batch 输出{'a': None, 'b': tensor([10,...,14])},a不循环;min_size:共 2 个 batch,a耗尽即停止;sequential:共 5 个 batch,先完整输出a的 2 个 batch(dataloader_idx=0),再输出b的 3 个 batch(dataloader_idx=1)。
这些行为在仓库测试 test_combined_loader.py 中也有覆盖(例如test_combined_dataset验证_dataset_length()按模式返回min/max聚合结果)。注意:CombinedLoader的__len__需要先调用iter(combined_loader)才会返回批次数,否则抛出RuntimeError。
limits 与长度控制
CombinedLoader提供limits属性(combined_loader.py),可以按迭代器设置批次上限:
- 传入单个数值会广播到所有迭代器;
- 传入列表时长度必须与扁平化后的迭代器数量一致,否则抛出
ValueError。
Trainer在setup_data阶段会结合limit_train_batches/limit_val_batches/limit_test_batches/limit_predict_batches等参数为每个迭代器计算出实际num_batches并写入combined_loader.limits(fit_loop.py)。各模式的__len__都会在存在 limits 时对长度做min(length, limit)截断后再聚合(如_Sequential.__len__返回sum(min(length, limit)))。
sequential 模式下钩子的 dataloader_idx 参数
使用"sequential"模式时,批次来自不同的 DataLoader,因此需要在部分钩子中额外添加dataloader_idx参数,Lightning 会在缺失时抛出错误提示这一要求。
涉及dataloader_idx的钩子定义在 core/hooks.py,主要包括:
on_validation_batch_start(batch, batch_idx, dataloader_idx=0)(L94)与on_validation_batch_end;on_test_batch_start(batch, batch_idx, dataloader_idx=0)(L117)与on_test_batch_end;on_predict_batch_start(batch, batch_idx, dataloader_idx=0)(L138)与on_predict_batch_end;- 批次传输相关钩子:
transfer_batch_to_device(batch, device, dataloader_idx)(L565)、on_before_batch_transfer(batch, dataloader_idx)(L614)、on_after_batch_transfer(batch, dataloader_idx)(L642)。
在sequential模式下,dataloader_idx表示当前批次来自第几个 DataLoader,可用于按来源区分处理逻辑(例如transfer_batch_to_device中对不同来源的 batch 做不同的设备搬移处理)。在非 sequential 模式下(如训练阶段的max_size_cycle),dataloader_idx恒为 0,因为每步都会取所有迭代器的批次并合并为一个 batch 结构。
在 LightningDataModule 中使用多个 DataLoader
在LightningDataModule中,可以通过数据加载钩子同时设置多个 DataLoader,Lightning 会自动选取对应的那一个:
class DataModule(LightningDataModule): def train_dataloader(self): # 任意可迭代对象或可迭代对象的集合 return DataLoader(self.train_dataset) def val_dataloader(self): # 任意可迭代对象或可迭代对象的集合 return [DataLoader(self.val_dataset_1), DataLoader(self.val_dataset_2)] def test_dataloader(self): # 任意可迭代对象或可迭代对象的集合 return DataLoader(self.test_dataset) def predict_dataloader(self): # 任意可迭代对象或可迭代对象的集合 return DataLoader(self.predict_dataset)例如上面的val_dataloader返回了两个DataLoader的列表,验证阶段会以sequential模式依次消费这两个验证集;若你在validation_step中需要区分批次来源,记得在钩子签名中加入dataloader_idx。
在 LightningModule 钩子中使用多个 DataLoader
与LightningDataModule完全相同的代码也可以在LightningModule中工作——直接覆写LightningModule的train_dataloader、val_dataloader、test_dataloader、predict_dataloader方法即可,返回单个迭代器或迭代器集合均被支持:
class MyModel(LightningModule): def train_dataloader(self): # 返回 DataLoader、原生迭代器、字典或列表均可 return {"main": DataLoader(self.dataset_a), "aux": DataLoader(self.dataset_b)} def training_step(self, batch, batch_idx): # batch 形如 {'main': ..., 'aux': ...} loss = ... return loss直接传入 Trainer 的数据加载参数
上述对任意可迭代对象(或可迭代对象集合)的支持同样适用于Trainer.fit、Trainer.validate、Trainer.test、Trainer.predict的 dataloader 参数,也就是说你完全可以把数据加载逻辑从模块中剥离,直接在调用时传入:
from lightning.pytorch import Trainer from lightning.pytorch.utilities import CombinedLoader trainer = Trainer() # fit:直接传字典形式的多个训练 DataLoader trainer.fit(model, train_dataloaders={"a": DataLoader(...), "b": DataLoader(...)}) # validate/test/predict:传列表,按 sequential 模式依次消费 trainer.validate(model, dataloaders=[DataLoader(val_1), DataLoader(val_2)]) trainer.predict(model, dataloaders=[DataLoader(pred_1), DataLoader(pred_2)]) # 需要自定义模式时,先构造 CombinedLoader 再传入 trainer.fit(model, train_dataloaders=CombinedLoader({"a": DataLoader(...)}, mode="max_size"))多 DataLoader 场景下的工程细节
- 状态保存与恢复:
CombinedLoader为内部实现了_Stateful接口(state_dict/load_state_dict)的 DataLoader 提供状态保存与恢复能力(_state_dicts/_load_state_dicts,见 combined_loader.py)。恢复时若 stateful 迭代器数量与 checkpoint 中的状态数量不一致,会抛出RuntimeError提示你保持与保存时一致的 DataLoader 定义。 - worker 清理:
CombinedLoader.reset()会重置内部迭代器并关闭各 DataLoader 的 worker(L361-L367)。_Sequential迭代器在同一时刻只启动一个 DataLoader 的 worker 集合,避免不必要的进程开销(对应 CHANGELOG 中「sequential 模式下按需启动 DataLoader workers」的改进)。 - 无长度迭代器:当某个可迭代对象没有
__len__(例如纯 iterable-style 数据集)时,长度计算将其视为float("inf")(_get_iterables_lengths,L404-L405);若所有数据集都是 iterable-style 且无法求长度,_dataset_length()会抛出NotImplementedError(对应测试 test_combined_loader.py)。 - 分布式采样器:传入的多个 DataLoader 会逐个经过
_process_dataloader处理(如自动挂接DistributedSampler),因此多 GPU 训练时每个迭代器都能获得正确的分布式采样行为。
小结
Trainer对任意可迭代对象及多迭代器集合的支持,让数据编排变得极其灵活:训练阶段默认以max_size_cycle让最长数据驱动 epoch 步数,验证/测试/预测阶段默认以sequential依次消费;需要精细控制时,可直接使用CombinedLoader选择min_size、max_size、sequential等模式,并结合limits控制每个迭代器的批次上限。理解这四种模式的合并语义、各阶段默认值与限制(fit不支持sequential、predict仅支持sequential),以及sequential模式下dataloader_idx钩子参数的约定,是编写多数据集 Lightning 程序的坚实基础。
【免费下载链接】pytorch-lightningPretrain, finetune ANY AI model of ANY size on 1 or 10,000+ GPUs with zero code changes.项目地址: https://gitcode.com/gh_mirrors/py/pytorch-lightning
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考