PyTorch Lightning 任意可迭代对象与多 DataLoader 支持:CombinedLoader 模式详解
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它们都继承自_ModeIteratorcombined_loader.py统一返回(batch, batch_idx, dataloader_idx)三元组_MaxSizeCycleL67-L106某个迭代器抛出StopIteration时标记为已耗尽若还有未耗尽的迭代器则用iter(self.iterables[i])重新创建迭代器继续循环取数_MaxSizeL184-L205抑制StopIteration耗尽后把该位置的输出置为None直到所有迭代器都耗尽_SequentialL123-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(modesequential) mode.)prediction_loop.py。trainer.fit不支持sequential模式fit_loop对传入的CombinedLoader模式有显式校验其余模式的组合方式由_SUPPORTED_MODES决定。手动选择模式如果默认模式不满足需求可以直接使用CombinedLoader并指定mode再把它传给Trainerfrom lightning.pytorch.utilities import CombinedLoader iterables {a: DataLoader(), b: DataLoader()} combined_loader CombinedLoader(iterables, modemin_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_size4), b: DataLoader(range(15), batch_size5)}即a有 2 个 batch、b有 3 个 batchmax_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 个 batcha耗尽即停止sequential共 5 个 batch先完整输出a的 2 个 batchdataloader_idx0再输出b的 3 个 batchdataloader_idx1。这些行为在仓库测试 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.limitsfit_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_idx0)L94与on_validation_batch_endon_test_batch_start(batch, batch_idx, dataloader_idx0)L117与on_test_batch_endon_predict_batch_start(batch, batch_idx, dataloader_idx0)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_cycledataloader_idx恒为 0因为每步都会取所有迭代器的批次并合并为一个 batch 结构。在 LightningDataModule 中使用多个 DataLoader在LightningDataModule中可以通过数据加载钩子同时设置多个 DataLoaderLightning 会自动选取对应的那一个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_dataloadersCombinedLoader({a: DataLoader(...)}, modemax_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 的 workerL361-L367。_Sequential迭代器在同一时刻只启动一个 DataLoader 的 worker 集合避免不必要的进程开销对应 CHANGELOG 中「sequential 模式下按需启动 DataLoader workers」的改进。无长度迭代器当某个可迭代对象没有__len__例如纯 iterable-style 数据集时长度计算将其视为float(inf)_get_iterables_lengthsL404-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),仅供参考
上一篇/下一篇内容由系统自动关联
返回资讯列表 →