新闻详情

新闻详情

首页 / 资讯中心 / 详情

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

发布时间:2026/9/2 18:26:35来源:尧图网络
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 早期依赖 GraphDefJAX 则通过jax.jit把函数追踪成 XLA HLO。编译器层是 TPU 软件栈最核心的部分。XLA 会做算子融合、内存分配、指令调度和张量布局选择。训练任务最终是否高效很大程度上取决于 XLA 能否消除冗余拷贝、把多个小算子合并成一个高效内核。运行时层是容易忽略的部分。PJRT 提供了设备抽象让 JAX 可以接到不同后端。TPU 后端通过 libtpu 与真实硬件通信。如果安装 JAX 时没有把 TPU 插件装好jax.devices()就只会返回 CPU 设备。1.3 为什么 JAX 是上手 TPU 的首选TensorFlow 也能跑 TPUPyTorch 也有 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_count8模拟多设备逻辑但真实 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 \ --zoneus-central1-b \ --accelerator-typev5e-4 \ --versiontpu-vm-base \ --projectyour-project-id版本字段tpu-vm-base只是一个示例。不同区域、不同芯片型号支持的镜像版本可能不同创建前先执行下面的命令确认当前可选项下面示例用于说明思路实际命令以后续官方文档为准。gcloud compute tpus tpu-vm versions list \ --zoneus-central1-b \ --projectyour-project-id创建完成后通过 SSH 登录到 TPU VM。gcloud compute tpus tpu-vm ssh tpu-demo \ --zoneus-central1-b \ --projectyour-project-id进入机器后先确认设备和系统信息。lsblk看数据盘nvidia-smi在这里不适用TPU VM 可以不用关心 GPU 驱动。2.2 在 TPU VM 上安装 JAX、TensorFlow 与 PyTorch/XLATPU 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-py、protobuf等公共库直接装最新版可能覆盖对方依赖。建议先创建独立虚拟环境再按项目锁版本安装。pip install -U tensorflow-cpuPyTorch/XLA 的安装命令会随版本变化一般需要通过 PyTorch 官方指示安装。这里给出一个通用参考pip install torch torchvision pip install torch_xlaPyTorch/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(id0, process_index0, coords(0,0,0), core_on_chip0), TpuDevice(...)] device count: 4如果输出只包含CpuDevice说明 JAX 没有识别到 TPU。从下面几个方向检查检查项命令或位置判断标准环境变量echo $TPU_NAME、$TPU_DRIVER_MODETPU_DRIVER_MODE1时通常使用 v5e 的 PjRt 路径插件文件python -c import jax.tools.collect_profile不应出现导包错误libtpupip show libtpu应能看到已安装版本JAX 后端python -c import jax; print(jax.default_backend())应为tpu或tpu相关插件注意不要只凭“程序能运行”判断环境正确。先执行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_rate0.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(fepoch {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 可执行的 HLOJAX 中看到的矩阵乘法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 与 SMEMGPU 开发者习惯用显存大小估算可训练模型规模但 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。如果后续调用传入不同 shapeXLA 会丢弃缓存并重新编译。动态 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_get或np.asarray这会强制把数据从 TPU 复制回 CPU造成 host 和 device 之间双向阻塞。正确做法是让复制发生在 batch 边界训练 step 内部只操作 JAX 数组。5.2 用 jax.sharding 把参数和梯度切到多张 TPU 上JAX 提供了jax.sharding体系用来描述数据如何分布到多个设备。最常用的是NamedSharding通过Mesh和PartitionSpec定义分片方式。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_namedata) params jax.tree_util.tree_map( lambda p, g: p - 0.01 * g, params, grads ) return params, lossjax.lax.pmean表示跨data轴对梯度取均值所有设备执行完成后得到一致梯度这是一般数据并行训练的核心通信原语。XLA 会根据设备拓扑决定 all-reduce 的执行方式不需要用户手写通信算子。这个示例用于说明思路实际项目需要结合自研模型结构调整分片维度。比如参数矩阵并行时需要把权重按列的维度切到不同设备并引入pall、psum等原语。5.3 多机训练与 PJRT 运行时需要注意什么当训练规模超过单 VM 的 TPU 芯片数量时需要把任务扩展到多台 TPU VM。JAX 在多机下的运行时由 PJRT 管理。PJRT 是一种可移植运行时抽象JAX 通过它统一访问 CPU、GPU 和 TPU 后端。多机训练需要关注三件事第一任务要显式指定机器的 coordinator 地址和 rank。JAX 常见做法是通过启动脚本设置环境变量例如JAX_COORDINATOR_ADDRESS、JAX_COORDINATOR_PORT和JAX_EXPECTED_DEVICES_PER_PROCESS。不同版本变量名可能不同落地前要查阅当前版本文档。第二进程数和设备数要对齐。一台 TPU VM 上启动多个 Python 进程时每个进程默认只能看到一部分设备。如果发现设备总数不对优先检查进程数是否为 VM 芯片数量的整数倍。第三数据集要按 global batch 切分。多卡训练时不能每张卡读取同一份 batch否则等于把同一个 batch 复制了 N 次。建议使用tf.data或torch.utils.data.DistributedSampler按process_index和process_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_getXlaRuntimeError: 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.cond或jax.lax.switch。用jax.jit的static_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 数必须写 checkpointcheckpoint 里至少包含参数、优化器状态、当前 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, argsocp.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 切分不存在数据重复或漏读。[ ] 混合精度配置明确bfloat16和float32的边界清楚。[ ] 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(optimizeradam, lossmse) model.fit(train_dataset, epochs10)TPUClusterResolver会从环境变量或 gcloud 信息中读取 TPU 地址。迁移时最容易踩的坑是数据集没有按 batch 预先切分以及 Keras 的steps_per_epoch设置不准确。建议在迁移前先跑一个很小的数据集合确认分布式设备被正确初始化。8.3 大模型训练FSDP、张量并行与第三方生态TPU 软件栈已经支持大规模模型训练。数据并行只是第一层大模型通常还需要分片数据并行把参数、梯度和优化器状态切到多卡。张量并行把单个矩阵乘法切成多个设备执行。流水线并行把不同层放到不同设备组上。JAX 生态中jax.sharding配合 XLA 的 SPMD 编译器可以表达这些并行模式。Google 开源的 MaxText 项目就是基于 JAX 构建的大模型训练框架支持 GPT、Gemma、Llama 等架构在 TPU 上训练。对大模型训练感兴趣的同学可以按这样的顺序学习先用jax.sharding.NamedSharding做数据并行。再用PartitionSpec对权重矩阵做二维分片实现张量并行。最后用jax.lax.scan把 decoder 层展开成可编译的训练主循环。结合 profiling 查看不同并行配置下的通信开销。TPU 软件栈的门槛主要在“理解编译器如何把你的代码变成硬件指令”这一层。跨过这一层之后多卡扩展、混合精度、故障恢复这些问题都有成熟的 JAX 生态工具可以复用。对新手来说最有价值的练习不是立刻复现一个大模型而是把一个 MLP、一个 ResNet 或者一个小 GPT 在 TPU 上从单卡逐步扩展到多卡记录每一步的编译时间、显存峰值和吞吐变化直到能解释每一个数字背后的原因。
网站建设高端定制企业官网
RELATED

相关资讯

更多精彩内容,欢迎继续阅读

较早相关资讯

最新相关资讯

JSBSim-1.0源码实操指南:从编译到六自由度飞行仿真 2026/9/2 19:14:42

JSBSim-1.0源码实操指南:从编译到六自由度飞行仿真

简介:JSBSim 1.0是一套开源飞行模拟框架的完整源代码,基于美国国家航空航天局公开的飞行力学数据构建,面向飞行仿真研究、航空航天教学、无人机控制与航电系统开发等场景。压缩包共461个文件,大小约1.35兆字节,主要包含…

阅读更多 →
时间轴联动插件:用TCP/UDP/串口实现视频与设备同步控制 2026/9/2 19:14:42

时间轴联动插件:用TCP/UDP/串口实现视频与设备同步控制

在实际的舞台灯光、会议系统和多媒体展项中,时间轴联动插件最典型的任务是:视频播放到某个时间点时,向不同 IP 地址的灯光控制器、音频处理器、交换机或播放器发送 TCP、UDP、串口命令,从而让视频内容与硬件设备保持同步。这类需求…

阅读更多 →
从Qwen大会看大模型技术趋势:部署、微调与Agent开发实战指南 2026/9/2 19:14:42

从Qwen大会看大模型技术趋势:部署、微调与Agent开发实战指南

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

阅读更多 →
Claude批处理实战:从JSONL到结果拉取的全流程指南 2026/9/2 19:14:42

Claude批处理实战:从JSONL到结果拉取的全流程指南

我们做 Claude API 开发时都会遇到同一个问题:单条请求处理少量文本很轻松,可一旦要跑几千条甚至几万条文本,比如批量翻译商品描述、给用户评论打标签、对日志做安全巡检,一条一条同步调用 API 不仅慢,而且成本容易被忽…

阅读更多 →
用AI Skill从代码自动生成架构图:告别手动维护 2026/9/2 19:14:42

用AI Skill从代码自动生成架构图:告别手动维护

如果你维护过一个超过半年、迭代了十几个版本的项目,大概率经历过这样的场景:文档里的架构图还是三个月前的单体结构,代码仓库里却已经拆出了四五个微服务模块。新同事入职,对着过期的架构图研究了半天,最后只能默默打…

阅读更多 →
三极管核心原理与实战:从水龙头比喻到开关放大电路设计 2026/9/2 19:11:41

三极管核心原理与实战:从水龙头比喻到开关放大电路设计

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

阅读更多 →

今日资讯

本周资讯

本月资讯

看完文章仍有疑问?

联系尧图顾问,获取一对一建站咨询

立即免费咨询 📞 400-888-8888
📞