news 2026/9/2 18:26:23

TPU软件栈实战:从JAX环境搭建到多卡训练与性能优化

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
TPU软件栈实战:从JAX环境搭建到多卡训练与性能优化

实际 AI 项目中,GPU 是默认选择,但 Google TPU 软件栈在大模型训练和推理规模化场景中正越来越常见。所谓 TPU 软件栈,不只是一组驱动或一个 Python 库,而是从高层框架、编译器到芯片运行时的一整套分层体系。只有理解这条链路,才能解释为什么同一段 JAX 代码在 GPU 上跑通后,放到 TPU 上会遇到设备未识别、编译时间过长、内存不足等问题;也只有理解这条链路,才能让训练任务真正利用多张 TPU,而不是把几十片加速器当成摆设。

这篇文章从一个可复现的最小训练任务入手,先搭建 Cloud TPU VM 环境,再解释 XLA 编译、JAX 设备抽象、TPU 内存模型和 SPMD 多卡编程,最后给出性能排查路径和生产落地建议。适合三类读者:刚接触 TPU 但已经熟悉 JAX 或 TensorFlow 的同学;从 GPU 训练迁移到 TPU 的算法工程师;以及需要维护 TPU 训练平台、排查训练故障的平台团队。

1. 先理解 TPU 软件栈在 AI 规模化投入中的位置

1.1 TPU 硬件形态决定了软件栈不是可选项

TPU 是 Google 设计的专用 AI 加速器芯片。它和 GPU 的一个关键差别是,GPU 可以直接运行 CUDA/ROCm 等通用并行计算内核,而 TPU 是一个 ASIC,计算单元、内存布局和指令集都是为神经网络算子高度定制的。普通 C++ 或 Python 代码不能直接“落到”TPU 上执行,必须经过编译器把它翻译成 TPU 的指令。

从硬件连接方式看,TPU 并不独立存在。Cloud TPU 的常见形态是一个 TPU VM:一个 x86 宿主机通过高速互联挂载 4 片或 8 片 TPU 芯片,片上还带脉动阵列、向量单元、标量单元和片内高速互联。宿主机的职责是运行 Python 进程、加载数据、调度函数,而真正的矩阵运算发生在 TPU 侧。这两个处理器之间的通信,必须依赖运行时层去管理。

因此,TPU 软件栈不是“装个驱动就能用”的简单组件。它承担了至少四件事:

  • 把高层模型代码转换为计算图。
  • 把计算图编译成 TPU 可执行的指令序列。
  • 管理 TPU 内存分配和 host 与 TPU 之间的数据复制。
  • 在多卡或多机训练时协调通信和梯度同步。

这四件事分别落在框架层、编译器层和运行时层,任何一层不合格,最终训练任务都会失败或者性能很弱。

1.2 软件栈的分层:框架层、编译器层、运行时层

为了排查问题时不迷路,可以把 TPU 软件栈按层拆开。

层级典型组件职责常见问题
框架层JAX、TensorFlow、PyTorch/XLA定义模型、优化器、训练循环API 版本不匹配、数据形状不一致
编译器层XLA、HLO、XLA Runtime把算子编译成 TPU 指令并做融合优化动态 shape 导致反复重编译、编译时间过长
运行时层PJRT、libtpu、TPU 驱动设备枚举、内存分配、主机与 TPU 通信找不到 libtpu.so、设备列表为空、OOM
硬件层TPU v4/v5e/v5p 等芯片执行矩阵乘法和激活函数超出约束、芯片故障、互联异常

框架层是用户接触最多的部分。JAX 和 TensorFlow 都会把 Python 函数转换成中间表示。TensorFlow 早期依赖 GraphDef,JAX 则通过jax.jit把函数追踪成 XLA HLO。

编译器层是 TPU 软件栈最核心的部分。XLA 会做算子融合、内存分配、指令调度和张量布局选择。训练任务最终是否高效,很大程度上取决于 XLA 能否消除冗余拷贝、把多个小算子合并成一个高效内核。

运行时层是容易忽略的部分。PJRT 提供了设备抽象,让 JAX 可以接到不同后端。TPU 后端通过 libtpu 与真实硬件通信。如果安装 JAX 时没有把 TPU 插件装好,jax.devices()就只会返回 CPU 设备。

