跳转至

TPU

张量处理单元 (Tensor Processing Unit, TPU) 是 Google 为机器学习设计的专用加速器。TPU 以大规模矩阵计算、高带宽内存和高速芯片互联为核心,通常通过 XLA 编译器栈使用,适合结构规则、可编译的大型训练与推理任务。

快速开始

在可用 TPU 的 JAX 环境中检查设备并运行矩阵乘法:

import jax
import jax.numpy as jnp

print(jax.devices())
x = jnp.ones((4096, 4096), dtype=jnp.bfloat16)
y = (x @ x).block_until_ready()
print(y.shape)

block_until_ready() 用于等待异步设备计算完成。基准测试若省略同步,测得的往往只是任务提交时间。

体系结构

TPU 的核心路径围绕矩阵乘法单元构建。大规模乘加运算以 Tile 形式进入阵列,数据在计算单元之间复用,从而减少反复访问外部内存。不同代际的 TPU 在矩阵单元、向量单元、HBM 容量、芯片互联和支持精度上有所差异。

TPU 不是只包含矩阵阵列。归一化、激活、索引和通信等操作仍需其他执行单元与内存系统配合。端到端性能取决于整个计算图能否被有效编译和分片。

XLA 编译

TPU 常通过 JAX、TensorFlow 或 PyTorch/XLA 使用。XLA 将高层张量程序转换为设备执行图,并执行算子融合、布局选择、内存规划和通信编排。

编译带来两个重要特征:

  • 稳定形状和规则控制流更容易优化;
  • 首次运行可能包含明显编译开销。

动态形状、依赖 Python 的数据分支或频繁改变序列长度,可能造成重新编译或性能退化。实践中常使用 Padding、Bucketing 和固定 Batch 降低形状变化。

Pod 与互联

多个 TPU 芯片可通过专用互联组成 Slice 或 Pod。训练框架把模型参数、数据和计算映射到设备网格,使用数据并行、张量并行或流水线并行扩展。

设备数增加后,集合通信与分片策略决定扩展效率。局部矩阵计算很快时,错误的网格映射会让通信成为主要瓶颈。

数据类型

TPU 广泛使用 BF16 等低精度格式。BF16 保留与 FP32 相同的指数位宽,动态范围较大,适合训练;有效尾数较短,累加与敏感算子仍需更高精度或稳定性处理。

具体代际可能支持 FP8、INT8 等格式。是否获得加速取决于编译器、模型形状和算子支持,不能只根据数据类型名称判断。

分片案例

以 Transformer 训练为例,可以构建设备网格:一个维度切分数据 Batch,另一个维度切分模型隐藏维度。数据并行维度需要同步梯度,模型并行维度在层内交换激活。

选择网格时应考虑:

  • 单层参数是否能放入单芯片;
  • Attention 与 FFN 的通信量;
  • 芯片互联拓扑;
  • Batch 是否足以填满数据并行副本;
  • 检查点能否在不同 Slice 规模间恢复。

适用场景

TPU 适合:

  • Google Cloud 中的大规模 JAX / TensorFlow 训练;
  • 形状稳定的 Transformer 和矩阵密集任务;
  • 能通过编译器表达的多设备计算;
  • 需要专用高速互联的规模化训练。

以下场景要谨慎评估:

  • 依赖大量 CUDA 自定义 Kernel 的项目;
  • 动态控制流和频繁变化的输入形状;
  • 只能在本地或特定云外运行的系统;
  • 团队缺少 XLA 调试和分片经验。

性能与成本评测

评测 TPU 时应把编译与稳态运行分开报告。至少记录:

  • 首次编译时间和缓存复用情况;
  • 每步训练时间、Token/s 与 MFU;
  • HBM 峰值和重计算开销;
  • 集合通信时间;
  • 失败恢复与检查点时间;
  • 按完成目标训练量计算的总成本。

按小时价格更低不代表任务总成本更低;编译、迁移和利用率都会影响结果。

常见问题

  • TPU 不是只能运行 TensorFlow,JAX 和 PyTorch/XLA 也可使用;
  • XLA 编译成功不代表分片和数据布局最优;
  • GPU Kernel 不能直接在 TPU 上运行,需要等价算子或重写;
  • 小模型、小 Batch 可能无法发挥大规模矩阵阵列;
  • 不同 TPU 代际的内存、精度和拓扑不能混为一谈。

与通用 GPU 的差异见 GPU,训练并行方法见 分布式训练。