Google TPU v6 (Trillium) 架构深度解析:从 MXU 矩阵引擎到 JAX 万卡分布式训练

为什么 TPU 仍是 AI 训练的基础设施标杆

在 NVIDIA GPU 主导的 AI 训练市场之外,Google TPU 一直是"隐形的冠军"。从 2016 年第一代 TPU v1 到 2024 年发布的第六代 TPU v6(代号 Trillium),Google 构建了从芯片到编译器、从单机到超大规模集群的完整 AI 训练栈。

TPU v5p(2024 年大规模部署)和 TPU v6/Trillium(2024 年底发布,2025 年量产)是当前 Google Cloud TPU 的主力型号。不同于 GPU 的 SIMT 执行模型,TPU 采用脉动阵列(Systolic Array)架构,为矩阵乘法这一 AI 核心运算提供了极致的硬件效率。

本文将从架构微结构出发,深入解析 MXU(Matrix Unit)矩阵引擎、ICI(Inter-Chip Interconnect)芯片间互连拓扑、HBM 内存子系统,然后过渡到 JAX/XLA 编译栈的分布式训练实战,最后探讨 TPU 与 GPU 在超大规模训练中的工程权衡。


TPU 架构演进:从 v1 到 Trillium

Google TPU 经历了六代演进,每一代都在矩阵引擎、内存带宽、互连拓扑三个维度持续突破:

代数代号发布时间矩阵引擎HBM峰值算力 (BF16)
v1-2016256×256 MXU8GB DDR392 TOPS (INT8)
v2-2017128×128 MXU ×264GB HBM245 TFLOPS (BF16)
v3-2018128×128 MXU ×232GB HBM2420 TFLOPS (BF16)
v4-2021128×128 MXU ×232GB HBM2e275 TFLOPS (BF16)
v5p-2023256×224 MXU ×296GB HBM3~459 TFLOPS (BF16)
v6Trillium2024256×256 MXU ×264GB HBM3e~500+ TFLOPS (BF16)

TPU v4 是一个分水岭:引入了光学电路交换(OCS,Optical Circuit Switching)和 3D 环面拓扑。TPU v5p 进一步将 MXU 从 128×128 扩展到 256×224,并提升至 96GB HBM3 内存。Trillium(v6)则在能效和稀疏计算上做了关键改进。


MXU 微结构:脉动阵列的数学原理

TPU 的核心计算单元是 MXU(Matrix Unit),本质上是一个二维脉动阵列(Systolic Array)。理解脉动阵列的原理,是理解 TPU 编程模型的关键。

脉动阵列的工作方式

以 4×4 脉动阵列计算矩阵乘法 D = A × B 为例:

b0  b1  b2  b3

↓ ↓ ↓ ↓ a0 → [pe00]→[pe01]→[pe02]→[pe03] a1 → [pe10]→[pe11]→[pe12]→[pe13] a2 → [pe20]→[pe21]→[pe22]→[pe23] a3 → [pe30]→[pe31]→[pe32]→[pe33]

每个处理单元(PE)在时钟周期内执行:

  1. 接收左侧输入(A 矩阵元素)
  2. 接收上方输入(B 矩阵元素)
  3. 执行乘累加:partial_sum += a × b
  4. 将 a 向右传递,b 向下传递

关键特性:数据在阵列中"流动",不需要寄存器文件或缓存访问。每个 PE 只需与邻居通信,实现了极高的能效比(每瓦算力)。

TPU v5p/Trillium 的 MXU 设计

TPU v5p 和 Trillium 的 MXU 为 256×224 和 256×256 的脉动阵列。这意味着每个 MXU 在单个时钟周期内可以执行 256×224(或 256×256)次乘累加操作。

每个 TPU 芯片包含 2 个 MXU(称为"TensorCore"),因此每芯片每周期可执行约 2 × 256 × 256 × 2 = 262,144 FLOPs(BF16 计算,因乘法和加法各算 1 FLOP)。

以 v5p 约 1.5GHz 时钟估算:

  • 单 MXU 每周期:256 × 224 × 2 = 114,688 FLOPs
  • 双 MXU 每周期:229,376 FLOPs
  • 单芯片峰值:229,376 × 1.5×10⁹ ≈ 344 TFLOPS (BF16)
  • 实际标称 ~459 TFLOPS(含稀疏加速)