1.3 为什么 JAX 是上手 TPU 的首选

TensorFlow 也能跑 TPU,PyTorch 也有 PyTorch/XLA 分支,但 JAX 有三个天然优势:

第一,JAX 以函数式编程为核心。模型参数和梯度是普通数据结构,训练循环可以写成纯函数,这让设备上的数据放置、编译和并行都更容易描述。

第二,JAX 与 XLA 绑定最深。jax.jit本质上就是“追踪函数,生成 XLA 计算,再执行编译产物”。从 JAX 到 TPU 的路径最短,出问题时最容易被理解。

第三,JAX 在 Google 内部和开源大模型生态中被大量使用。很多 TPU 上的开源模型示例、性能调优工具和大规模分布式训练方案都是基于 JAX 写的,后续参考资料多。

下面的内容都以 JAX 为主。PyTorch/XLA 和 TensorFlow 的差别会在最后一节单独说明。

2. 搭建一套可复现的 TPU 开发环境

2.1 创建 TPU VM 前需要确认的规格

TPU 开发不建议直接在本地模拟。虽然 JAX 可以运行在 CPU 上,也能用XLA_FLAGS=--xla_force_host_platform_device_count=8模拟多设备逻辑,但真实 TPU 的编译行为、内存模型和分布式通信无法完全模拟。更稳妥的做法是创建一个 Cloud TPU VM。

创建前需要确认四件事:

  • 区域和可用配额。
  • TPU 芯片类型:v4、v5e、v5p 等。
  • 加速器形态:例如 v5e-4、v5e-8、v4-8,末尾数字表示训练 Pod 中可见的芯片数量。
  • TPU VM 运行环境镜像版本。

使用 gcloud 创建的基本命令如下:

gcloud compute tpus tpu-vm create tpu-demo \ --zone=us-central1-b \ --accelerator-type=v5e-4 \ --version=tpu-vm-base \ --project=your-project-id

版本字段tpu-vm-base只是一个示例。不同区域、不同芯片型号支持的镜像版本可能不同,创建前先执行下面的命令确认当前可选项,下面示例用于说明思路,实际命令以后续官方文档为准。

gcloud compute tpus tpu-vm versions list \ --zone=us-central1-b \ --project=your-project-id

创建完成后,通过 SSH 登录到 TPU VM。

gcloud compute tpus tpu-vm ssh tpu-demo \ --zone=us-central1-b \ --project=your-project-id

进入机器后先确认设备和系统信息。lsblk看数据盘,nvidia-smi在这里不适用,TPU VM 可以不用关心 GPU 驱动。

2.2 在 TPU VM 上安装 JAX、TensorFlow 与 PyTorch/XLA

TPU VM 通常预置了 Python 环境和一些依赖,但建议在虚拟环境里重新安装,避免系统包互相污染。

JAX 的 TPU 版本安装命令如下:

python -m venv ~/venv source ~/venv/bin/activate pip install -U pip setuptools wheel pip install -U "jax[tpu]" -f https://storage.googleapis.com/jax-releases/libtpu_releases.html

这条命令会安装 JAX、XLA 相关 Python 包以及 libtpu 插件。libtpu是 JAX 连接 TPU 硬件的关键组件,没有它,jax.devices()看不到 TPU 设备。

如果项目里还需要 TensorFlow,可以按需安装。需要注意的是,JAX 和 TensorFlow 都依赖absl-pyprotobuf等公共库,直接装最新版可能覆盖对方依赖。建议先创建独立虚拟环境,再按项目锁版本安装。

pip install -U tensorflow-cpu

PyTorch/XLA 的安装命令会随版本变化,一般需要通过 PyTorch 官方指示安装。这里给出一个通用参考:

pip install torch torchvision pip install torch_xla

PyTorch/XLA 本身也是一个 XLA 编译器的前端,它把 PyTorch 计算图编译到 XLA,再执行到 TPU。安装后可以通过import torch_xla.core.xla_model as xm访问 XLA 设备。

2.3 验证环境:设备可见性最快检查路径

环境是否可用,不需要先跑一个完整模型。把下面这段脚本放到check_tpu.py中执行即可。

import jax devices = jax.devices() print("devices:", devices) print("device count:", len(devices)) for i, d in enumerate(devices): print("device", i, "=", d, "kind:", getattr(d, "platform", "unknown"))

