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.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_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 自动构造并注入DistributedSampler,world_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(), 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.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 或 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}" )几个值得注意的要点:
- 合法取值:
'auto'、1、设备总数(如 8),或形如[0]~[7]的单元素列表用于指定使用某一个具体的 TPU 核心。其他数值会抛出ValueError。 - 设备总数并非固定为 8:
auto_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。 - 字符串形式:
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=1 | torch.float16 | 启用 FP16 |
"bf16-true" | XLA_USE_BF16=1 | torch.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_BF16与XLA_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 训练的大致流程如下:
- 进程启动:
_XLALauncher调用torch_xla.distributed.xla_multiprocessing.spawn(xmp.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(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),仅供参考