Trillium 通过架构优化(包括更大的 MXU 和更高频率)进一步提升至 500+ TFLOPS。

数据类型支持

数据类型精度峰值倍率(vs BF16)
BF1616-bit 浮点1×
FP8 (E4M3)8-bit 浮点~2× (通过稀疏可达更高)
FP8 (E5M2)8-bit 浮点~2×
INT88-bit 整数~2×
INT44-bit 整数~4×(推理场景)

TPU v5p 和 Trillium 原生支持 BF16 和 FP8 训练,这是现代大模型训练的关键。FP8 的 2× 算力提升意味着在 Trillium 上使用 FP8 可达到 ~1 PFLOPS 的等效算力。


ICI 拓扑:从 2D Mesh 到光学交换

超大规模训练的关键瓶颈往往不是算力,而是芯片间通信带宽。TPU 的 ICI(Inter-Chip Interconnect)是其与 GPU NCCL 竞争的护城河之一。

3D 环面拓扑(TPU v4/v5p)

TPU v4 和 v5p 采用 3D Torus(环面)拓扑:

  • 每个 TPU 芯片有 6 个 ICI 端口(前、后、左、右、上、下)
  • 每个端口带宽 64 GB/s(双向)
  • 单个芯片总 ICI 带宽:6 × 64 = 384 GB/s(双向)
  • 4×4×4 的 Torus 共 64 颗芯片组成一个"Pod Slice"

环面拓扑的优势:任意两颗芯片间的平均跳数低,在多 workers AllReduce 时通信效率高。

光学电路交换(OCS)

TPU v4 引入的 OCS 是革命性的创新:

  • 传统电交换:固定拓扑,改拓扑需物理接线
  • 微机电系统(MEMS)镜面阵列:可在毫秒级重构光路
  • 一个 OCS 可连接数千个光纤端口,重构后形成任意拓扑

在训练中,OCS 允许同一个 TPU Pod 在不同作业间动态重配置:训练 GLaM 时用大环面,切换到小模型时用多个独立的 2×2×4 子 Pod 并行。

Trillium 的互连演进

TPU v6/Trillium 在 ICI 上的改进包括:

  1. 更高单端口带宽:从 64 GB/s 提升至 ~128 GB/s
  2. 稀疏通信原语:硬件级支持不规则通信模式(如 MoE 的路由分发)
  3. 改进的 OCS 切换时间:从毫秒级降至亚毫秒级

对于 MoE(Mixture of Experts)模型,稀疏通信至关重要。传统环面拓扑在 All-to-All 模式下效率较低,Trillium 的优化使其在 MoE 训练上相比 v5p 有约 30-50% 的加速。


内存子系统:HBM 层次结构

三级存储层次

TPU 的内存系统分为三个层次:

  1. VMEM(向量内存):软件管理的 SRAM,约 16-32 MB/芯片

  • Load/Store 到 VMEM 由 DMA 引擎执行
  • 向量核心(VREG)直接操作

  1. HBM(高带宽内存):片外 DRAM

  • v5p: 96GB,带宽 ~2.4 TB/s
  • Trillium: 64GB,带宽 ~2.8 TB/s
  • 虽然容量减少,但带宽显著提升(HBM3e 升级)

  1. 远程内存(通过 ICI):其他芯片的 HBM

  • 通过 Remote Direct Memory Access (RDMA) 访问
  • 延迟较高,但提供了逻辑上的统一内存空间

编程启示

在 JAX 中,数据在这些层次间的搬运是隐式由 XLA 编译器优化的。但理解层次结构有助于写出更优的代码:

  • VMEM 是瓶颈:每个 VMEM 访问需要 50-100 个时钟周期
  • Tile 策略:将计算数据切成小块(Tile),使其适配 VMEM 容量
  • 软件流水线:使用 jax.lax.prefetch 预取数据到 VMEM,隐藏 HBM 访问延迟


JAX/XLA:TPU 的编译堆栈

TPU 的编程体验通过 JAX 库和 XLA(Accelerated Linear Algebra)编译器实现。这不是简单的"写 TensorFlow 然后跑在 TPU 上",而是一个完全不同的计算范式。