预期输出是一组TpuDevice,例如:

devices: [TpuDevice(id=0, process_index=0, coords=(0,0,0), core_on_chip=0), TpuDevice(...)] device count: 4

如果输出只包含CpuDevice,说明 JAX 没有识别到 TPU。从下面几个方向检查:

检查项命令或位置判断标准
环境变量echo $TPU_NAME$TPU_DRIVER_MODETPU_DRIVER_MODE=1时通常使用 v5e 的 PjRt 路径
插件文件python -c "import jax.tools.collect_profile"不应出现导包错误
libtpupip show libtpu应能看到已安装版本
JAX 后端python -c "import jax; print(jax.default_backend())"应为tputpu相关插件

注意:不要只凭“程序能运行”判断环境正确。先执行jax.devices()确认设备列表,再跑真实训练。否则后续所有性能问题都会被错误环境掩盖。

3. 用一个最小训练任务跑通 TPU 软件栈

3.1 任务设计:两层 MLP 识别 MNIST

环境确认后,跑一个足够小但完整的训练任务。这里选择两层 MLP 识别 MNIST,而不是直接上 Transformer。

选择这个任务有三个原因:

  • 数据下载简单,通过 Keras 自带数据集即可拿到。
  • 网络结构简单,参数可以通过普通 Python 列表维护。
  • 训练规模小,能在一台 4 卡 TPU VM 上快速跑完,适合验证软件栈是否正常。

完整代码如下,核心逻辑分为三个部分:初始化参数、前向预测、训练步骤。

import numpy as np import jax import jax.numpy as jnp from jax import random, jit, value_and_grad def init_params(rng, layer_sizes): params = [] keys = random.split(rng, len(layer_sizes) - 1) for key, in_size, out_size in zip(keys, layer_sizes[:-1], layer_sizes[1:]): w = random.normal(key, (in_size, out_size)) * jnp.sqrt(2.0 / in_size) b = jnp.zeros(out_size) params.append((w, b)) return params def predict(params, x): h = x for w, b in params[:-1]: h = jax.nn.relu(jnp.dot(h, w) + b) w, b = params[-1] return jnp.dot(h, w) + b def loss_fn(params, x, y): logits = predict(params, x) return jnp.mean(jnp.sum((logits - y) ** 2, axis=-1)) @jit def train_step(params, x, y, learning_rate=0.01): loss, grads = value_and_grad(loss_fn)(params, x, y) params = jax.tree_util.tree_map( lambda p, g: p - learning_rate * g, params, grads ) return params, loss

这个示例没有使用任何复杂包装,目的就是让训练流程可以直接被读懂。train_step@jit装饰,因此第一次调用会被编译,后面相同 shape 的调用会复用编译结果。

3.2 用 JAX 数据加载和训练主循环

MNIST 数据使用 Keras 加载,再转换成 JAX 数组。TPU 训练对数据加载路径要求不高,但要注意 batch 大小固定,最好使用numpy数组预取。

import tensorflow as tf (x_train, y_train), _ = tf.keras.datasets.mnist.load_data() x_train = x_train.reshape(-1, 784).astype(np.float32) / 255.0 y_train = jax.nn.one_hot(np.array(y_train), 10) batch_size = 128 num_batches = x_train.shape[0] // batch_size train_data = [] for i in range(num_batches): xb = np.asarray(x_train[i * batch_size : (i + 1) * batch_size]) yb = np.asarray(y_train[i * batch_size : (i + 1) * batch_size]) train_data.append((xb, yb)) rng = random.key(0) params = init_params(rng, [784, 256, 10]) for epoch in range(5): total_loss = 0.0 for xb, yb in train_data: xb = jnp.asarray(xb) yb = jnp.asarray(yb) params, loss = train_step(params, xb, yb) total_loss += float(loss) print(f"epoch {epoch}, avg loss = {total_loss / num_batches:.4f}")

每个 batch 都执行一次jnp.asarray,把 NumPy 数据拷贝到 JAX 可访问的内存中。训练主循环保持在 Python 层,每次调用train_step进入已经编译好的 XLA 可执行对象,这是 JAX 最小可运行闭环的推荐写法。

如果希望把多个 step 放进同一个编译单元,可以把整个 epoch 循环放入jax.lax.scan,这样能减少 Python 调用开销,但对这个演示任务不是必须的。

