Google TPU 软件栈深度解析:从 XLA、PJRT 到分布式训练实战

📅 发布时间:2026/9/2 15:56:39
Google TPU 软件栈深度解析:从 XLA、PJRT 到分布式训练实战 最近两年AI 训练资源的争夺已经从“能不能买到 GPU”演变成“如何让每一块加速卡都高效运转”。很多团队盯着 GPU 集群的排队、价格和扩展性却往往忽略了一个长期存在而且正在快速成熟的选项Google TPU。如果只看 TPU 的纸面参数很容易把它简单理解成“一个专用芯片”但真正让 TPU 在大模型训练与 AI 规模化落地中站住脚的是它背后完整的一整套软件栈——从 XLA 编译器、PJRT 运行时到 TPU VM、队列资源调度。这篇文章想讲清楚的核心问题只有一个TPU 软件栈到底是什么样的结构一个普通开发团队如何从零接入以及落地过程中最值得注意的坑在哪里。先说判断TPU 的软件栈已经不是最初那个“TensorFlow 专属”的封闭体系。从公开资料看JAX、PyTorch、TensorFlow 都已经通过 PJRT 这样的统一运行时接入 TPU官方也提供了 TPU VM 一键环境开发者可以像使用 GPU 云节点一样直接开始训练。这传递了一个重要信号TPU 的竞争点已经从单芯片算力转向了生态和软件栈成熟度。所以这篇文章不会停留在硬件介绍层面而是会从 AI 负载的特殊性出发把 TPU VM、XLA、PJRT、分布式策略、JAX 训练示例、PyTorch 迁移方法、常见问题和工程最佳实践一次讲透。如果你正在做大模型预训练、微调或者已经厌倦了 GPU 集群的排队和成本压力这篇文章值得读完。读完你至少能回答三个问题TPU 软件栈包含哪些核心层次如何在 TPU VM 上跑通一个 JAX 训练任务存量 PyTorch 模型迁移到 TPU 会遇到哪些坎。由于很多细节依赖实际环境和版本文中涉及具体版本号的地方我会用保守写法请以官方控制台或命令帮助信息为准。1. AI 时代对芯片的特殊需求决定了 TPU 的定位要理解 TPU 软件栈首先要回到一个更基础的问题AI 时代的芯片到底特殊在哪里传统软件负载的特点是逻辑复杂、分支多、数据访问局部性强对延迟敏感所以 CPU 这种通用指令处理器最合适。而 AI 训练负载的特点是另一套逻辑由大量矩阵乘法和卷积组成计算模式高度重复数据量巨大对单次延迟没有那么敏感但非常依赖吞吐量和内存带宽。更重要的是当模型分布到多卡甚至多机后通信开销会以超线性方式增长。因此AI 芯片设计要同时满足“算得快”“搬得快”“通信快”三个条件这和传统 CPU 的设计哲学完全不同。GPU 凭借大规模并行核心和成熟的通用计算生态成为了当前 AI 训练的主流选择这是事实。但 GPU 本身是为图形渲染和通用并行计算设计的它需要在兼容性和性能之间做权衡。TPU 则是 Google 专门为神经网络设计的 ASIC 芯片从最早的 TPU v1 到后来的 v2、v3、v4再到主打性价比的 v5e、面向大模型训练的 v5p以及最新的 v6e 系列每一代都在强化矩阵运算、降精度计算和芯片间高速互联。这种专用性让 TPU 在特定负载下能获得不错的能效比但代价是软件栈必须足够成熟否则硬件能力无法被上层框架用好。如果把 CPU、GPU、TPU、NPU 放在一张表里对比差异会更直观芯片类型通用性设计目标典型场景软件入口CPU最高通用指令执行操作系统、业务逻辑、IO编译器 / 指令集GPU高并行图形与通用计算渲染、矩阵运算、AI 训练CUDA / ROCm / OpenCLTPU中神经网络矩阵运算大规模 AI 训练与推理XLA PJRT JAX/TF/PyTorchNPU较低AI 推理加速端侧、边缘设备推理厂商 SDK / ONNX Runtime这张表想说明的是它们不是简单的替代关系。实际工程中CPU 负责数据预处理和调度GPU 承接灵活的训练任务NPU 解决端侧推理而 TPU 的价值在于大规模训练和推理场景中的专用性能与规模化调度能力。从“AI 芯片名词盘点”这个角度看TPU 是最能体现“软硬件协同设计”思路的代表之一。对开发团队来说真正值得重估 TPU 的理由有三个第一它能按秒计费、按需创建适合训练任务密集但不想长期养 GPU 集群的团队第二它在 JAX 生态里是一等公民能让 SPMD 编程和自动并行变得更加自然第三Google 在 TPU 上积累了从 TPU Pod 到多切片训练的整套方案这些能力正在逐步开放给云上用户。也就是说TPU 的规模化能力并不只是“把芯片堆起来”而是通过软件栈完成资源调度、编译优化和容错。2. 从 TPU VM 到 XLATPU 软件栈的整体分层理解 TPU 软件栈最好从它的部署形态说起。如果你使用 Google Cloud最直接的入口是创建 TPU VM。TPU VM 是一台同时包含 CPU 和 TPU 芯片的虚拟机你通过 SSH 登录后可以直接在虚拟机里运行框架代码。这和旧式的 TPU Node 架构不同后者把 TPU 放在另一个网络节点中需要通过 TCP 访问维护起来比较麻烦。TPU VM 的模型更接近普通云服务器大大降低了使用门槛。现在的 Cloud TPU 环境基本都基于 TPU VM 形态这也是新手最容易上手的入口。在 TPU VM 里软件栈可以自顶向下分成五个层次第一层是框架层包含 JAX、TensorFlow、PyTorch。过去 TPU 和 TensorFlow 绑定得很深但现在 JAX 已经成为 TPU 上的主流编程接口PyTorch 也可以通过 PyTorch/XLA 接入TensorFlow 则继续通过 TPUStrategy 支持大规模训练。第二层是运行时层。这里最重要的概念是 PJRT全称是 Portable JAX Runtime。它的作用是屏蔽硬件细节让上层框架能以一种统一的方式提交计算图、管理设备、执行任务。JAX、TensorFlow、PyTorch/XLA 都可以基于 PJRT 接入 TPU这是 TPU 生态开放的关键一步。过去有开发者被 XRT 老运行时困扰现在新版本已经逐步向 PJRT 迁移。第三层是编译器层也就是 XLA。XLA 是一个针对线性代数的编译器它会把框架生成的中间计算图编译成适合目标硬件的机器码。在 TPU 上XLA 并不是一个可选优化而是必经路径。框架代码先被转换成 HLO 中间表示XLA 对 HLO 进行算子融合、内存规划、自动并行等优化再生成 TPU 可执行程序。第四层是驱动和内核层包括 libtpu 等底层库负责和 TPU 芯片直接交互。这一层普通开发者很少接触但版本兼容很重要。如果jaxlib或torch_xla版本和 TPU VM 镜像不匹配很容易出现设备找不到或运行时报错。第五层是资源调度层包括 Queued Resources、GKE 等。它不是单个程序的环节而是决定“何时创建 TPU、如何排队、如何扩缩容”的管理层。对于规模化落地来说这一层的重要性甚至超过单芯片性能。这五个层次构成了完整的 TPU 软件栈。写代码时你通常只接触第一层和第三层的边界但真正理解整个栈排错和调优时才能有的放矢。尤其是当你从 GPU 迁移到 TPU最大的变化并不是换一个device名称而是“程序如何被编译和执行”的逻辑变了。3. 核心构建块XLA 编译与 PJRT 运行时如果说软件栈是一栋房子XLA 和 PJRT 就是承重墙。先看 XLA。XLA 解决的核心问题是算子融合和内存复用。在 GPU 上运行一个“矩阵乘 激活函数 损失函数”的链路通常会被拆成多个 kernel 逐个执行每个 kernel 都要从显存读取中间结果、算完再写回这是很大的带宽浪费。XLA 会把整条链路融合成一个计算图尽量让中间结果在片上传递减少 HBM 的读写。这对 TPU 特别重要因为 TPU 的设计逻辑是“数据搬运越少越好”。XLA 的好处很直观但代价是编译时间。第一次运行一段 JAX 代码时XLA 需要把计算图编译成设备可执行文件这个过程可能比实际训练一个 batch 还要久。编译结果会被缓存所以第二次跑相同 shape 的任务会快很多。真正容易踩坑的是动态 shape如果你的输入尺寸在训练过程中变化XLA 缓存会失效程序可能反复触发重新编译表现为“偶发性卡顿”或“第一次非常慢、后面时快时慢”。再看 PJRT。PJRT 的定位类似硬件抽象层它对上提供统一的设备发现、计算提交、执行、传输接口对下对接不同的加速硬件。对框架开发者来说只要实现一次 PJRT 插件就能让一套框架跑在多种硬件上。对普通开发者来说PJRT 带来的直观好处是同一段 JAX 代码在 CPU、GPU、TPU 上往往只需要切换设备或环境变量就能运行。下面这个示例可以在刚创建的 TPU VM 上执行用来验证环境是否正常import jax print(jax.default_backend()) print(jax.devices())在 TPU VM 上预期大概率会打印出tpu和一串TpuDevice列表在 GPU 机上则会看到gpu和GpuDevice。这个简单的输出能快速确认驱动、运行时和设备之间的链路是否接通是排错的第一步。从工程角度理解 XLA 和 PJRT我的建议是不要把它们当成黑盒。当你遇到“编译慢”“动态 shape 触发重编译”“算子不支持”时首先要想到这些现象来自 XLA 编译层而不是框架 bug。多看看 XLA 日志中的 HLO 图形往往比盲目改代码更有效。4. TPU 环境准备与基础配置在正式接触 TPU 软件栈之前先判断一下它适不适合你当前的阶段。如果你满足下面任一条件TPU 值得认真尝试团队正在做大规模预训练或微调GPU 集群排队严重想深入使用 JAX 生态已经使用 Google Cloud希望获得更高能效比的训练资源准备把训练任务用 Kubernetes 编排希望把 TPU 纳入统一调度。如果你的项目还处于本地原型验证阶段没有海外云资源或预算那么先在本机用 CPU/GPU 跑通 JAX 代码之后再迁移到 TPU VM是更务实的路径。创建 TPU VM 的通用步骤如下。假设你已经完成了 Google Cloud 账号开通、项目创建和计费启用并在项目中启用了 TPU API可以通过gcloud创建。这里要注意不同区域的可用产品不同accelerator-type和version也会随官方维护周期变化。下面的命令属于通用模板实际值请以gcloud帮助信息或控制台为准。gcloud compute tpus tpu-vm create my-tpu \ --zoneus-central2-b \ --accelerator-typev4-8 \ --versiontpu-vm-jax-0.4.35-pod-pjrt命令中的my-tpu是实例名称--zone是区域--accelerator-type决定 TPU 的规模和形态--version决定 TPU VM 镜像。这里要特别提醒不要直接抄袭网上的历史版本号因为 TPU 镜像更新很快旧版本可能已经被下线或出现兼容问题。创建前建议执行gcloud compute tpus tpu-vm create --help或者在控制台查看当前可选版本。创建完成后通过 SSH 登录gcloud compute tpus tpu-vm ssh my-tpu --zoneus-central2-b登录后首先检查 JAX 能否看到 TPU 设备python3 -c import jax; print(jax.devices())如果返回结果包含多个TpuDevice说明软件栈基本可用。此时再安装项目额外依赖比如pip install jax[cuda]这种只在 GPU 环境需要的包就不必再装了TPU 镜像通常已经预装了匹配的jax和jaxlib。TPU VM 还有一个和普通云服务器一样的特性它不是永久环境。当你删除并重新创建 TPU VM或者切换镜像后之前通过pip install装的包会丢失。因此比较推荐的做法是写一个初始化脚本例如把安装依赖的命令放到 shell 脚本中创建 TPU VM 时通过启动脚本自动执行。这样每次重建环境都能保持一致避免出现“这台 TPU 能跑另一台跑不了”的环境漂移问题。5. 完整示例在 TPU 上跑通 JAX 训练 MNIST环境就绪后用一个最小可运行的例子验证全链路。JAX 语法贴近 NumPy自动微分和 XLA 编译都是原生能力是当前 TPU 上最顺手的编程框架。下面这个示例用 JAX 实现一个两层全连接网络在 MNIST 上训练 3 个 epoch。文件可以命名为mnist_jax_tpu.py直接在 TPU VM 上运行即可import jax import jax.numpy as jnp from jax import random, grad, jit import numpy as np from sklearn.datasets import fetch_openml from sklearn.model_selection import train_test_split def load_mnist(): X, y fetch_openml(mnist_784, version1, return_X_yTrue, as_frameFalse) X jnp.asarray(X, dtypejnp.float32) / 255.0 y jnp.asarray(y.astype(np.int32)) y jax.nn.one_hot(y, 10) return train_test_split(X, y, test_size0.2, random_state42) def init_params(rng, layer_sizes): params [] for i in range(len(layer_sizes) - 1): rng, subkey random.split(rng) w random.normal(subkey, (layer_sizes[i], layer_sizes[i 1])) * 0.1 b jnp.zeros((layer_sizes[i 1],)) params.append((w, b)) return params def forward(params, x): h x for i, (w, b) in enumerate(params): h jnp.dot(h, w) b if i len(params) - 1: h jax.nn.relu(h) return h def loss_fn(params, x, y): logits forward(params, x) return jnp.mean(jax.nn.softmax_cross_entropy(logits, y)) jit def train_step(params, x, y, lr0.01): grads grad(loss_fn)(params, x, y) return [(w - lr * dw, b - lr * db) for (w, b), (dw, db) in zip(params, grads)] def evaluate(params, X_test, y_test): pred jnp.argmax(forward(params, X_test), axis1) return jnp.mean(pred jnp.argmax(y_test, axis1)) if __name__ __main__: X_train, X_test, y_train, y_test load_mnist() params init_params(random.PRNGKey(0), [784, 256, 10]) batch_size 256 for epoch in range(3): for i in range(0, len(X_train), batch_size): xb X_train[i : i batch_size] yb y_train[i : i batch_size] params train_step(params, xb, yb) acc evaluate(params, X_test, y_test) print(fepoch {epoch 1}, test acc: {acc:.4f})这段代码的逻辑很直白init_params初始化权重forward做前向计算loss_fn计算交叉熵损失train_step通过grad拿到梯度并更新参数jit让 XLA 把训练步骤编译成高效执行程序evaluate在测试集上算准确率。运行命令python3 mnist_jax_tpu.py你会在终端看到类似下面的输出epoch 1, test acc: 0.91xx epoch 2, test acc: 0.92xx epoch 3, test acc: 0.93xx这里的准确率数值只和一个“可运行的模型”有关不值得记住。跑通这个 example 的真正意义在于框架、XLA 编译、TPU 设备、数据加载、训练循环这一整条链路已经打通。在进入大规模任务之前这条链路是否稳定决定了后面所有工作能不能顺利开展。从这个示例也能看出一个观念问题TPU 不是用来跑 MNIST 的。这类小任务在 CPU 上几秒就能结束在 TPU 上还要先花时间编译体验反而“更慢”。TPU 的优势要到模型规模变大、计算密集度变高、数据吞吐变大的时候才会显现。所以不要用“跑 MNIST 快不快”来判断 TPU 的能力而要看它在分布式大任务中的稳定性和扩展性。6. 存量模型怎么迁移PyTorch / TensorFlow 接入 TPU很多团队手里已经用 PyTorch 写好了模型最关心的问题自然是不改代码直接上 TPU 行不行答案是不可能完全不改但改动量通常比想象中要小。PyTorch/XLA 项目把 PyTorch 的计算图转换为 XLA 可执行的图从而在 TPU 上运行。迁移的核心是把原来“模型放 GPU、优化器在 GPU 上执行”的方式改成“模型放 XLA 设备、优化器通过 XLA 同步执行”。一个最小迁移示例是这样的。假设你原来有一个 PyTorch 训练循环import torch import torch.nn.functional as F import torch_xla import torch_xla.core.xla_model as xm # 原来是 device torch.device(cuda) device xm.xla_device() model Net().to(device) optimizer torch.optim.SGD(model.parameters(), lr0.01) for data, target in train_loader: data, target data.to(device), target.to(device) optimizer.zero_grad() output model(data) loss F.nll_loss(output, target) loss.backward() # 原来是 optimizer.step() xm.optimizer_step(optimizer)这里最关键的两行改动一个是xm.xla_device()另一个是xm.optimizer_step(optimizer)。在 PyTorch/XLA 的早期实现中整个计算过程是 Lazy 执行的符号张量要等到真正需要值时才会触发编译新版本支持了更接近 PyTorch 原生体验的执行模式但多核场景下的optimizer.step()同步问题依然要交给xm.optimizer_step处理。TensorFlow 的迁移路径更特殊。因为 XLA 本来就是从 TensorFlow 生态里成长起来的所以 TensorFlow 对 TPU 的支持非常原生。如果你习惯 Keras可以直接用TPUStrategyimport tensorflow as tf resolver tf.distribute.cluster_resolver.TPUClusterResolver() tf.config.experimental_connect_to_cluster(resolver) tf.tpu.experimental.initialize_tpu_system(resolver) strategy tf.distribute.TPUStrategy(resolver) with strategy.scope(): model tf.keras.Sequential([...]) model.compile(optimizeradam, losssparse_categorical_crossentropy, metrics[accuracy]) model.fit(train_dataset, epochs3)这段代码里TPUClusterResolver会自动发现当前 TPU VM 上的设备TPUStrategy负责把训练过程分发到所有 TPU 核心上。不过迁移的真实难点并不在 API 调用而在于模型的执行模式。PyTorch 中灵活的 Python 控制流、动态 shape、运行时 if/else在 TPU 上都会成为 XLA 编译的阻碍。常见的情况是代码在 GPU 上跑得很好一迁移到 TPU 就报XlaRuntimeError报错信息里包含INVALID_ARGUMENT或 shape mismatch。这类问题往往就是因为计算图里有不确定的 shape 分支。所以做迁移时记住一个原则让计算图尽可能静态化。固定 batch size、固定输入图片尺寸、把条件分支改写成jax.lax.cond或torch.where这类符号化控制流能避开绝大多数编译问题。7. 分布式训练与规模化落地的关键路径到这一步你已经跑通了单 TPU VM 上的训练任务。接下来要回答的问题是如何把 TPU 的规模化能力用起来TPU 的规模化并不只是“多插几块卡”那么简单。第一层是单 VM 多核心的并行例如 v4-8 里包含多个 TPU 芯片最自然的用法是数据并行把同一个 batch 分到不同核心上梯度聚合并更新。第二层是 TPU Pod成百上千个 TPU 芯片通过高带宽互联组成一个超级计算池适合训练非常大的模型。第三层是多切片Multislice训练它允许你跨多个 Pod 运行训练任务进一步突破单池的资源上限。每一层跨越软件栈的复杂度都在上升。框架侧对应的支持也在成熟。JAX 提倡 SPMD 编程模型用jax.sharding把张量分布到设备上TensorFlow 通过TPUStrategy做数据与模型并行PyTorch 生态则主要是 FSDP 结合 PyTorch/XLA让模型参数与梯度分片存储从而在内存有限的情况下训练超大模型。从工程角度看FSDP 是 PyTorch 用户从单机多卡迁移到 TPU 时最需要提前了解的概念因为它决定了模型能否在有限内存下训练。资源层面的调度也不可忽视。Queued Resources 允许你提前申请一段时间的 TPU 资源按需排队、按量计费适合确定性的训练计划。如果你的团队已经使用 KubernetesGKE 可以把 TPU 节点纳入集群调度让训练任务和推理任务共用一套编排体系。不过要注意TPU 节点的调度粒度、启动速度和 GPU 节点不完全相同实际生产中需要对排障和滚动升级方案做针对性设计。有一个工程上的强烈建议不要一开始就创建一个巨大的 TPU Pod。很多团队第一次用 TPU 就急着申请大集群结果环境变量、网络配置、数据分片任何一个环节出错都会造成大量资源浪费。更稳妥的路径是先在 v4-8 或 v5e 小规格上跑通训练再逐步扩大 accelerator-type最后再进入 Pod 或多切片模式。每一步只改变一个变量排错成本会低很多。8. 常见问题与排查思路TPU 环境的问题往往集中出现在几个固定环节编译阶段、数据加载阶段、分布式初始化阶段。下面这张表是实际项目中最常见的排查路径。问题现象可能原因排查方式解决方案第一次训练非常慢XLA 编译开销观察日志中“Compilation”阶段固定 shape、增大 batch、预热编译缓存运行时报XlaRuntimeError: INVALID_ARGUMENT动态 shape、算子不支持查看报错堆栈中的 shape 信息固定输入尺寸、避免 Python 运行时 if/else训练时报 OOMbatch 过大或编译内存开销高查看 TPU 内存日志逐步缩小 batch梯度累积、FSDP 参数分片数据加载成为瓶颈CPU 端预处理跟不上用 profiling 工具观察 DataLoading 耗时使用 TFDS 或预取把数据提前读入 host 内存多机训练时设备数量不对启动方式缺少 multi-host 参数检查每台主机的jax.devices()输出按官方 multi-host 启动方式配置 coordinator删除并重建 TPU VM 后环境丢失TPU VM 不是持久环境检查启动日志中的 pip 安装记录使用启动脚本或自定义镜像jax.devices()为空或报驱动错误jaxlib与镜像版本不匹配查看jax.__version__与 TPU 镜像版本重建镜像或升级/降级jaxlib遇到问题时第一步永远不是改代码而是先确认基础环境。比如先运行jax.devices()就是最简单的链路验证。第二步是打开 XLA 日志让编译过程可见。JAX 中可以通过环境变量开启 loggingTensorFlow 和 PyTorch/XLA 也有各自的 verbose 选项。看见问题才能定位问题。还要注意TPU 的报错信息很多时候并不直接指向真正的原因。OOM可能来自数据分片策略不当也可能来自编译过程中的内存峰值INVALID_ARGUMENT可能是动态 shape也可能是某个算子没有 TPU 实现。排查时要有耐心把“框架报错”和“底层原因”分开分析。9. 最佳实践与工程建议基于对 TPU 软件栈的理解下面这些建议能帮你少走弯路。第一把 TPU VM 当成无状态计算节点。不要在 TPU VM 上保存重要文件数据与 Checkpoint 应该放到持久化存储中。因为 TPU VM 删掉重建后本地文件会全部丢失而 AI 训练里的 Checkpoint 是真正的资产必须放在独立的存储服务里。第二固定镜像和依赖版本。TPU 软件栈迭代非常快XLA 的编译行为、JAX 的 API、PyTorch/XLA 的接口都在变化。如果你不固定版本可能上个月能跑的代码这个月创建的新镜像就编译出错。记录项目使用的jax、jaxlib、torch_xla、tensorflow版本和 TPU 镜像版本一起写入项目文档是有价值的习惯。第三坚持“先小后大”的扩缩容路线。从单 VM 到 Pod再到 Multislice每一步都要做完整的功能验证。尤其要注意数据加载和分布式初始化这两个环节它们最容易在扩规模时暴露问题。比如小规格 TPU 上数据加载用的是本机内存扩大规模后如果仍从单点拉数据网络 IO 就会成为瓶颈。第四先 profiling再优化。TPU 上的大多数性能问题不是“计算不够快”而是“编译等待太久”或“数据供给不上”。不要凭感觉调 batch size先看 profiling 数据。TPU 官方生态提供了 TensorBoard 插件和 profiling 工具能比较清楚地看到编译时间、算子耗时、数据加载耗时等指标。基于数据决策比拍脑袋靠谱得多。第五安全与权限要按最小权限原则设置。TPU VM 是可以 SSH 登录的计算资源如果 IAM 权限过大相当于把训练集群的入口暴露给不必要的人。建议按角色分配权限限制 SSH 来源打开审计日志对高危操作启用审批。不要在生产环境使用共享的管理员账号。第六成本控制要前置。TPU 按秒计费听起来很灵活但如果你创建了 Pod 规模资源后忘记释放成本同样很高。建议为 TPU 设置配额