JAX 的核心抽象

JAX 的 API 设计深受函数式编程影响,三个核心变换:

import jax

import jax.numpy as jnp

# 1. jit:XLA 编译,融合算子、优化内存布局 @jax.jit def forward(params, x): for layer in params: x = jnp.dot(x, layer['w']) + layer['b'] x = jax.nn.relu(x) return x

# 2. grad:自动微分 loss_fn = lambda params, x, y: jnp.mean((forward(params, x) - y) ** 2) grads = jax.grad(loss_fn)(params, x_batch, y_batch)

# 3. vmap:自动向量化/batch # 将单样本推理函数自动转换为批量推理 batched_forward = jax.vmap(forward, in_axes=(None, 0))

XLA 到 TPU 的执行流程

Python 代码 → JAX Tracing → HLO (High Level Optimizer) → LLO (Low Level Optimizer) → TPU 指令 → 硬件执行

关键优化发生在 LLO 阶段:

  • 算子融合(Kernel Fusion):将多个小算子融合成一个 TPU kernel,减少 HBM 访问
  • Layout 插入:自动插入数据布局转换(从 row-major 到 MXU 友好的 tiled layout)
  • 通信-计算重叠:将 AllReduce 与下一个计算步骤流水线化

实战:XLA 编译检查

# 查看 HLO 图,是 TPU 调优的必备技能

hlo_text = forward.lower(params, x).as_text() print(hlo_text[:5000])

# 使用 jax.profiler 追踪 TPU 执行 with jax.profiler.trace("/tmp/tpu_trace"): for _ in range(10): loss, grads = jax.value_and_grad(loss_fn)(params, x_batch, y_batch)


分布式训练实战

数据并行(Data Parallelism)

最简单的分布式策略,JAX 通过 pmap(已弃用)或 pjit(新 API)实现:

from jax.experimental import pjit

from jax.sharding import PartitionSpec, Mesh

# 定义设备网格:8×8 = 64 颗 TPU devices = np.asarray(jax.devices()).reshape(8, 8) mesh = Mesh(devices, axis_names=('data', 'model'))

# 数据并行:沿 data 维度分片数据,复制模型参数 data_sharding = pjit.PartitionSpec('data') # 数据按 batch 维分片 model_sharding = pjit.PartitionSpec() # 模型参数全复制

@partial(pjit, in_shardings=(model_sharding, data_sharding), out_shardings=data_sharding) def train_step(params, batch): # 计算梯度(每份数据独立计算) grads = jax.grad(loss_fn)(params, batch) # XLA 会自动插入 AllReduce 来同步梯度 params = optax.apply_updates(params, grads) return params

张量并行(Tensor Parallelism)

对于大模型(如 70B+),单芯片无法容纳模型权重,需要张量并行:

# 模型并行的简单示例:将线性层按列切分

@partial(pjit, in_shardings=(pjit.PartitionSpec('model'), pjit.PartitionSpec('data')), out_shardings=pjit.PartitionSpec('data')) def column_parallel_linear(weight_colshard, activation): # 每颗 TPU 只持有权重矩阵的一部分列 # 输出是局部结果,需要 AllGather 或 ReduceScatter return jnp.dot(activation, weight_colshard)

混合并行:数据 + 流水线 + 模型

实际的大规模训练通常结合三种并行策略。以训练一个 540B 参数模型为例:

# 3D 并行拓扑:2 (pipeline) × 4 (model) × 16 (data) = 128 TPU

devices = np.asarray(jax.devices()).reshape(2, 4, 16)

# 流水线并行:不同 pipeline stage 在不同行的 TPU 上 # 模型并行:同一行内的 TPU 协作计算同一层的不同部分 # 数据并行:同一列内的 TPU 处理不同的数据 batch