3.3 运行观察点:编译日志、设备数量和 step 耗时

运行脚本时重点观察三个现象:

第一,第一次执行train_step前会有几秒的编译阶段。日志中可能出现 JAX/XLA 的编译输出,这不是卡死,而是 XLA 正在把函数编译成 TPU 指令。

第二,epoch 0的耗时通常明显高于后面几个 epoch。原因是编译结果被缓存,后续相同 shape 的调用直接执行。

第三,任务结束后用jax.devices()再确认一次设备数量。如果本应使用 4 卡,却只有 1 个设备,说明 JAX 运行在多进程模式下没有正确加载多个 TPU 芯片,需要检查进程数或TPU_NAME配置。

一个正常运行的输出示例:

epoch 0, avg loss = 2.2861 epoch 1, avg loss = 1.3872 epoch 2, avg loss = 0.9258 epoch 3, avg loss = 0.6682 epoch 4, avg loss = 0.5137

如果把print加在train_step外面,第一次 step 打印间隔会比后续长很多,这是 JIT 编译的正常表现。

4. 把软件栈拆开看:XLA 编译与 TPU 内存模型

4.1 XLA 如何把高层算子变成 TPU 可执行的 HLO

JAX 中看到的矩阵乘法jnp.dot并不是在 TPU 上逐行执行,而是先被追踪成 HLO 计算图。HLO 是 XLA 的中间表示,类似 PyTorch 的 TorchScript 和 TensorFlow 的 GraphDef。

XLA 编译器在拿到 HLO 后,会执行一系列优化:

  • 算子融合:把dot + add + relu融合成一个内核,减少读内存次数。
  • 布局选择:决定矩阵按行主序还是列主序存储,TPU 对二维布局有强偏好。
  • 内存复用:在训练中复用临时缓冲,降低 HBM 峰值占用。
  • 指令调度:把可以并行的计算放到不同执行单元上。

这些优化对用户透明,但用户能通过环境变量导出优化前后的 HLO,帮助定位性能问题。

export XLA_FLAGS="--xla_dump_to=/tmp/hlo"

执行训练脚本后,/tmp/hlo中会生成.txt.dot文件。看到fusion字样的大量出现,说明 XLA 已经做了算子融合;如果看到大量独立的小 op,说明网络结构可能不适合 TPU,存在较多无法融合的边界操作。

4.2 TPU 内存体系:HBM、VMEM 与 SMEM

GPU 开发者习惯用显存大小估算可训练模型规模,但 TPU 的内存模型不太一样。TPU 至少涉及三类内存。

内存类型全称访问者特点
HBMHigh Bandwidth Memoryhost、计算核心容量最大,类似 GPU 显存,但带宽比片内内存低
VMEMVector Memory向量单元、脉动阵列容量小、带宽极高,用于激活、中间结果
SMEMScalar Memory标量单元容量最小,用于循环计数、标量参数等

实际训练中,绝大多数中间矩阵都希望放到 VMEM 里。如果一次性计算太大,XLA 编译器会把中间结果溢写到 HBM。溢写越多,性能越差。这也是为什么 TPU 上经常出现“显存占用看着不大,但编译后执行很慢”的情况。

常见的内存相关报错是OUT_OF_MEMORY。它不一定表示 HBM 用满,也可能是 XLA 在编译阶段无法为某个中间张量分配连续 VMEM 区域。处理思路不是盲目调小 batch,而是先查看编译生成的 HLO,确认是哪个算子造成了峰值内存。

4.3 bfloat16 与静态形状为什么是 TPU 调优的杠杆

TPU 对bfloat16支持非常高效。bfloat16只有 16 位,但保留与float32相同的指数位范围,只是减少了尾数精度,适合大多数神经网络训练。使用bfloat16之后,矩阵乘法所需的内存带宽约降低一半,更多中间结果能留在 VMEM 中,训练吞吐提升明显。

在 JAX 中,可以通过dtype控制:

x_bf16 = x.astype(jnp.bfloat16) w_bf16 = w.astype(jnp.bfloat16)

不过不能简单把所有参数都改成bfloat16。batch norm、聚合统计量、loss 计算通常建议保持float32,否则可能出现精度问题。常见做法是:模型权重和矩阵乘法主路径用bfloat16,优化器状态和 loss 用float32,并通过jax.lax.convert_element_type或混合精度策略管理。

