十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

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

PyTorch Lightning TPU 训练中级实战:分布式采样器、核心数配置与 16 位精度

PyTorch Lightning TPU 训练中级实战:分布式采样器、核心数配置与 16 位精度 PyTorch Lightning TPU 训练中级实战分布式采样器、核心数配置与 16 位精度【免费下载链接】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导读本篇中级指南面向希望在云端 TPUCloud TPU上运行 PyTorch Lightning 训练任务的开发者围绕三个核心实战点展开分布式采样器的自动处理机制、TPU 核心数1 或 8的配置方式以及 TPU 上的 16 位精度训练。读完本文你将理解 Lightning 如何自动为 TPU 插入 DistributedSampler、正确配置L.Trainer(acceleratortpu, devicesN)并掌握precision16-true与bf16-true的底层原理直接上手云 TPU 训练。本文基于仓库中的 TPU 中级指南docs/source-pytorch/accelerators/tpu_intermediate.rst撰写并结合src/lightning下的源码与测试用例进行深度验证。实验性功能警告本指南涉及的 TPUXLA训练属于实验性ExperimentalAPI接口与行为可能在未来版本中变化生产环境使用前请关注版本变更说明。一、使用云 TPU 前的环境认知在开始之前先明确 TPU 训练在 Lightning 中的技术定位。TPU 设备由 PyTorch/XLAtorch_xla提供支持Lightning 通过 XLA 加速器、XLA 策略与 XLA 精度插件三个层次将其接入训练流程。从源码结构看相关实现分布如下加速器src/lightning/fabric/accelerators/xla.pyFabric 基础实现与src/lightning/pytorch/accelerators/xla.pyPyTorch 扩展策略src/lightning/fabric/strategies/xla.py与src/lightning/pytorch/strategies/xla.py精度插件src/lightning/fabric/plugins/precision/xla.py与src/lightning/pytorch/plugins/precision/xla.py集群环境src/lightning/fabric/plugins/environments/xla.py在 XLAAccelerator 实现 中可以看到两个硬性环境约束必须安装torch_xla源码中要求torch_xla1.13见src/lightning/fabric/accelerators/xla.py且运行时必须使用 PJRT 运行时旧的 XRT 运行时已不再支持初始化时若未使用 PJRT 会直接抛出RuntimeError(The XLA XRT runtime is not supported anymore.)。适用前提以下所有配置均假设你已在云 TPU 虚拟机如 Google Cloud TPU VM上正确安装 PyTorch 与torch_xla且环境变量与运行时满足 PJRT 要求。二、分布式采样器Lightning 自动处理无需手动定义在原生 PyTorch 中当使用 TPU或 DDP进行多设备分布式训练时你需要手动构造torch.utils.data.distributed.DistributedSampler确保每个设备拿到属于自己的那一份数据分片。在 Lightning 中这一步完全不需要你操心——框架会在训练启动时自动插入正确的采样器把正确的数据块分配到对应的 TPU 核心上。注意不要在自定义的train_dataloader()中手动添加DistributedSamplerLightning 会自动完成这一操作。重复添加可能造成采样逻辑冲突。自动插入的机制可以从策略源码中得到印证。在 XLAStrategy 中distributed_sampler_kwargs属性显式返回了采样器所需的副本数与秩信息property def distributed_sampler_kwargs(self) - dict[str, int]: return {num_replicas: self.world_size, rank: self.global_rank}Lightning 内部正是基于num_replicas副本数与rank进程秩为你的 DataLoader 自动构造并注入DistributedSamplerworld_size与global_rank则由 XLAEnvironment 从 XLA 运行时读取如xr.world_size()、xr.global_ordinal()。万一确实需要手动构造采样器若出于某些特殊原因例如完全绕过 Lightning 的数据管线你仍然需要手动构造采样器可以参考以下示例取自原文档并保留完整上下文import torch_xla.core.xla_model as xm def train_dataloader(self): dataset MNIST(os.getcwd(), trainTrue, downloadTrue, transformtransforms.ToTensor()) # required for TPU support sampler None if use_tpu: sampler torch.utils.data.distributed.DistributedSampler( dataset, num_replicasxm.xrt_world_size(), rankxm.get_ordinal(), shuffleTrue ) loader DataLoader(dataset, samplersampler, batch_size32) return loader其中xm.xrt_world_size()返回 XLA 设备总数即副本数xm.get_ordinal()返回当前进程的全局序号即秩。需要说明的是这两个 API 属于torch_xla历史接口在较新版本torch_xla2.1中对应地可改用torch_xla.runtime的world_size()与global_ordinal()——XLAEnvironment 源码正是按此分支实现的见src/lightning/fabric/plugins/environments/xla.py。但如前所述在 Lightning 中通常不需要走到这一步。三、配置 TPU 核心数只能选 1 或 8在 Trainer 中配置 TPU 核心数非常简单。单个 TPU 设备如 TPU v2-8 单板提供 8 个核心因此devices参数只能取 1 或 8import lightning as L my_model MyLightningModule() trainer L.Trainer(acceleratortpu, devices8) trainer.fit(my_model)仅此而已你的模型将在这 8 个 TPU 核心上并行训练。若只想使用单核心将devices改为1即可。要使用整片 TPU Pod多台主机、更多核心请参考本文末尾指向的 TPU Pod 相关章节。设备校验的源码细节devices参数并非任意取值都能通过校验。在 XLA 加速器的设备解析与校验逻辑 中_check_tpu_devices_valid明确限制了合法取值def _check_tpu_devices_valid(devices: object) - None: device_count XLAAccelerator.auto_device_count() if ( # support number of devices isinstance(devices, int) and devices in {1, device_count} # support picking a specific device or isinstance(devices, (list, tuple)) and len(devices) 1 and 0 devices[0] device_count - 1 ): return raise ValueError( fdevices can only be auto, 1, {device_count} or [0-{device_count - 1}] for TPUs. Got {devices!r} )几个值得注意的要点合法取值auto、1、设备总数如 8或形如[0][7]的单元素列表用于指定使用某一个具体的 TPU 核心。其他数值会抛出ValueError。设备总数并非固定为 8auto_device_count()会根据torch_xla版本动态获取可用设备数——在torch_xla2.1时通过tpu.num_available_devices()查询旧版本则按 TPU 版本映射v2/v3 为 8v4 为 4见src/lightning/fabric/accelerators/xla.py中的device_count_on_version。因此 v4 TPU 上devices的合法值实际上是 1 或 4。字符串形式devices也支持字符串如8或0,1,2会被解析为整数或列表_parse_tpu_devices_str。单核心的运行时限制值得注意的是XLAStrategy在多进程启动模式下并不支持只在单设备上运行PJRT 运行时的限制。在 XLAStrategy.setup_distributed 与 Fabric 版的setup_environment见src/lightning/fabric/strategies/xla.py中当parallel_devices长度为 1 时会抛出NotImplementedError并提示改用SingleDeviceXLAStrategy策略。也就是说devices1在 Trainer 中应配合单设备策略使用8 核心全量并行才是XLAStrategy的主战场。四、TPU 上的 16 位精度训练Lightning 也支持在 TPU 上进行 16 位精度训练。默认情况下TPU 训练使用 32 位精度如需启用 16 位在 Trainer 中传入precision参数即可import lightning as L my_model MyLightningModule() trainer L.Trainer(acceleratortpu, precision16-true) trainer.fit(my_model)两种半精度模式fp16 与 bf16从当前仓库的源码看XLA 精度插件支持三种取值32-true默认、16-truefp16与bf16-truebfloat16定义在 Fabric XLAPrecision 的类型别名_PRECISION_INPUT中。其核心实现通过设置环境变量来切换 XLA 的精度行为if precision 16-true: os.environ[XLA_USE_F16] 1 self._desired_dtype torch.float16 elif precision bf16-true: os.environ[XLA_USE_BF16] 1 self._desired_dtype torch.bfloat16 else: self._desired_dtype torch.float32即precision取值环境变量期望 dtype说明32-true无默认torch.float32全精度TPU 训练的默认行为16-trueXLA_USE_F161torch.float16启用 FP16bf16-trueXLA_USE_BF161torch.bfloat16启用 bfloat16TPU尤其是 TPU v2/v3硬件对 bfloat16 有原生支持其动态范围与 fp32 相同是 TPU 上常用的半精度格式。原文档提到的 Under the hood the xla library will use the bfloat16 type 对应的是bf16-true模式当前实现同时提供了16-true的 fp16 选项两者都能显著降低显存占用并通常提升吞吐。校验与清理逻辑输入校验XLAPrecision.__init__会拒绝不支持的取值如16、16-mixed、bf16-mixed、64-true抛出ValueError。这一行为在 测试用例 tests/tests_fabric/plugins/precision/test_xla.py 中被完整覆盖验证。也就是说TPU 上不支持混合精度自动缩放AMP mixed-precision模式只支持真 16 位/真 32 位的显式精度。环境变量清理训练结束时teardown()会弹出XLA_USE_BF16与XLA_USE_F16环境变量见src/lightning/fabric/plugins/precision/xla.py避免污染后续任务。对应的清理行为同样有测试验证test_teardown。优化器步进优化XLA 精度插件还接管了optimizer_step在 Fabric 版中使用xm.optimizer_step(optimizer, optimizer_argskwargs, barrierTrue)——设置barrierTrue是因为在optimizer.step之后始终执行xm.mark_step()对性能更有利见src/lightning/fabric/plugins/precision/xla.py#L59-L68PyTorch 版则包装 closure 并在 step 后调用xm.mark_step()见src/lightning/pytorch/plugins/precision/xla.py#L62-L85。五、TPU 训练在 Lightning 中的底层工作流理解底层机制有助于排查问题。结合策略与启动器源码一次 TPU 训练的大致流程如下进程启动_XLALauncher调用torch_xla.distributed.xla_multiprocessing.spawnxmp.spawn在 N 个核心上启动工作进程启动方式为fork并要求入口脚本受if __name__ __main__保护见 src/lightning/fabric/strategies/launchers/xla.py。数据加载XLAStrategy.process_dataloader将用户 DataLoader 包装为torch_xla.distributed.parallel_loader.MpDeviceLoader使数据批次在 XLA 设备上并行加载见src/lightning/pytorch/strategies/xla.py#L179-L190。模型同步setup阶段通过broadcast_master_param将主进程参数广播到所有核心保持各核心模型初始状态一致见src/lightning/pytorch/strategies/xla.py#L155-L161。梯度与保存梯度归约在optimizer.step内完成Fabric 版_backward_sync_control None即为此设计见src/lightning/fabric/strategies/xla.py#L58保存 checkpoint 前会先xm.mark_step()同步所有待执行的惰性张量避免集体操作挂起见src/lightning/pytorch/strategies/xla.py#L299-L308。需要说明以上流程细节是基于src/lightning当前源码结构推断的调用关系具体行号以仓库实际内容为准。六、延伸阅读与注意事项完整使用整片 TPU Pod多主机多核心的场景不在本文中级指南范围内请参考仓库中的 TPU 高级指南 与 TPU 基础指南。训练中可通过XLAAccelerator.get_device_stats()获取每个 XLA 设备的空闲内存与峰值内存统计见 src/lightning/pytorch/accelerators/xla.py配合 profiler 分析 TPU 利用率。TPU 上不支持跳过 backward即在training_step中返回None的自动优化模式会触发MisconfigurationException见src/lightning/pytorch/plugins/precision/xla.py#L76-L84编写训练逻辑时需留意。XLA 策略在进程未启动时访问root_device、world_size等属性会抛出RuntimeError相关属性在_launched为假时返回占位值见src/lightning/pytorch/strategies/xla.py这表明分布式信息的读取被严格限定在 spawn 出的工作进程内。小结本篇围绕 TPU 训练的三大中级主题给出了可直接落地的配置方案分布式采样器完全交给 Lightning 自动处理distributed_sampler_kwargs自动注入num_replicas与rank核心数配置为L.Trainer(acceleratortpu, devices1|8)实际上限取决于 TPU 版本v4 为 416 位精度通过precision16-true或bf16-true开启底层由XLA_USE_F16/XLA_USE_BF16环境变量驱动。掌握这三项能力即可在云 TPU 上以接近零样板代码的方式启动并行训练任务。【免费下载链接】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),仅供参考
返回列表