常见陷阱与性能调优

  1. Tile 大小不匹配:MXU 要求矩阵维度是 128 的倍数。如果不是,XLA 会自动 padding,但会带来有效算力下降。实践建议:设计模型时确保 hidden_size 是 128 的整数倍。

  1. HBM 带宽受限:BF16 训练的计算强度(Arithmetic Intensity)≈ 1 FLOP / 2 bytes = 0.5 FLOP/byte。TPU v5p 的算力/带宽比约为 459 TFLOPS / 2.4 TB/s ≈ 191 FLOP/byte。启示:BF16 训练永远是计算受限的,不用担心带宽。但 FP8 训练(~2 PFLOPS)则会变成带宽受限。

  1. ICI 拥塞:在 MoE 模型的 token dispatch 中,All-to-All 通信可能导致 ICI 拥塞。解决:使用 jax.lax.psum_scatter 替代全 AllReduce,或使用 Trillium 的路由优化硬件。

  1. VMEM 溢出(VMEM OOM):TPU 的 VMEM 仅 16-32 MB。如果使用过大的 attention head 或 batch size,会导致 XLA 频繁的 spills/fills。监控:通过 jax.profiler 观察 "MemoryBandwidthUtilization" 指标。


TPU vs GPU:工程权衡

TPU 的优势

  • 算力/功耗比更高:脉动阵列的确定性执行没有 warp divergence,硬件利用率通常 85-95%,而 GPU 实际利用率往往只有 50-70%(即使做了很好的 kernel 融合)。
  • 通信拓扑更优:环面 + OCS 提供比 NVLink 更灵活的拓扑重构能力。
  • 编译器自动化程度更高:XLA 在 TPU 上的优化空间比 CUDA 更大,因为 TPU 的硬件行为更确定。

TPU 的劣势

  • 灵活性受限:脉动阵列对非矩阵运算效率低。自定义 CUDA kernel(如 FlashAttention-3)在 GPU 上更容易实现。
  • 生态系统:PyTorch/XLA 虽有进展,但成熟度远不如 PyTorch/CUDA。自定义算子在 GPU 上可以用 Triton/Treeverse 快速开发。
  • 可用性:TPU 主要限于 Google Cloud,且需提前申请配额。GPU 在 AWS/GCP/Azure 都有广泛供给。

2025-2026 趋势观察

  1. Google 自用 vs Cloud 供给:Google 自身训练 Gemini、Gemma 等模型使用的 TPU 数量远超 Cloud 对外供给。Cloud TPU 的可用性正在改善(尤其是 Trillium)。

  1. MoE vs Dense:MoE 模型的稀疏计算在 Trillium 上有更好的硬件支持,这可能推动 TPU 在 MoE 训练上的市场份额。

  1. 推理场景:TPU v5p/Trillium 在推理上也有竞争力(尤其是 FP8 + INT8),与 GPU 的推理专用卡(L40S、Nvidia 下一代)形成竞争。


总结与展望

Google TPU 从 v1 到 Trillium 的演进,展示了一条与 GPU 不同的硬件设计路线:通过牺牲通用性换取矩阵运算的极致效率。在 AI 训练日益集中于大规模密集矩阵乘法的趋势下,TPU 的架构选择具有合理性。

对于工程师而言,理解 TPU 不仅仅是学习一个新的加速器,更是理解"当硬件为特定计算模式优化时,软件层需要如何重新设计"。JAX/XLA 的函数式编程模型、自动布局优化、全局通信规划,这些思想正在反哺 GPU 编程(如 PyTorch 的 torch.compile 和 Triton 编译器)。

下一代 TPU(v7 / "Bluejay")据传将进一步提升三维集成密度,可能在封装内集成 HBM4 和光学 I/O。与此同时,行业也在观望 UALink(Universal Accelerator Link)能否挑战 NVLink 在 GPU 间通信的地位——TPU 的 ICI + OCS 已经证明了替代方案的可行性。

对于想要在 TPU 上开始的开发者,建议路径是:先掌握 JAX 的函数式思维,然后通过 jax.profiler 理解 XLA 编译结果,最后在实际通信拓扑上调试分布式策略。TPU 的学习曲线陡峭,但一旦跨过门槛,你会发现它的确定性行为让性能调优比 GPU 更可预测。


本文基于公开的 TPU 架构论文(尤其是 ISCA 2024/2025 相关论文)、Google Cloud 技术博客、JAX 官方文档以及作者对 TPU v5p/Trillium 的工程实践经验撰写。

点赞(0) 打赏

评论列表 共有 0 条评论

暂无评论
立即
投稿

微信公众账号

微信扫一扫加关注

发表
评论
返回
顶部