另一个重要杠杆是静态形状。JAX 的@jit会在首次调用时记录输入 shape。如果后续调用传入不同 shape,XLA 会丢弃缓存并重新编译。动态 shape 是 TPU 训练中常见的性能杀手,例如对变长文本做动态 padding、在循环体内改变矩阵维度,都会导致反复重编译。

注意:在 TPU 上,一个看似无害的arr.shape[0]也可能导致 Python 跟踪逻辑产生不同分支。训练输入应尽量使用固定 shape,或者显式 pad 到固定长度。

5. 从单卡到多卡:SPMD、PJRT 与分布式训练

5.1 数据在哪:device_put、device_get 与隐式复制

单卡任务跑通后,下一个阶段是让多个 TPU 芯片同时工作。JAX 中的数组总是存在于某个设备上,调试时要先问“这个数组在哪个设备”。jax.device_put可以指定放置设备,jax.device_get把数组取回 host。

import jax devices = jax.devices() print("device count:", len(devices)) x = jnp.ones((8, 8)) x_on_device = jax.device_put(x, devices[0]) print(x_on_device.device) x_back = jax.device_get(x_on_device) print(type(x_back))

在多卡训练中,一个最常见的错误是在 Python 循环内部频繁执行device_getnp.asarray,这会强制把数据从 TPU 复制回 CPU,造成 host 和 device 之间双向阻塞。正确做法是让复制发生在 batch 边界,训练 step 内部只操作 JAX 数组。

5.2 用 jax.sharding 把参数和梯度切到多张 TPU 上

JAX 提供了jax.sharding体系,用来描述数据如何分布到多个设备。最常用的是NamedSharding,通过MeshPartitionSpec定义分片方式。

from jax.sharding import Mesh, PartitionSpec, NamedSharding devices = jax.devices() mesh = Mesh(devices, ("data",)) sharding = NamedSharding(mesh, PartitionSpec("data"))

当数组被放到这个 sharding 下时,JAX 会按“数据维度”把数组切分到所有设备上。训练循环中把输入数据用对应 sharding 放置,XLA 会自动插入必要的 all-reduce 通信来完成梯度同步。

示例:

@jit def train_step_sharded(params, x, y): loss, grads = value_and_grad(loss_fn)(params, x, y) # 梯度也是分片状态,通信由 XLA 根据 sharding 自动生成 grads = jax.lax.pmean(grads, axis_name="data") params = jax.tree_util.tree_map( lambda p, g: p - 0.01 * g, params, grads ) return params, loss

jax.lax.pmean表示跨data轴对梯度取均值,所有设备执行完成后得到一致梯度,这是一般数据并行训练的核心通信原语。XLA 会根据设备拓扑决定 all-reduce 的执行方式,不需要用户手写通信算子。

这个示例用于说明思路,实际项目需要结合自研模型结构调整分片维度。比如参数矩阵并行时,需要把权重按列的维度切到不同设备,并引入pallpsum等原语。

5.3 多机训练与 PJRT 运行时需要注意什么

当训练规模超过单 VM 的 TPU 芯片数量时,需要把任务扩展到多台 TPU VM。JAX 在多机下的运行时由 PJRT 管理。PJRT 是一种可移植运行时抽象,JAX 通过它统一访问 CPU、GPU 和 TPU 后端。

多机训练需要关注三件事:

第一,任务要显式指定机器的 coordinator 地址和 rank。JAX 常见做法是通过启动脚本设置环境变量,例如JAX_COORDINATOR_ADDRESSJAX_COORDINATOR_PORTJAX_EXPECTED_DEVICES_PER_PROCESS。不同版本变量名可能不同,落地前要查阅当前版本文档。

第二,进程数和设备数要对齐。一台 TPU VM 上启动多个 Python 进程时,每个进程默认只能看到一部分设备。如果发现设备总数不对,优先检查进程数是否为 VM 芯片数量的整数倍。

第三,数据集要按 global batch 切分。多卡训练时不能每张卡读取同一份 batch,否则等于把同一个 batch 复制了 N 次。建议使用tf.datatorch.utils.data.DistributedSamplerprocess_indexprocess_count切分,保证每张卡收到的数据不重叠。

