news 2026/9/19 13:01:39

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

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
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

导读:本篇中级指南面向希望在云端 TPU(Cloud TPU)上运行 PyTorch Lightning 训练任务的开发者,围绕三个核心实战点展开:分布式采样器的自动处理机制、TPU 核心数(1 或 8)的配置方式,以及 TPU 上的 16 位精度训练。读完本文,你将理解 Lightning 如何自动为 TPU 插入 DistributedSampler、正确配置L.Trainer(accelerator="tpu", devices=N),并掌握precision="16-true""bf16-true"的底层原理,直接上手云 TPU 训练。

本文基于仓库中的 TPU 中级指南(docs/source-pytorch/accelerators/tpu_intermediate.rst)撰写,并结合src/lightning下的源码与测试用例进行深度验证。

实验性功能警告:本指南涉及的 TPU(XLA)训练属于实验性(Experimental)API,接口与行为可能在未来版本中变化,生产环境使用前请关注版本变更说明。


一、使用云 TPU 前的环境认知

在开始之前,先明确 TPU 训练在 Lightning 中的技术定位。TPU 设备由 PyTorch/XLA(torch_xla)提供支持,Lightning 通过 XLA 加速器、XLA 策略与 XLA 精度插件三个层次将其接入训练流程。从源码结构看,相关实现分布如下:

  • 加速器:src/lightning/fabric/accelerators/xla.py(Fabric 基础实现)与src/lightning/pytorch/accelerators/xla.py(PyTorch 扩展)
  • 策略:src/lightning/fabric/strategies/xla.pysrc/lightning/pytorch/strategies/xla.py
  • 精度插件:src/lightning/fabric/plugins/precision/xla.pysrc/lightning/pytorch/plugins/precision/xla.py
  • 集群环境:src/lightning/fabric/plugins/environments/xla.py

在 XLAAccelerator 实现 中可以看到两个硬性环境约束:必须安装torch_xla(源码中要求torch_xla>=1.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()中手动添加DistributedSampler,Lightning 会自动完成这一操作。重复添加可能造成采样逻辑冲突。

自动插入的机制可以从策略源码中得到印证。在 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_sizeglobal_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(), train=True, download=True, transform=transforms.ToTensor()) # required for TPU support sampler = None if use_tpu: sampler = torch.utils.data.distributed.DistributedSampler( dataset, num_replicas=xm.xrt_world_size(), rank=xm.get_ordinal(), shuffle=True ) loader = DataLoader(dataset, sampler=sampler, batch_size=32) return loader

其中xm.xrt_world_size()返回 XLA 设备总数(即副本数),xm.get_ordinal()返回当前进程的全局序号(即秩)。需要说明的是,这两个 API 属于torch_xla历史接口,在较新版本(torch_xla>=2.1)中对应地可改用torch_xla.runtimeworld_size()global_ordinal()——XLAEnvironment 源码正是按此分支实现的(见src/lightning/fabric/plugins/environments/xla.py)。但如前所述,在 Lightning 中通常不需要走到这一步


三、配置 TPU 核心数:只能选 1 或 8

在 Trainer 中配置 TPU 核心数非常简单。单个 TPU 设备(如 TPU v2-8 单板)提供 8 个核心,因此devices参数只能取 1 或 8

import lightning as L my_model = MyLightningModule() trainer = L.Trainer(accelerator="tpu", devices=8) 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( f"`devices` can only be 'auto', 1, {device_count} or [<0-{device_count - 1}>] for TPUs. Got {devices!r}" )

几个值得注意的要点:

  1. 合法取值'auto'1、设备总数(如 8),或形如[0][7]的单元素列表用于指定使用某一个具体的 TPU 核心。其他数值会抛出ValueError
  2. 设备总数并非固定为 8auto_device_count()会根据torch_xla版本动态获取可用设备数——在torch_xla>=2.1时通过tpu.num_available_devices()查询;旧版本则按 TPU 版本映射(v2/v3 为 8,v4 为 4,见src/lightning/fabric/accelerators/xla.py中的device_count_on_version)。因此 v4 TPU 上devices的合法值实际上是 1 或 4。
  3. 字符串形式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策略。也就是说:devices=1在 Trainer 中应配合单设备策略使用,8 核心全量并行才是XLAStrategy的主战场。


