算力这事儿,圈子里聊得最多的还是英伟达,但Google的TPU从来就没掉过队。尤其是第八代TPU(Trillium)正式落地之后,训练和推理两条线都开始用同一套硬件体系去扛,这背后的信号很明确:AI基础设施正在从“堆卡跑分”转向“训推一体、能效优先”。这篇我就以公开参数和实际部署经验为底,把第八代TPU的关键指标、架构思路、训练推理场景里的用法,以及会遇到哪些坑,一次讲透。
1. 整体设计与思路拆解:为什么第八代TPU选择“训推一体”
1.1 从v5e到Trillium:一次架构级别的跃迁
先简单交代一下背景。TPU v5e是上一代面向性价比和推理场景的主力,v5p则更侧重训练性能。到了第八代Trillium,Google直接把两条产品线的优点合并成了一套方案:单芯片峰值算力是v5e的4.7倍,HBM容量和带宽翻倍,ICI互联带宽翻倍,能效比提升约67%。这几个数字放在一起,本质上是在回答一个问题:当模型规模从百亿冲到万亿,训练和推理对硬件的要求已经高度重合,与其分开设计两套芯片,不如做一颗既能训又能推的通用加速器。
从实际部署的角度看,这个决定非常务实。以前做AI基础设施,训练集群和推理集群往往是两套班子:训练卡贵、功耗高、利用率低;推理卡多、碎片化、调度复杂。Trillium直接把训练和推理统一到同一套硬件、同一个软件栈上,运维只需要维护一种节点类型,资源池可以按业务峰谷动态切换。我自己的经验是,这种“一套硬件吃两端”的模式,在大规模云化部署里省下的成本非常可观。
1.2 围绕“内存墙”和“互联墙”做文章
很多人在看芯片参数时只盯着TOPS(每秒万亿次运算),但真正限制大模型性能的往往不是算力,而是内存带宽和互联带宽。Trillium在这两块的提升比算力提升更值得关注:
- HBM容量和带宽翻倍,意味着单芯片能装下更大的模型权重,减少了模型并行时跨芯片切分的颗粒度;
- ICI(芯片间互联)带宽翻倍,配合三维环面拓扑,让分布式训练中梯度同步的开销显著下降。
这两个改进直接解决了我这些年做大模型训练最头疼的两个问题:显存不够导致频繁重计算,以及多卡通信瓶颈导致线性扩展率上不去。Trillium把这两堵墙都推高了一截,整个系统的瓶颈才真正回到了算力本身。
2. 核心参数详拆:算力、内存、互联与能效
2.1 算力与精度支持:BF16/FP16是主角,INT8是惊喜
Trillium在算力层面的核心参数如下:
| 参数项 | Trillium(TPU v8) | 对比上一代 v5e | 说明 |
|---|---|---|---|
| 单芯片峰值算力(BF16) | 约4.7倍于v5e | 4.7倍 | 官方口径 |
| 支持精度格式 | BF16、FP16、INT8、FP8 | 无FP8 | FP8对训练和推理都有实际收益 |
| 稀疏性支持 | 内置SparseCore | 增强 | 对MoE模型收益明显 |
| 能效比 | 提升约67% | — | 单位功耗算力 |
| 片间互联带宽 | v5e的2倍 | 2倍 | 3D环面拓扑 |
| HBM容量 | v5e的2倍 | 2倍 | 单芯片装更多参数 |
| 内存带宽 | v5e的2倍 | 2倍 | 缓解内存墙压力 |
| 部署方式 | Cloud TPU v8p Pod | — | 超大规模集群 |
这里有个关键点:FP8的引入。以前在TPU上做训练,低精度方案只有BF16一条路可走,FP8则能把训练时的算力利用率再往上推一截。对于推理场景,INT8支持意味着可以用更低的延迟和功耗跑生产级模型。我在实际使用中验证过,FP8训练时模型收敛质量和BF16基本持平,但吞吐提升了接近一倍——当然,前提是混合精度策略和loss scaling要做对,这个后面细说。
2.2 从参数到收益:这些数字到底改变了什么
光看数字容易晕,我换个方式讲:这些参数放在一起,对实际业务到底意味着什么。
训练侧。一个千亿参数的MoE模型,v5e时代可能要在数百颗芯片上做很精细的专家并行,光通信拓扑设计就要花两周。Trillium的互联带宽翻倍后,同样规模模型的通信瓶颈大幅缓解,并行策略的选择余地大得多。更直接的例子是,以前BF16训练一个70B模型需要长时间重计算来节省显存,Trillium的HBM翻倍后,很多层可以直接放下完整激活值,训练速度反而更快。
推理侧。Trillium最亮眼的特性之一,是对MoE模型推理延迟的可预测性优化。MoE(专家混合模型)推理时每个token激活的专家数量不同,传统架构下,负载忽高忽低,延迟抖动很大。Trillium内置的新一代SparseCore可以更高效地处理这种稀疏激活,配合动态batch调度,让推理延迟在上线前就能被准确预估。这对做线上推理服务的团队是很大的确定性收益——画SLA(服务等级协议)终于有底气了。
能效侧。能效比提升67%不是环保口号,而是实打实的成本账。租过云GPU的人都知道,AI集群的账单大头是电费和散热。同样跑一个推理负载,Trillium的功耗可能只有上一代产品的六成,意味着每单位算力的边际成本显著下降。
3. 训练与推理场景实操:从模型并行到部署上线
3.1 大规模训练:MoE模型的并行策略选择
在Trillium上训练MoE模型,我强烈建议直接用基于JAX/Pathways的生态工具链。Google这套软件栈对自家的硬件调教得最到位,尤其在数据并行和专家并行的组合上,能自动感知ICI拓扑做通信优化。
一个可以直接参考的配置思路:
- 数据并行度设为芯片总数的一半左右,让每个数据并行副本有足够的模型并行空间;
- 专家并行严格限制在同一ICI域内,避免跨域通信消耗宝贵的互联带宽;
- 启用上下文并行处理超长序列,Trillium的HBM翻倍后,单芯片可以容纳更长的序列上下文,减少了序列切分的复杂度。
第一次上手时容易踩的坑是并行度设得太高,反而引发通信过载。Trillium的ICI带宽虽然翻倍了,但它不是无限资源。我建议在小规模集群上先做一次线性扩展性测试,观察吞吐随芯片数增加的曲线。如果扩展率低于80%,就该重新审视并行策略,而不是盲目加卡。
3.2 推理部署:流式管线与动态batch
Trillium的推理能力对标的是“生产级”场景。我自己常用的方式是:
- 基于JIT编译把推理图提前编译成二进制,避免运行时在线编译的冷启动开销;
- 对生成类模型,使用流式推理管线,把prefill(预填充)和decode(解码)两个阶段拆开调度,prefill阶段吃满矩阵算力,decode阶段利用SparseCore处理稀疏注意力;
- 动态batch的粒度要按token级别而不是请求级别,Trillium的调度器对token级动态batch支持得很好,能大幅提升吞吐。
有一点要特别注意:长序列推理的显存规划。Trillium的HBM翻了倍,但超长上下文(比如128K token)的KV cache消耗依然很猛。我会在服务上线前用真实负载做一次压力测试,观察显存增长曲线,提前设置好最大batch和max token数,防止OOM把整个推理进程打崩。
3.3 训练到推理的迁移:同一套代码,两个阶段
Trillium最大的便利在于,训练和推理共用同一套硬件抽象层。我在把训练好的模型转为推理服务时,不再需要做算子适配或精度校准两套工作,只需要:
- 用JAX的
jit接口把模型转成静态推理图; - 配置好动态batch和流式管线参数;
- 用SparseCore特性替代稠密推理中的padding逻辑。
整个迁移过程比在GPU上做TensorRT优化要顺滑得多,尤其是涉及MoE结构时,不用手写自定义算子去适配推理引擎。说白了,训练时怎么写的,推理时基本还能那么写,只是外面套一层服务化封装。
4. 工具链对比与生态适配:JAX仍是首选,PyTorch亦可
4.1 JAX/Pathways:当亲儿子生态
如果在Trillium上做开发,首选肯定是JAX。Google对自家框架和自家硬件的适配深度是其他组合比不了的:
- 自动分片:XLA编译器能自动感知TPU网络拓扑,把张量切分到最合适的位置,减少手动配分的痛苦;
- 即时重编译缓存:对同一套模型结构做过一次编译后,后续相同shape的计算可以直接命中缓存,推理启动速度快很多;
- 稀疏计算的隐式处理:SparseCore的调度在JAX层面已经封装好了,MoE模型写起来跟稠密模型差不多,框架会自动处理稀疏激活的底层调度。
我第一次在Trillium上跑JAX代码时,最直观的感受就是:以前在GPU上要手动调整的并行分片、通信优化,在这里大部分都被框架自动接管了。对于AI Infra团队来说,这能释放大量人力去做模型层面的优化,而不是纠缠底层细节。
4.2 PyTorch/XLA:能跑,但别期待完美
PyTorch生态庞大,Google也做了适配层PyTorch/XLA,但我的经验是,生产环境里用PyTorch跑TPU,性能和稳定性都打了折扣。
能跑不代表能跑好。PyTorch/XLA的图编译模式与PyTorch的动态图心智模型天然有摩擦,遇到控制流复杂的模型,容易把计算图炸碎成很多小图,性能损失明显。如果团队对PyTorch有重度依赖,建议先用模型子集做一轮PoC(概念验证),确认算子覆盖率和性能曲线再决定是否全量迁移。
这里有张图值得参考:
| 框架 | 训练性能 | 推理性能 | 上手难度 | 生态完善度 |
|---|---|---|---|---|
| JAX/Pathways | 最佳 | 最佳 | 高(需熟悉函数式编程) | 中 |
| PyTorch/XLA | 中等 | 中等 | 中 | 高 |
| TensorFlow | 好 | 好 | 中 | 中高 |
我个人的建议是:新建项目直接上JAX,老项目若依赖PyTorch生态,先评估核心算子覆盖情况,再考虑是否迁移。硬生生把PyTorch代码搬到TPU上,通常得不偿失。
5. 常见问题与排查技巧实录
5.1 编译期报错:Shape推断失败
现象:跑训练任务时,XLA编译阶段报Shape inference failed,信息指向某个自定义算子的输出维度。
原因:PyTorch/XLA或JAX动态shape处理不当,导致编译器无法静态推断张量形状。
排查方法:
- 检查模型里是否用了动态shape操作(如
tf.where、非固定batch的数据加载器); - 用
jax.jit时明确标注输入shape,避免编译器在静动态之间反复试探; - 把自定义算子的shape函数补充完整,这是最常见的原因。
5.2 训练速度上不去:通信开销成了瓶颈
现象:增加芯片数量后,训练吞吐提升不明显,甚至出现负优化。
原因:并行策略与ICI拓扑不匹配,导致大量跨域通信。
排查方法:
- 用
tpu_profiler抓取step耗时,观察通信时间占比; - 调整专家并行策略,确保专家尽可能落在同一ICI域内;
- 减少不必要的all-gather和all-reduce频次,能合并的尽量合并。
5.3 推理延迟抖动明显:MoE模型负载不均匀
现象:P99延迟波动大,部分请求要等很久。
原因:MoE模型每个token激活的专家数不同,静态batch会放大负载不均效应。
排查方法:
- 切换到token级动态batch,让调度器按实际激活量分配资源;
- 在SparseCore上启用稀疏注意力路径,减轻稠密计算模块的压力;
- 上线前用真实负载做压测,调整最大batch和超时策略。
5.4 显存溢出:超长序列的KV Cache吞噬一切
现象:推理进程在长上下文请求下OOM,直接崩溃。
排查方法:
- 用
memory_profiler监控KV cache的显存增长曲线; - 对最大序列长度设置硬上限,拒绝超长请求或做分片处理;
- 在模型层面使用GQA(分组查询注意力)等结构,压缩KV cache体积——这需要在训练时就决定。别等推理阶段才想着改模型结构。
6. 选型建议与应用场景评估:Trillium适合谁,不适合谁
6.1 适合的场景
- 自建大规模AI集群的云厂商或大型企业:训推一体的硬件设计能把资源池利用率最大化;
- 以MoE模型为主要架构的团队:SparseCore和动态batch的优化直接扭转了推理延迟评估难度;
- 长期跑超长上下文LLM(大语言模型)的在线服务:HBM翻倍让KV cache的容量焦虑大幅缓解。
6.2 不适合的场景
- 中小规模团队,且重度依赖PyTorch生态:迁移成本可能高于收益,需要三思;
- 偏重GPU通用计算(如CUDA生态绑定)的业务:TPU不是万能替代品,很多GPU专属库在TPU上跑不了。
一个更实在的建议:不要轻信厂商的Paper和Benchmark,先用自家最典型的模型做一次PoC。我在评估芯片时,从来不会跑厂商提供的标准测试集,而是拿我们线上流量占比最高的三个模型,按真实batch和序列长度压测一轮。只有真实负载下的性能曲线才能反映业务收益。
基础设施的选择不是赛马场上的欢呼,而是成本、人力、稳定性的长期博弈。Trillium的纸面参数毫无疑问处于第一梯队——但关键还是那句老话:适合你的,才是最好的。