6. 性能排查:现象、根因与处理路径

6.1 训练慢但 CPU 不忙,先看 host-device 复制和编译缓存

现象:训练 step 耗时很长,CPU 占用率不高,TPU 也没有跑满。

优先检查:

  • 是否在训练循环内执行np.asarray(jax_array)jax.device_get
  • 是否每次循环都出现重新编译,观察日志里是否有Compiling字样。
  • 是否存在动态 shape 导致 XLA 缓存失效。

常见解决方法:

  • 把数据下划线放到 TPU 后,保持所有中间结果在 JAX 数组中。
  • jax.jit包裹更粗粒度的函数,减少 Python 调用次数。
  • jax.profiler导出 trace,看 host 与 device 之间的时间线是否有明显空洞。
import jax.profiler as prof prof.start_trace("/tmp/trace") # 执行若干训练 step for _ in range(10): params, loss = train_step(params, xb, yb) prof.stop_trace()

生成的 trace 文件可以用 Perfetto 在浏览器中加载。重点看 host 负责的数据加载、预取、D2H/H2D复制段,以及 TPU 执行段。如果每个 step 里都有一段很长的 host 数据处理,瓶颈大概率在数据管道而不是 TPU。

6.2 XlaRuntimeError、OUT_OF_MEMORY 和 libtpu 相关报错

下面这张表列出了 TPU 训练中最高频的几类报错,按现象分类给出处理方向。

报错现象常见原因检查方式处理建议
NotFoundError: libtpu.so not found未安装 jax[tpu] 或 libtpu 版本不匹配pip show libtpu重新安装jax[tpu],确认插件路径
RuntimeError: XlaRuntimeError: UNKNOWN编译失败或设备通信异常看完整堆栈中的 HLO 文件名导出 HLO 定位失败算子,检查算子是否受支持
Resource exhausted: Out of memoryHBM/VMEM 峰值超限看 HLO 内存分配日志调小 batch,启用梯度累积,检查是否反复 device_get
XlaRuntimeError: FAILED_PRECONDITION设备未就绪或通信组不一致检查多进程参数、coordinator 地址统一进程启动参数,确认所有 rank 使用相同配置
训练最后随机挂起多卡数据不一致导致 all-reduce 永远等待查看进程日志是否停留在某个 collective检查每个进程数据量是否相同,检查 mesh 和 sharding 是否匹配

遇到OUT_OF_MEMORY时,不要第一时间只调小batch_size。先在代码里导出 HLO,查看编译日志中峰值内存出现的位置。如果峰值来自某个巨大的中间结果,可以考虑算子融合、混合精度或把部分逻辑拆分到 CPU 端。

6.3 编译时间过长或反复重编译

现象:每次 step 都重新编译,或者第一个 epoch 要等很久,之后却很快。

TPU 上编译时间从数十秒到数分钟都算正常,尤其是小算子很多、需要跨设备通信的网络。但“反复重编译”通常由下面三类原因引起:

  • 输入 shape 不稳定,比如dyanmic pad后序列长度不一致。
  • Python 函数内部存在非静态控制流,例如依赖arr.shape[0]if
  • 不同batch size被混用,导致多个编译缓存并存。

推荐做法:

  • 固定训练输入 shape,统一到同一个 batch size。
  • @jit内部避免 Python 原生if,改用jax.lax.condjax.lax.switch
  • jax.jitstatic_argnums管理少数真正需要重编译的 Python 参数,不要把所有 Python 参数都设为静态。
  • 如果数据需要 padding,先统一到固定大小,再由 XLA 保证计算效率,而不是每次按 batch 内最大长度动态计算。

6.4 profiling 数据怎么采、怎么看

当系统能跑但性能不符合预期时,profiling 是最重要的证据来源。JAX 支持两种采集方式:一种是jax.profiler导出 trace,适合查看单进程时间线;另一种是jax.profiler.profile上下文管理器,适合只采集一段训练代码。

with jax.profiler.profile("/tmp/profile"): for step in range(20): params, loss = train_step(params, xb, yb)

打开 trace 后按时间线从上到下看:

  • host行:CPU 上执行的数据预处理、Python 调度。
  • TPU行:真正在 TPU 上执行的算子。
  • D2H/H2D行:TPU 和 host 之间的拷贝。

