ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

PyTorch Lightning 数据管理实战:从零掌握 LightningDataModule 的完整 API 与分布式数据管线

PyTorch Lightning 数据管理实战:从零掌握 LightningDataModule 的完整 API 与分布式数据管线 人工智能深度学习机器学习预训练分布式训练微调【免费下载链接】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 仓库中的 LightningDataModule 官方指南 为骨架系统讲解如何用LightningDataModule把 PyTorch 数据管线中下载、清洗、切分、变换、装载五个步骤封装为可共享、可复用、可热切换的数据模块。读完本文你将掌握prepare_data、setup、四个*_dataloader钩子与prepare_data_per_node等核心 API 的调用时机与最佳实践理解其在多进程/多节点训练下的底层行为并学会在 Trainer、纯 PyTorch 代码与 checkpoint 三种场景中正确使用 DataModule。为什么需要 DataModule让数据准备告别散落各处在常规 PyTorch 代码里数据清洗与准备工作通常分散在多个文件、多个脚本中导致切分方式变换方式归一化参数tokenize 细节无法跨项目共享与复用。LightningDataModule就是为解决这个问题而生的可共享、可复用的类它把 PyTorch 数据处理涉及的五步全部封装起来下载 / tokenize / 预处理Download / tokenize / process清洗并可选保存到磁盘Clean and maybe save to disk加载到torch.utils.data.Dataset中应用变换旋转、tokenize 等包装进torch.utils.data.DataLoader。封装完成后同一个类可以在任何地方被共享和使用——换数据集、换项目都无需重写数据逻辑model LitClassifier() trainer Trainer() imagenet ImagenetDataModule() trainer.fit(model, datamoduleimagenet) cifar10 CIFAR10DataModule() trainer.fit(model, datamodulecifar10)如果你在项目交接或复现实验时问过这些问题你用的是什么切分你用的什么变换归一化参数是什么数据是怎么预处理/tokenize 的——那么 DataModule 正是为你准备的。什么是 LightningDataModuleLightningDataModule实现在 src/lightning/pytorch/core/datamodule.py是 PyTorch Lightning 中管理数据的便捷类它封装了训练、验证、测试、预测四类 dataloader以及数据处理、下载、变换所需的全部步骤。借助它你可以开发与具体数据集解耦的模型、在数据集之间热切换、跨项目共享数据切分与变换。从源码结构看LightningDataModule继承自DataHooks与HyperparametersMixin见 datamodule.py因此它既拥有全套数据钩子也天然支持超参数保存同时它还持有一个指向Trainer的引用self.trainer供钩子内部访问训练状态。普通 PyTorch 写法 vs DataModule 写法先看常规 PyTorch 的典型代码# regular PyTorch test_data MNIST(my_path, trainFalse, downloadTrue) predict_data MNIST(my_path, trainFalse, downloadTrue) train_data MNIST(my_path, trainTrue, downloadTrue) train_data, val_data random_split(train_data, [55000, 5000]) train_loader DataLoader(train_data, batch_size32) val_loader DataLoader(val_data, batch_size32) test_loader DataLoader(test_data, batch_size32) predict_loader DataLoader(predict_data, batch_size32)等价的 DataModule 只是把同样的代码组织起来但使其可跨项目复用class MNISTDataModule(L.LightningDataModule): def __init__(self, data_dir: str path/to/dir, batch_size: int 32): super().__init__() self.data_dir data_dir self.batch_size batch_size def setup(self, stage: str): self.mnist_test MNIST(self.data_dir, trainFalse) self.mnist_predict MNIST(self.data_dir, trainFalse) mnist_full MNIST(self.data_dir, trainTrue) self.mnist_train, self.mnist_val random_split( mnist_full, [55000, 5000], generatortorch.Generator().manual_seed(42) ) def train_dataloader(self): return DataLoader(self.mnist_train, batch_sizeself.batch_size) def val_dataloader(self): return DataLoader(self.mnist_val, batch_sizeself.batch_size) def test_dataloader(self): return DataLoader(self.mnist_test, batch_sizeself.batch_size) def predict_dataloader(self): return DataLoader(self.mnist_predict, batch_sizeself.batch_size) def teardown(self, stage: str): # Used to clean-up when the run is finished ...当处理流程变复杂引入更多变换、多 GPU 训练时可以让 Lightning 替你处理分布式细节同时数据集依旧可复用、可共享mnist MNISTDataModule(my_path) model LitClassifier() trainer Trainer() trainer.fit(model, mnist)一个更完整的实战示例下面是一个更贴近真实项目的完整 DataModule注意运行需要安装 torchvisionimport lightning as L from torch.utils.data import random_split, DataLoader # Note - you must have torchvision installed for this example from torchvision.datasets import MNIST from torchvision import transforms class MNISTDataModule(L.LightningDataModule): def __init__(self, data_dir: str ./): super().__init__() self.data_dir data_dir self.transform transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,))]) def prepare_data(self): # download MNIST(self.data_dir, trainTrue, downloadTrue) MNIST(self.data_dir, trainFalse, downloadTrue) def setup(self, stage: str): # Assign train/val datasets for use in dataloaders if stage fit: mnist_full MNIST(self.data_dir, trainTrue, transformself.transform) self.mnist_train, self.mnist_val random_split( mnist_full, [55000, 5000], generatortorch.Generator().manual_seed(42) ) # Assign test dataset for use in dataloader(s) if stage test: self.mnist_test MNIST(self.data_dir, trainFalse, transformself.transform) if stage predict: self.mnist_predict MNIST(self.data_dir, trainFalse, transformself.transform) def train_dataloader(self): return DataLoader(self.mnist_train, batch_size32) def val_dataloader(self): return DataLoader(self.mnist_val, batch_size32) def test_dataloader(self): return DataLoader(self.mnist_test, batch_size32) def predict_dataloader(self): return DataLoader(self.mnist_predict, batch_size32)这个示例同时演示了stage参数的正确用法fit阶段只需要训练集与验证集test与predict阶段再按需惰性加载测试/预测数据避免一次性把所有数据载入内存。LightningDataModule 完整 API 深度解析定义一个 DataModule核心就是实现下面这些用于创建 train/val/test/predict dataloader 的方法其基类定义可在 src/lightning/pytorch/core/hooks.py 的DataHooks中查看prepare_data如何下载、tokenize 等setup如何切分、定义数据集等train_dataloaderval_dataloadertest_dataloaderpredict_dataloaderprepare_data只在单进程执行的下载/预处理多进程分布式环境下同时下载或保存数据会造成数据损坏。Lightning 保证prepare_data只在单个 CPU 进程内调用因此可以安全地把下载逻辑放进去。多节点训练时它的执行范围由prepare_data_per_node决定。prepare_data之后才会调用setup两者之间有一个 barrier确保所有进程在数据准备完成、可被使用之后再进入setup。prepare_data适合做这些事下载数据——只在磁盘上从单进程下载一次tokenize——一次性过程不适合在全部进程上重复执行其他一次性 IO 操作。class MNISTDataModule(L.LightningDataModule): def prepare_data(self): # download MNIST(os.getcwd(), trainTrue, downloadTrue, transformtransforms.ToTensor()) MNIST(os.getcwd(), trainFalse, downloadTrue, transformtransforms.ToTensor())警告prepare_data是在主进程中被调用的不建议在这里给状态赋值如self.x y。因为它只在一个进程上执行这里赋的状态对其他进程不可见。需要给每个进程都可见的状态请放在setup里。这一点在 hooks.py 的 docstring 中也有明确示例good: download_data/tokenizebad: self.split ...。从源码看prepare_data的调用逻辑位于 data_connector.py它会读取datamodule.prepare_data_per_node在_InfiniteBarrier()保护下仅当(prepare_data_per_node and local_rank_zero)或(not prepare_data_per_node and global_rank_zero)时才真正调用钩子——这正是每个节点只跑一次或全局只跑一次两种模式的实现依据。setup在每个进程上执行的切分与状态构建setup适合在每个 GPU 进程上执行的操作统计类别数量构建词表vocabulary执行 train/val/test 切分创建数据集应用在 datamodule 中显式定义的变换。import lightning as L class MNISTDataModule(L.LightningDataModule): def setup(self, stage: str): # Assign Train/val split(s) for use in Dataloaders if stage fit: mnist_full MNIST(self.data_dir, trainTrue, downloadTrue, transformself.transform) self.mnist_train, self.mnist_val random_split( mnist_full, [55000, 5000], generatortorch.Generator().manual_seed(42) ) # Assign Test split(s) for use in Dataloaders if stage test: self.mnist_test MNIST(self.data_dir, trainFalse, downloadTrue, transformself.transform)对于 NLP 任务典型做法是prepare_data里做 tokenize 并保存到磁盘setup里再加载回来class LitDataModule(L.LightningDataModule): def prepare_data(self): dataset load_Dataset(...) train_dataset ... val_dataset ... # tokenize # save it to disk def setup(self, stage): # load it back here dataset load_dataset_from_disk(...)setup方法要求一个stage参数用来区分trainer.{fit,validate,test,predict}的不同初始化逻辑。注意setup会在所有节点上的每个进程中被调用在这里设置状态是推荐做法。同样地teardown也由每个节点的每个进程调用用于清理状态。四个 dataloader 钩子四个钩子的用法完全对称——通常在setup中定义好数据集后在这里用DataLoader包装返回train_dataloader生成训练 dataloader被Trainer.fit使用。val_dataloader生成验证 dataloader被Trainer.fit和Trainer.validate使用。test_dataloader生成测试 dataloader被Trainer.test使用。predict_dataloader生成预测 dataloader被Trainer.predict使用。import lightning as L class MNISTDataModule(L.LightningDataModule): def train_dataloader(self): return DataLoader(self.mnist_train, batch_size64) def val_dataloader(self): return DataLoader(self.mnist_val, batch_size64) def test_dataloader(self): return DataLoader(self.mnist_test, batch_size64) def predict_dataloader(self): return DataLoader(self.mnist_predict, batch_size64)从DataHooks的基类实现看这四个方法在未实现时会抛出MisconfigurationException如 hooks.py提示必须实现才能与 Lightning Trainer 配合使用。基类 docstring 还说明Lightning 会自动为分布式环境添加正确的 sampler如DistributedSampler无需你手动设置如果没有 test/val 数据集及对应 step 方法也可以不实现对应钩子。transfer_batch_to_device自定义 batch 的设备迁移当你的DataLoader返回的是自定义数据结构而非开箱即支持的torch.Tensor、list、dict、tuple及其任意嵌套时重写此钩子来定义数据如何迁移到目标设备CPU、GPU、TPU 等。def transfer_batch_to_device(self, batch, device, dataloader_idx): if isinstance(batch, CustomBatch): # move all tensors in your custom data structure to the device batch.samples batch.samples.to(device) batch.targets batch.targets.to(device) elif dataloader_idx 0: # skip device transfer for the first dataloader or anything you wish pass else: batch super().transfer_batch_to_device(batch, device, dataloader_idx) return batch基类默认实现调用move_data_to_device(batch, device)见 hooks.py。官方建议该钩子只负责迁移数据、不要修改数据也不要把数据迁到参数之外的其他设备。on_before_batch_transfer 与 on_after_batch_transfer批量级变换on_before_batch_transfer在 batch 被迁移到设备之前执行适合做批量级增广augmentation例如在 CPU 上对batch[x]施加变换。on_after_batch_transfer在 batch 被迁移到设备之后执行适合做 GPU 上的变换如gpu_transforms。def on_before_batch_transfer(self, batch, dataloader_idx): batch[x] transforms(batch[x]) return batch def on_after_batch_transfer(self, batch, dataloader_idx): batch[x] gpu_transforms(batch[x]) return batch注意结合self.trainer.training / validating / testing / predicting标志可以在同一钩子内针对不同阶段施加不同逻辑。state_dict 与 load_state_dictDataModule 状态随 checkpoint 保存state_dict在保存 checkpoint 时被调用返回需要持久化的 datamodule 状态字典默认空字典load_state_dict在加载 checkpoint 时被调用用于恢复状态。源码实现在 datamodule.py。原始文档在结尾通过 extensions/datamodules_state.rst 给出了配套示例import lightning as L class LitDataModule(L.LightningDataModule): def state_dict(self): # track whatever you want here state {current_train_batch_index: self.current_train_batch_index} return state def load_state_dict(self, state_dict): # restore the state based on what you tracked in (def state_dict) self.current_train_batch_index state_dict[current_train_batch_index]一旦你的 DataModule 定义了这两个方法checkpoint 就会自动追踪并恢复DataModule 的状态无需额外手工处理。仓库测试 test_datamodules.pytest_dm_checkpoint_save_and_load验证了保存后 checkpoint 中会出现以 DataModule 类名dm.__class__.__qualname__为键的状态字典加载后load_state_dict会被正确调用以恢复状态。teardown运行结束时的清理teardown(stage)在 fittrain validate、validate、test 或 predict 结束时调用适合删除临时文件、释放资源等清理工作并且会在每个进程上执行。prepare_data_per_node控制 prepare_data 的执行范围设为True默认prepare_data()在每个节点的LOCAL_RANK0进程上调用设为False只在NODE_RANK0, LOCAL_RANK0的全局主进程上调用一次适合共享文件系统场景。class LitDataModule(LightningDataModule): def __init__(self): super().__init__() self.prepare_data_per_node True该属性定义在DataHooks.__init__中hooks.py默认值为True。配套还有一个allow_zero_length_dataloader_with_multiple_devices属性默认False用于控制多设备下是否允许某个 local rank 返回零长度的 dataloader。分布式下prepare_data的调用时序在 data_connector.py 中有精确实现_InfiniteBarrier保证各进程在数据就绪前不越界。使用 DataModule推荐方式与进阶技巧推荐方式直接传给 Trainerdm MNISTDataModule() model Model() trainer.fit(model, datamoduledm) trainer.test(datamoduledm) trainer.validate(datamoduledm) trainer.predict(datamoduledm)需要先构建模型时手动调用 prepare_data / setup当模型需要依赖数据集信息如类别数、输入宽度、词表大小来构建时可以先手动运行prepare_data和setupLightning 仍会保证它们在正确的设备上执行dm MNISTDataModule() dm.prepare_data() dm.setup(stagefit) model Model(num_classesdm.num_classes, widthdm.width, vocabdm.vocab) trainer.fit(model, dm) dm.setup(stagetest) trainer.test(datamoduledm)从 Trainer 反向访问数据可以通过trainer.datamodule访问当前正在使用的 DataModule并通过 Trainer 的train_dataloader、val_dataloaders、test_dataloaders、predict_dataloaders属性访问当前 dataloader。这些关联在_DataConnector.attach_datamodule中建立——该函数把 datamodule 实例挂到各 loop 的_data_source上同时设置trainer.datamodule datamodule与datamodule.trainer trainer见 data_connector.py。脱离 Lightning 使用 DataModule纯 PyTorch 场景DataModule 只是普通的 Python 类自然也可以在纯 PyTorch 代码中使用# download, etc... dm MNISTDataModule() dm.prepare_data() # splits/transforms dm.setup(stagefit) # use data for batch in dm.train_dataloader(): ... for batch in dm.val_dataloader(): ... dm.teardown(stagefit) # lazy load test data dm.setup(stagetest) for batch in dm.test_dataloader(): ... dm.teardown(stagetest)即使不引入 TrainerDataModule 也通过把数据集的所有细节统一进一个结构提升了实验的可复现性。DataModule 中的超参数与 LightningModule 相同的 API与 LightningModule 一样DataModule 支持通过save_hyperparameters保存超参数import lightning as L class CustomDataModule(L.LightningDataModule): def __init__(self, *args, **kwargs): super().__init__() self.save_hyperparameters() def configure_optimizers(self): # access the saved hyperparameters opt optim.Adam(self.parameters(), lrself.hparams.lr)关于save_hyperparameters的更多细节可参考 LightningModule 文档。这种能力来自HyperparametersMixindatamodule.py它让 DataModule 与 LightningModule 共享同一套 hparams 机制仓库测试test_hyperparameters_savingtest_datamodules.py也验证了这一行为。进阶能力from_datasets、load_from_checkpoint 与字符串表示除文档主线外从源码还可以看到几个值得掌握的便捷能力均在 datamodule.py 中实现from_datasetsdatamodule.py无需手写任何钩子直接从一个或多个torch.utils.data.Dataset构造 DataModule。batch_size与num_workers会被自动注入到各 dataloader仅当你的__init__接受这些参数时额外参数通过**datamodule_kwargs透传。训练 loader 自动开启shuffleTrue并设置pin_memoryTrue。测试见test_dm_init_from_datasets_*系列test_datamodules.py。dm LightningDataModule.from_datasets(train_ds, val_ds, test_ds, predict_ds, batch_size32, num_workers4)load_from_checkpointdatamodule.py类方法从 checkpoint 恢复 DataModule。Lightning 保存 checkpoint 时会把__init__参数存到datamodule_hyper_parameters键下**kwargs可覆盖已保存的超参数也可用hparams_file传入.yaml/.csv补充缺失参数。注意必须用 DataModule类调用而非实例否则会抛TypeError。datamodule MyLightningDataModule.load_from_checkpoint(path/to/checkpoint.ckpt) datamodule MyLightningDataModule.load_from_checkpoint(PATH, batch_size32, num_workers10)__str__字符串表示datamodule.py打印 DataModule 时会输出四个 dataloader 及其数据集大小信息例如{Train dataloader: size55000}数据集不可用时显示None。对应测试为test_datamodule_string_*系列test_datamodules.py。小结LightningDataModule的核心价值在于把数据管线的五个步骤收敛到一个可共享、可复用、可热切换的类中单进程 IO 进prepare_data由prepare_data_per_node控制每节点或全局执行一次中间有 barrier 保证数据就绪每进程状态构建进setup按stagefit/validate/test/predict惰性创建对应数据集四个 dataloader 钩子只负责把数据集包装成DataLoader分布式 sampler 由 Lightning 自动处理批量级变换通过transfer_batch_to_device、on_before_batch_transfer、on_after_batch_transfer三个钩子完成状态持久化通过state_dict/load_state_dict随 checkpoint 自动保存与恢复超参数则与 LightningModule 共用save_hyperparametersAPI。无论你是在单卡调试、多 GPU/多节点训练还是纯 PyTorch 脚本中组织数据把数据逻辑封装进 DataModule 都是提升可复现性与协作效率的推荐做法。文中所有结论均有 docs/source-pytorch/data/datamodule.rst、src/lightning/pytorch/core/datamodule.py、src/lightning/pytorch/core/hooks.py、src/lightning/pytorch/trainer/connectors/data_connector.py 与 tests/tests_pytorch/core/test_datamodules.py 作为仓库内依据可进一步阅读源码验证。赞分享人工智能深度学习机器学习预训练分布式训练微调【免费下载链接】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 数据管道完全指南用 LightningDataModule 封装可复用的数据处理流程PyTorch Lightning 数据管道完全指南用 LightningDataModule 封装可复用的数据处理流程 LightningDataModulAI 技能科研生物信息学数据科学5分钟快速掌握分布式数据分片技术从零到实战完整指南5分钟快速掌握分布式数据分片技术从零到实战完整指南 在当今数据爆炸的时代 分布式数据分片技术 已成为企业应对海量数据和高并发访问的 核心解决方案 。通过数据低代码后端前端AI 应用大模型RAG工作流自动化Vitess分布式数据库完全指南从零开始掌握大规模MySQL集群管理Vitess分布式数据库完全指南从零开始掌握大规模MySQL集群管理 Vitess是一个用于大规模数据库管理的开源系统基于MySQL构建提供高性能、可扩展数据库分布式数据库云原生后端数据存储创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表