
市面上讨论 Google TPU 的文章大多停在“TPU 很厉害”“算力比 GPU 强”这种层面真正讲清楚软件栈怎么用、怎么把模型跑起来、遇到问题怎么排查的内容反而不多。这篇文章要解决的就是“Google TPU 软件栈到底怎么玩”的问题。先给结论TPU 不是一个单独硬件名词它是“TPU 芯片 编译器 运行时 上层框架”一整套软件栈协同工作的结果。你要真正上手至少会碰到 XLA、PJRT、JAX、PyTorch/XLA 这几个关键词。适合看这篇文章的人主要是三类准备在云上跑大模型训练或微调的同学、想把已有 PyTorch 代码迁移到 TPU 上的团队、以及刚拿到 TPU 配额但不确定从哪里开始动手的人。最值得关注的点不是某个 API 怎么调而是这套软件栈里头有几条关键路径数据怎么进 TPU、计算图怎么编译、多芯怎么切分、日志和性能指标怎么解读。把这四条路径理顺了规模化落地才算有底。下面按实际使用顺序拆开讲。1. 先搞清楚 TPU 软件栈的运行条件和适用范围1.1 硬件和云环境是前提不是可选配置TPU 不像 GPU 那样买了显卡就能插在自己机器上。绝大多数人接触 TPU是通过云服务按需租用常见形态是 TPU VM 或 TPU Node。前者更像一台“自带加速卡”的虚拟机你可以用 SSH 登录直接装环境、跑代码后者则是一台独立的主机加 TPU需要通过网络去调用。因此确认资源时你不仅要看 TPU 型号还要看配套 CPU、内存、磁盘和网络带宽。举例来说TPU v4 的单芯算力强但如果你要跑的是大规模稠密推荐模型瓶颈可能不在算力而在内存带宽和嵌入表访问。TPU v5e 更偏向训练和推理适合中等规模 transformer 模型配置上相对均衡。TPU v2/v3 是较早的型号入门学习成本低但显存和带宽都有限不适合拿来硬跑超大模型。原项目材料没有给出具体版本对比表这里不写死。落地时需要去云平台确认当前计费、配额、可用区域。注意TPU 的“配额”通常不只看总量还看区域。有些区域 v5e 紧张换个区域反而能快速拿到资源。1.2 软件栈核心构成XLA、PJRT、JAX、PyTorch/XLATPU 软件栈可以拆成四层理解最底层是运行时负责把计算任务提交到 TPU 芯片。现在越来越常见的是 PJRT统一运行时接口它的作用是让 JAX、PyTorch 这类框架用同一种方式跟 TPU 对话。第二层是编译器核心是 XLAAccelerated Linear Algebra。模型代码不会直接变成 TPU 指令而是先转成计算图再经过 XLA 优化最后生成能在 TPU 上高效执行的底层代码。XLA 的编译质量直接决定你的模型能不能吃得满算力。第三层是框架主流有两种选择JAX 和 PyTorch/XLA。JAX 在学术界和 Google 生态里很常见写法更像 NumPy但支持自动微分和 JIT 编译PyTorch/XLA 则让现有 PyTorch 代码尽量少改动就能迁移到 TPU。第四层是数据与调度层比如数据加载、队列、故障重试、多任务并发。这一层最容易被忽略但规模化落地时它最影响稳定性和吞吐。这四层里大多数情况下你最需要跟 XLA 编译日志和 PJRT 运行时打交道。模型跑不出来、跑得慢、内存爆掉几乎都能从这两层找到原因。1.3 适合哪些任务不适合哪些任务结合实测和社区反馈TPU 软件栈更适合这些场景大规模矩阵计算密集的任务典型是大模型预训练、微调、多模态模型训练。需要高吞吐、批处理能力稳定的方案比如一次任务纳管上万条训练样本。已深度使用 JAX 或愿意把代码改成 JAX 风格的项目。已有 PyTorch 项目但团队愿意投入精力做一次迁移适配的场景。不适合的场景也要说清楚依赖大量自定义 CUDA kernel 的项目迁移成本高因为很多底层算子需要重写或替换。对生态依赖极强的任务比如某些只在 PyTorch 生态里完善的库TPU 上可能没有对应优化。小规模、短期实验申请资源、写分布式的成本可能比 GPU 单卡方案更高。这里必须提醒TPU 支持不等于所有模型开箱即用。支撑和优化是两码事。跑一次能通过和批量跑 100 次都稳定隔着整整一个工程化阶段。2. 第一次跑通 TPU 训练从按需分配 VM 到最小训练循环2.1 选择 TPU 类型和确认配额我第一次使用 TPU 时第一反应是创建一台机器结果发现第一步其实是查配额。你需要在云控制台里确认目标区域是否有对应 TPU 型号的库存同时确认项目配额是否足够。这一步做不好后面无论代码怎么写都启动不了。选择类型时先问自己三个问题模型规模多大。如果是几亿参数大概率的测试v5e 单芯足够如果模型到了几十亿、几百亿参数就要考虑多芯切分。训练模式是单机还是多机。单机多芯和多机多芯的配置方式完全不同网络拓扑、数据并行策略都会受影响。团队对 JAX 和 PyTorch 的熟悉程度。不要只追求“算力最强”选择团队能快速跑通的组合更实际。配额确认完之后再考虑环境镜像。当前 TPU VM 一般会提供包含 JAX、PyTorch/XLA 等预装组件的镜像。第一次上手我建议直接用官方镜像而不是自己在裸系统上一层层装否则环境问题会耗掉大量时间。2.2 创建 TPU VM 的通用步骤下面给出的是一个通用流程示例实际命令要以云平台控制台和当前版本为准# 1. 检查项目配额 gcloud compute project-info describe --projectyour-project-id # 2. 创建 TPU VM gcloud compute tpus tpu-vm create tpu-test \ --zoneus-central1-b \ --accelerator-typev5litepod-4 \ --versiontpu-vm-v4-base \ --projectyour-project-id # 3. SSH 登录 TPU VM gcloud compute tpus tpu-vm ssh tpu-test --zoneus-central1-b创建前要确认accelerator-type是否拼写正确、当前区域是否支持该型号。如果创建失败先查配额再查型号拼写最后查网络和权限。登录后可以先看设备和驱动信息ls /dev/accel* sudo lspci | grep -i google如果ls /dev/accel*输出为空说明 TPU 驱动没加载或者镜像版本有问题这时候先不要跑训练代码先解决环境。2.3 最小训练循环先跑一个能出数的 JAX 示例第一次跑通不需要复杂模型我的习惯是先用一个很小的 MLP 或者线性回归验证数据链路和 TPU 通信链路。下面是一个 JAX 示例用来验证基础训练流程import jax import jax.numpy as jnp from jax import grad, jit # 检查 TPU 是否可用 print(devices:, jax.devices()) # 构造最小线性回归数据 x jnp.linspace(0, 1, 128).reshape((128, 1)) y 3.0 * x 0.5 def loss_fn(params): w, b params pred x * w b return jnp.mean((pred - y) ** 2) jit def train_step(params, lr0.01): grads grad(loss_fn)(params) return params - lr * grads params (jnp.array(0.0), jnp.array(0.0)) for step in range(20): params train_step(params) if step % 5 0: loss loss_fn(params) print(fstep {step}, loss {loss:.4f})输出会显示 JAX 当前可见的设备列表。如果只有 CPU没有 TPU说明 PJRT 没识别到 TPU后续要检查依赖版本和环境变量。运行这个示例的意义不是训练精度而是确认整条链路通的代码能进入 XLA 编译流程编译后的计算能调度到 TPU张量能正常返回给 Python日志能正常打印如果这个示例都跑不通别急着上大模型。先排查环境。2.4 PyTorch 迁移时的最小验证如果团队用的是 PyTorch第一次验证建议用 PyTorch/XLA 跑一段非常小的模型。示例逻辑大致是这样import torch import torch_xla import torch_xla.core.xla_model as xm dev xm.xla_device() print(device:, dev) class TinyModel(torch.nn.Module): def __init__(self): super().__init__() self.fc torch.nn.Linear(8, 1) def forward(self, x): return self.fc(x) model TinyModel().to(dev) opt torch.optim.SGD(model.parameters(), lr0.01) x torch.randn(16, 8).to(dev) y torch.randn(16, 1).to(dev) for step in range(10): opt.zero_grad() loss torch.nn.functional.mse_loss(model(x), y) loss.backward() xm.optimizer_step(opt) if step % 2 0: print(fstep {step}, loss {loss.item()})这里最值得注意的就是xm.optimizer_step(opt)。它不只是做优化器更新还会触发 XLA 计算图的执行。如果直接调用opt.step()可能不会真正把梯度更新计算提交到 TPU这是新手最容易踩的坑。3. 分布式训练和批量落地数据布局与 XLA 编译3.1 数据布局先想清楚张量怎么切单张量在单芯上跑通之后下一个问题是多个 TPU 芯片之间数据到底怎么切。这个决策直接影响训练速度和稳定性。常见切分维度有四种数据并行每个芯片放一份完整模型副本把不同 batch 分给不同芯片。张量并行把单个大矩阵按行或按列拆开到不同芯片适合超大模型单层参数放不下的情况。流水线并行把模型按层切分不同芯片负责不同层适合极深的模型。混合并行以上几种组合使用。JAX 里常用jax.sharding和jax.device_put来控制数据如何分布。比如import jax from jax.sharding import Mesh, PartitionSpec # 假设有 4 个 TPU 芯片 devices jax.devices() mesh Mesh(devices.reshape((2, 2)), (data, model)) def shard_array(arr): return jax.device_put(arr, mesh) x jnp.ones((128, 1024)) x_sharded shard_array(x)用Mesh和PartitionSpec的目的是显式告诉 XLA 编译器“哪个维度切到数据并行哪个维度切到模型并行”。没有这一步XLA 可能会自动生成一个能跑但效率很低的布局导致通信开销剧增。3.2 XLA 编译优化不要忽略编译时间XLA 有一个显著特点第一次跑某个 shape 的训练步时会花较长时间做编译后续步数会明显变快。这个“前面的慢”不是死循环也不是 bug。为了减少首次编译时间尽量保持 batch size 和输入维度固定否则每次 shape 变化都会触发重新编译。执行逻辑上可以考虑这样固定模型输入 shape包括 batch size、序列长度、图像分辨率。先跑 10 到 20 步确认编译完成且 loss 在下降。再开启完整训练循环避免在正式训练阶段反复编译。如果使用动态 shape 会显著增加编译次数除非必要否则不做。注意在 PyTorch/XLA 里不要在每个 step 都调用xm.mark_step()。过度标记会让 XLA 无法合并计算图导致整体吞吐下降。3.3 批量训练时的任务管理、失败重试和输出命名从 Demo 到规模化最大的变化不是模型更大了而是任务管理复杂度上来了。批量训练至少要考虑四个点输入文件列表管理。建议把训练样本清单写成 manifest 文件而不是直接用os.listdir读目录。理由是可以记录已经处理过哪些文件失败时能断点续跑。输出命名。如果多个任务并发写同一个输出目录文件名必须包含唯一 ID否则会出现互相覆盖。失败重试。单条数据出错不要整个任务崩掉。最常见的方法是捕获异常后把失败的样本 ID 写到单独日志文件最后统一重跑。资源上限。即使 TPU 显存够大也不要让单任务吃满全部内存。批量任务需要预留一部分资源给其他进程否则并发一多整个 VM 可能直接 OOM。训练脚本里可以这样组织批量任务import glob import json import os failed_ids [] output_dir ./outputs os.makedirs(output_dir, exist_okTrue) for sample_path in glob.glob(./input_data/*.json): sample_id os.path.basename(sample_path).replace(.json, ) try: result train_one_sample(sample_path) out_path os.path.join(output_dir, f{sample_id}_{os.getpid()}.json) with open(out_path, w) as f: json.dump(result, f, ensure_asciiFalse) except Exception as e: failed_ids.append({id: sample_id, error: str(e)}) print(ffailed: {sample_id}, error: {e}) with open(failed_ids.json, w) as f: json.dump(failed_ids, f, ensure_asciiFalse)这个结构看起来简单但它是批量任务最稳定的底座。等失败清单落盘后再针对这些样本单独排查。4. 性能分析与参数调优先看编译图再调参数4.1 从哪些指标看训练是否健康性能调优最容易犯的错误是一上来就改学习率、改 batch size却不知道当前瓶颈是显存、编译、通信还是数据加载。我的建议是先看这五类指标设备利用率TPU 核心的 busy 比例是不是持续在高位。如果长期低于 50%很可能在等待数据或等待编译。编译时间占比如果编译时间占整体时间超过 20%需要检查 shape 是否频繁变化。输入数据吞吐每秒能从磁盘或网络加载多少样本。通信耗时多芯片之间的 all-reduce 聚合时间。Step time单步训练耗时。固定 shape 下step time 应该平稳如果突然跳动优先怀疑网络或磁盘抖动。这些指标可以通过 XLA 自带的 profiling 工具查看也可以自己写简单的计时逻辑。第一次统计时不用太复杂先记录 step time 和编译耗时基本就能定位大多数问题。4.2 核心参数怎么取舍参数没有绝对最优只有相对你的模型和任务最优。下面给一组通用判断维度参数作用调大后可能的影响调小时的影响batch size决定每个 step 处理多少样本提高吞吐但显存占用上升降低显存压力但训练变慢learning rate控制参数更新幅度收敛快但不稳定稳定但可能收敛慢gradient accumulation模拟更大 batch可突破显存限制单步吞吐下降precision精度模式低精度速度快但可能损失稳定性高精度更稳定但更慢num_epochs训练轮数可能过拟合可能欠拟合max_shard_size单芯承载参数大小减少切分次数可能增加通信开销我一般建议从小 batch 开始比如 16 或 32确认正确性和稳定性之后再逐步翻倍。不要一上来就跑大 batch尤其当模型很大时显存爆掉后排查成本很高。4.3 实测中常见的性能瓶颈真实测试中TPU 训练慢不一定是因为算力不够。我遇到过的典型情况有以下几种。第一种是数据加载慢。TPU 计算很快如果数据由 CPU 从磁盘读取再传输到 TPU这个过程如果跟不上训练速度设备利用率会被拖低。解决办法是使用异步数据加载、把数据放到更快存储、或提前打包成适合随机读取的格式。第二种是矩阵形状不规整。比如 NLP 任务里 padding 过多导致 30% 的计算都在处理无效 token。此时设备利用率看着很高但实际有效吞吐很低。先做 batch 内长度排序或者更紧凑的打包比调 XLA 参数更有效。第三种是gradient accumulation和batch size双高导致 XLA 内存峰值暴涨。低精度与梯度累积叠加时尤其明显。遇到 OOM先降低 batch size 或 accumulation steps而不是盲目减少模型参数量。第四种是 checkpoint 保存和验证评测太频繁。每轮都全量保存大模型磁盘写入开销很大而且保存时可能阻塞训练循环。建议保存频率降低或者使用异步保存。5. 容易踩坑的排查清单5.1 启动阶段创建失败、设备不可见、驱动未加载启动阶段报错先从外到内排查配额是否足够区域是否支持该型号。SSH 是否能正常登录权限是否足够。ls /dev/accel*是否能看到设备。镜像是否包含所需运行时比如 JAX 或 PyTorch/XLA 版本是否匹配。依赖版本是否与当前 TPU 镜像兼容。我遇到过最典型的情况是jax.devices()只返回 CPU不返回 TPU。原因通常是 PJRT 插件没安装或者环境变量没指向 TPU。查这个问题时先看pip list | grep jax、pip list | grep pjrt再看jax.config是否启用 TPU。5.2 运行阶段编译卡住、训练时 OOM、速度特别慢运行阶段的问题按这条链路排查先看日志有没有报错再确认输入数据 shape 是否发生变化然后看显存和内存占用再看 XLA 编译图是否频繁重编译最后看网络通信是否异常。编译卡住通常不是死机而是某个算子正在做耗时的 XLA 优化。如果同一段代码编译时间超过 20 分钟可以检查是否有自定义算子、是否有动态 shape、是否用了大量不支持的高阶操作。OOM 则要区分是“显存不够”还是“内存不够”。TPU 报 OOM优先降低 batch size、降低精度、减少中间张量保存数量如果是主机内存 OOM要查数据加载时是否一次性读入太多文件。速度特别慢时先看是不是数据加载卡顿。如果 CPU 加载数据跟不上 TPU你会看到设备利用率低但 CPU 占用高。这种问题改编译器参数没用要去优化数据管线。注意多芯片训练里all-reduce 通信延迟会被放大到所有芯片的同步等待里。如果发现网络抖动整个训练进程都会卡顿表现就是 step time 突然飙升。5.3 输出不稳定loss 不下降、loss 为 NaN、结果不可复现Loss 不下降先看数据预处理和标签是否正确不要先怀疑 TPU。Loss 为 NaN优先检查学习率是否过大、是否有除零操作、低精度模式下是否有溢出。结果不可复现可能来自随机种子管理不严格也可能来自并行聚合顺序不一致。建议固定好 numpy、Python 和框架的随机种子并且把数据加载顺序固定下来。如果多机训练时结果不稳定很大概率是全局 batch size 变化之后没有同步调整学习率。这些看起来是模型层面的问题但实际上和分布式策略强相关。5.4 日志和诊断信息怎么看规模化落地时日志比性能指标更救命。每条训练日志建议至少包含当前时间戳全局 step 和 epoch当前任务 ID 或样本 IDloss、learning rate设备利用率、内存占用、step time是否触发重编译错误摘要不要把日志全部堆在标准输出里。建议分两类常规进度写入文件错误和异常单独写到错误日志。失败重跑时直接以错误日志为主不解析 stdout。6. PyTorch 迁移到 TPU 时的时间投入和方案取舍6.1 先评估代码里有哪些算子需要改造不少团队的实际模型是用 PyTorch 写的硬迁到 JAX 成本太高所以更现实的选择是 PyTorch/XLA。但 PyTorch/XLA 不是万能兼容层有些操作支持得很好有些操作会退化成 CPU 执行或者直接报错。迁移前建议做一次代码扫描重点看是否依赖自定义 CUDA extension。是否用了大量torch.where、高阶 mask、动态索引。是否依赖某些只在 CUDA 上有优化的库比如 apex。模型内部有没有频繁改变 tensor 的 shape。有没有分布式训练代码比如 DDP需要替换成 XLA 对应的并行 API。扫描结果出来后先评估“改造工作量”和“预期收益”。如果只是单卡训练迁移收益可能不大如果是要跑大规模预训练才值得投入。6.2 PyTorch/XLA 的常用适配点PyTorch 代码迁移到 PyTorch/XLA 时需要改动的地方通常集中在四处设备初始化用xla_device()替代cuda。分布式封装不能用普通 DDP要使用 XLA 对应接口。优化器更新使用xm.optimizer_step替代optimizer.step()。checkpoint 保存和加载需要先xm.mark_step()确保计算图执行完成再保存模型状态。这里给一个简化的 checkpoint 保存逻辑import torch import torch_xla.core.xla_model as xm # 训练若干步后保存 xm.mark_step() # 触发计算图执行 state { model: model.state_dict(), optimizer: opt.state_dict(), step: step, } xm.save(state, fcheckpoint_{step}.pt)xm.save会处理 TPU 上张量和 CPU 之间的拷贝避免你在保存时踩到设备不一致的报错。6.3 什么时候该考虑用 JAX 而不是 PyTorch这不是要挑起框架之争而是实际项目里必须做出的选择。我的判断标准很简单如果团队只熟悉 PyTorch项目周期紧那么尽量用 PyTorch/XLA不要临时换 JAX。如果项目是从零开始、模型主力是 transformer、且未来要在 TPU 上长期迭代那么 JAX 值得认真考虑。如果要做大量自定义并行策略JAX 的 sharding 抽象更灵活适合深度调优。如果只是跑现成代码PyTorch 生态里能找到更多现成实现。选择原则不是“谁更好用”而是“谁更可能让你在一个月内跑通并持续迭代”。7. 规模化落地的最后几个提醒真正把 TPU 软件栈用于规模化训练时我建议把注意力从“模型能不能跑”转移到“任务能不能稳定跑完”上来。以下三件事经常被忽略但影响最大。第一预先定义好“什么叫跑成功”。比如标准是 1000 个 step 内 loss 降到某个范围还是平均 step time 低于多少秒。没有验收标准性能优化和目标设定就会变得很散。第二定期备份任务元数据。至少保存模型结构、数据版本、代码 commit、关键超参、TPU 型号和软件栈版本。这样出问题时能快速复现而不是靠记忆。第三把失败当作正常情况设计。数据处理偶发异常、网络偶尔抖动、存储偶尔变慢这些都不是“不可能发生”而是“大概率发生”。用 manifest、重试、失败清单、输出 ID 这套组合才能把偶发异常消化在系统内部。如果要长期跑 TPU 训练不要把目录结构、日志命名、任务队列这些东西临时拼凑。先把它固定下来后面每次训练只是往里追加数据不用反复造轮子。写在最后Google TPU 软件栈给我的整体感受是入门门槛不在硬件资源而在对 XLA 编译、PJRT 运行时和分布式切分的理解。第一次跑通最小示例不难难的是稳定地跑完大规模任务。如果你现在正准备上 TPU建议先把“单任务跑稳”作为第一个里程碑不要直接冲多机大模型。等日志、输出目录、失败重试、性能指标都正常之后再去扩展模型规模和集群数量。踩过几次坑之后回头看很多问题不是 TPU 能力不够而是软件栈的前置环境没有理顺或者输入数据形态不稳定。只要把这些基础打牢TPU 在训练吞吐和大规模并行方面的优势才能真正显现出来。