如果TPU行持续空白,说明任务在等数据。如果host行大量时间花在numpy操作,说明数据预处理没处理好。如果TPU行有算子但 gap 很大,说明 XLA 编译后的指令之间仍有明显依赖阻塞,可以从算子融合和 shape 入手。

7. 生产落地的关键约束与工程实践

7.1 学习环境与生产环境的差别

学习环境只要能跑通,生产环境则必须考虑稳定性、可观测性和成本。

维度学习/实验环境生产训练环境
数据路径内存或本地文件分布式存储、预取、缓存
失败恢复重跑脚本checkpoint 记录、自动重启、跳过已完成的 step
日志print基本够用结构化日志、指标上报
监控TPU 利用率、内存峰值、编译时间、all-reduce 耗时
配额少量按需创建提前申请配额、制定缩容策略
异常处理出错了就修自动重试、死信告警、人工介入

生产环境第一原则是:不要只验证模型能训练,要验证训练在任意时刻中断后都能恢复。为此,每个 epoch 或固定 step 数必须写 checkpoint,checkpoint 里至少包含参数、优化器状态、当前 epoch/step、随机数状态和数据偏移。

7.2 checkpoint、容错与配额

JAX 生态推荐使用 Orbax 管理 checkpoint。它支持同步和异步保存,也能处理多条jax.tree结构的数据。一个简化的保存思路如下:

from orbax import checkpoint as ocp ckpt_dir = "gs://your-bucket/mnist-demo" checkpointer = ocp.Checkpointer(ocp.PyTreeCheckpointHandler()) async def save_state(step, params, opt_state, rng): await checkpointer.async_save( ckpt_dir, args=ocp.args.PyTreeSave({"params": params, "opt_state": opt_state, "rng": rng}), )

这个示例用于说明思路。生产项目要结合自己的存储路径、训练状态结构和恢复策略调整。关键点是 checkpoint 保存必须是原子的,避免写到一半时进程被杀造成损坏。

配额是 TPU 生产落地的特殊约束。TPU 属于高价值加速器资源,尤其是多芯片 Pod 需要提前在项目中申请配额。上线时间敏感的大规模训练任务之前,先确认目标区域是否有足够的 v4/v5e/v5p 资源,避免临到发版才发现配额不足。

7.3 上线前检查清单

每次把 TPU 训练任务送入生产前,按下面的清单逐项确认:

  • [ ]jax.devices()返回真实的TpuDevice列表,数量符合预期。
  • [ ] 训练函数加入@jax.jit,且已验证不会因动态 shape 频繁重编译。
  • [ ] 数据管道支持按process_index和全局 batch 切分,不存在数据重复或漏读。
  • [ ] 混合精度配置明确,bfloat16float32的边界清楚。
  • [ ] checkpoint 包含参数、优化器状态、随机数状态和 epoch/step 信息。
  • [ ] 中断恢复演练通过:杀掉训练进程,重新启动后能从最新 checkpoint 恢复。
  • [ ] profiling 采集过一次,确认 host 到 TPU 之间没有大量等待空洞。
  • [ ] 日志能输出编译耗时、step 耗时、损失值和 TPU 内存峰值。
  • [ ] 配额和区域资源已确认,不会在训练中途因资源不足失败。
  • [ ] 训练输出具备可重复性:指定随机种子,并在 checkpoint 中保存 RNG 状态。

这个清单可以当作 TPU 训练服务发布前的准入标准。每一项都不难,但漏掉任何一项,都可能让故障出现在深夜训练启动后。

8. 扩展方向

8.1 PyTorch/XLA 适合什么场景

如果团队已经用 PyTorch 写了大量模型,迁移成本是不得不考虑的问题。PyTorch/XLA 让 PyTorch 代码可以编译到 XLA 并运行在 TPU 上,但它不是无缝替代。

PyTorch/XLA 的常见问题包括:

  • 某些算子编译效率不如 JAX 原生路径。
  • 动态 shape 对 XLA 编译器不友好,需要显式避免。
  • 分布式通信需要额外学习xm.optimizer_step等抽象。

适合选择 PyTorch/XLA 的场景是:模型已经成熟、短期不重构、团队熟悉 PyTorch。适合选择 JAX 的场景是:新项目、大规模训练、希望深入控制编译和分布式行为。