四、TPU 上的 16 位精度训练

Lightning 也支持在 TPU 上进行 16 位精度训练。默认情况下,TPU 训练使用 32 位精度;如需启用 16 位,在 Trainer 中传入precision参数即可:

import lightning as L my_model = MyLightningModule() trainer = L.Trainer(accelerator="tpu", precision="16-true") trainer.fit(my_model)

两种半精度模式:fp16 与 bf16

从当前仓库的源码看,XLA 精度插件支持三种取值:"32-true"(默认)、"16-true"(fp16)与"bf16-true"(bfloat16),定义在 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-true"XLA_USE_F16=1torch.float16启用 FP16
"bf16-true"XLA_USE_BF16=1torch.bfloat16启用 bfloat16

TPU(尤其是 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_BF16XLA_USE_F16环境变量(见src/lightning/fabric/plugins/precision/xla.py),避免污染后续任务。对应的清理行为同样有测试验证(test_teardown)。
  • 优化器步进优化:XLA 精度插件还接管了optimizer_step,在 Fabric 版中使用xm.optimizer_step(optimizer, optimizer_args=kwargs, barrier=True)——设置barrier=True是因为在optimizer.step之后始终执行xm.mark_step()对性能更有利(见src/lightning/fabric/plugins/precision/xla.py#L59-L68);PyTorch 版则包装 closure 并在 step 后调用xm.mark_step()(见src/lightning/pytorch/plugins/precision/xla.py#L62-L85)。

五、TPU 训练在 Lightning 中的底层工作流

理解底层机制有助于排查问题。结合策略与启动器源码,一次 TPU 训练的大致流程如下:

  1. 进程启动_XLALauncher调用torch_xla.distributed.xla_multiprocessing.spawnxmp.spawn)在 N 个核心上启动工作进程,启动方式为fork,并要求入口脚本受if __name__ == "__main__"保护(见 src/lightning/fabric/strategies/launchers/xla.py)。
  2. 数据加载XLAStrategy.process_dataloader将用户 DataLoader 包装为torch_xla.distributed.parallel_loader.MpDeviceLoader,使数据批次在 XLA 设备上并行加载(见src/lightning/pytorch/strategies/xla.py#L179-L190)。
  3. 模型同步setup阶段通过broadcast_master_param将主进程参数广播到所有核心,保持各核心模型初始状态一致(见src/lightning/pytorch/strategies/xla.py#L155-L161)。
  4. 梯度与保存:梯度归约在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_deviceworld_size等属性会抛出RuntimeError,相关属性在_launched为假时返回占位值(见src/lightning/pytorch/strategies/xla.py),这表明分布式信息的读取被严格限定在 spawn 出的工作进程内。

小结

本篇围绕 TPU 训练的三大中级主题给出了可直接落地的配置方案:分布式采样器完全交给 Lightning 自动处理distributed_sampler_kwargs自动注入num_replicasrank);核心数配置为L.Trainer(accelerator="tpu", devices=1|8)(实际上限取决于 TPU 版本,v4 为 4);16 位精度通过precision="16-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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/19 12:59:01

Cocos Creator构建Windows桌面版:从exe到安装包全流程实战指南

最近项目要发布Windows桌面版&#xff0c;产品那边要求给客户一个能双击安装的exe安装包&#xff0c;而不是让用户自己解压文件夹去点运行。我本来以为Cocos Creator构建个Windows平台也就点两下的事&#xff0c;真正走了一遍才发现&#xff0c;从编辑器构建出exe到做出一个合格…

作者头像 李华
网站建设 2026/9/19 12:52:51

Visual Studio 2022 实战安装指南:工作负载、版本选型与故障排查

1. 这不是“点下一步”的安装指南&#xff0c;而是你真正需要的 Visual Studio 2022 实战部署手册Visual Studio 2022 是目前 Windows 平台上最成熟、最完整的集成开发环境&#xff08;IDE&#xff09;&#xff0c;它远不止是一个“写 C# 的工具”。如果你正在为一个新项目搭建…

作者头像 李华