8.2 TensorFlow/Keras 迁移到 TPU 的做法

TensorFlow 2 可以通过TPUStrategy将 Keras 模型迁移到 TPU。基本流程是先定义分布策略,再在策略作用域内构建模型和数据集。

import 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(optimizer="adam", loss="mse") model.fit(train_dataset, epochs=10)

TPUClusterResolver会从环境变量或 gcloud 信息中读取 TPU 地址。迁移时最容易踩的坑是数据集没有按 batch 预先切分,以及 Keras 的steps_per_epoch设置不准确。建议在迁移前先跑一个很小的数据集合,确认分布式设备被正确初始化。

8.3 大模型训练:FSDP、张量并行与第三方生态

TPU 软件栈已经支持大规模模型训练。数据并行只是第一层,大模型通常还需要:

  • 分片数据并行,把参数、梯度和优化器状态切到多卡。
  • 张量并行,把单个矩阵乘法切成多个设备执行。
  • 流水线并行,把不同层放到不同设备组上。

JAX 生态中,jax.sharding配合 XLA 的 SPMD 编译器可以表达这些并行模式。Google 开源的 MaxText 项目就是基于 JAX 构建的大模型训练框架,支持 GPT、Gemma、Llama 等架构在 TPU 上训练。

对大模型训练感兴趣的同学,可以按这样的顺序学习:

  1. 先用jax.sharding.NamedSharding做数据并行。
  2. 再用PartitionSpec对权重矩阵做二维分片,实现张量并行。
  3. 最后用jax.lax.scan把 decoder 层展开成可编译的训练主循环。
  4. 结合 profiling 查看不同并行配置下的通信开销。

TPU 软件栈的门槛主要在“理解编译器如何把你的代码变成硬件指令”这一层。跨过这一层之后,多卡扩展、混合精度、故障恢复这些问题都有成熟的 JAX 生态工具可以复用。对新手来说,最有价值的练习不是立刻复现一个大模型,而是把一个 MLP、一个 ResNet 或者一个小 GPT 在 TPU 上从单卡逐步扩展到多卡,记录每一步的编译时间、显存峰值和吞吐变化,直到能解释每一个数字背后的原因。

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

ZKEYS V6.0.0深度解析:服务器销售管理系统从部署到实战

简介:阿帕云ZKEYS公有云服务器销售管理系统V6.0.0正式版是面向云服务提供商的一站式销售管理工具,覆盖云服务器、虚拟主机与域名等业务场景,帮助技术团队完成产品上架、订单流转、用户权限、计费账单、资源监控及工单客服等环节,适…

作者头像 李华
网站建设 2026/9/2 18:15:31

libtiff预编译库集成指南:32位与64位DLL配置与排错

简介:编译好的libtiff动态库与静态库资源,涵盖32位与64位两个版本,专为需要在C/C项目中快速集成TIFF读写功能的开发者准备。该库支持TIFF图像的读取、写入与修改,可处理多种压缩算法与色彩空间。通过预编译的DLL与LIB,…

作者头像 李华
网站建设 2026/9/2 18:12:59

XP11免费DC-10插件试飞全流程:从安装到冷舱启动指南

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/2 18:12:25

ThinkPad T480黑苹果OpenCore 0.6.6引导配置实战

简介:这是一份面向ThinkPad T480用户的OpenCore 0.6.6黑苹果引导EFI,适用于搭载i5-8250U、UHD 620核显的20L系列机型,旨在解决macOS安装引导与硬件驱动问题。压缩包共95个文件,约24.62MB,包含plist配置文件、aml主板补…

作者头像 李华
网站建设 2026/9/2 18:06:12

下一代模型快100倍?本地部署与推理优化实战

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/2 18:04:17

单片机计算机毕设之基于 STM32 单片机的 OLED 实时环境数据显示安防系统 基于 STM32 的消防险情感知与水泵、风机协同控制系统(012606)

博主介绍:✌️码农一枚 ,专注于大学生项目实战开发、讲解和毕业🚢文撰写修改等。全栈领域优质创作者,博客之星、掘金/华为云/阿里云/InfoQ等平台优质作者、专注于嵌入式单片机,Java、小程序技术领域和毕业项目实战 ✌️…

作者头像 李华