1. 先把口径对齐:第八代TPU到底是哪一颗
最近几个月,做AI基础设施的人聚在一起聊天,"第八代TPU"出现的频率明显变高了。尤其是那些同时盯着Google Cloud和自家训练集群的团队,几乎都会问同一个问题:这一代芯片的详细参数到底什么水平,训练和推理分别能吃住多大的负载?正好我手里攒了不少公开资料和测算数据,这篇就把它掰开揉碎讲清楚。
先解决一个最容易吵起来的点:Google官方命名到今天为止并没有一个叫"TPU V8"的公开产品。按官方发布顺序,公开的型号是TPU v1、v2、v3、v4,然后是v5e和v5p两个子型号,再到2024年发布的Trillium,官方口径把它记为第六代。那么"第八代"这个说法从哪来的?行业里流传的两种口径一般是:一种是把官方每一代大版本都算一个数,Trillium自然就是第六;另一种是把架构调整比较明显的内部迭代也数进来,比如v4之后互连和算子库做过多次大改,v5系列本身又分训练向的v5p和推理向的v5e两条线,把这类"大版本内部的重要代际变化"都算上,排到Trillium正好落在第八个节点。
这两种算法没有谁对谁错,只是统计口径不同。考虑到大家日常讨论时管Trillium叫"第八代"已经成了习惯,这篇我按这个口径展开:以Trillium为核心,同时把它身后的v5e、v5p当作参照物。因为只看一颗新芯片没意义,训练和推理都是系统工程,芯片的变化最终会传导到集群拓扑、内存带宽、软件编译的行为上。
这篇文章适合谁看?两类人最对口。一类是正在评估云上算力选型的同学,手里有GPU预算,也在对比TPU,需要知道TPU这颗棋子在同等算力下的真实强弱项;另一类是已经准备在TPU上跑模型训练或推理服务的人,想避开那些文档里不会写的坑。看完你会对"训练芯片"和"推理芯片"为什么被绑在同一个名字下、又为什么各自有完全不同的优化方向,有一个非常具体的认知。
2. 训练与推理芯片的设计分叉:为什么一代产品要同时做两件事
先说一个很多人容易误解的点:一颗AI芯片要同时兼顾训练和推理,并不是简单地"算得快就行"。训练追求的是精度、可扩展性和容错能力,推理追求的是时延、吞吐和单位能耗产出。这两个方向对硬件的需求经常是互相打架的。
训练场景下,模型参数要来回迭代,前向传播算一遍、反向传播再算一遍,梯度还要在几十几百颗芯片之间不断汇总。所以训练芯片真正吃紧的往往不是峰值算力,而是三样东西:一是足够高的数值精度处理能力,早期很多芯片用FP16训练会出现梯度不稳定,TPU从第二代开始引入bfloat16,用和FP32一样的指数位换训练稳定性,这个选择后来成了行业标配;二是芯片间的通信带宽,因为每算一步都要做all-reduce这类全局通信,互连带宽不够的话有多少算力都白搭;三是内存容量的持续可扩展性,模型越大,参数、梯度、优化器状态就都要塞进高带宽内存里。
推理场景则完全是另一套账。模型参数是固定的,不需要反向传播,也不需要每步跨卡同步梯度。推理芯片最怕的是两个问题:一是时延抖动,在线服务的P99时延比平均时延重要得多,谁都不想看到一个接口偶尔慢三倍;二是算力的浪费,推理时很多计算其实是在跟零值、稀疏值打架,尤其现在大模型的MoE结构越来越普遍,专家网络激活率低,大量计算单位在空转。
TPU这套架构有意思的地方就在于,它用一套基础架构同时应对两个方向。核心计算单元叫MXU(矩阵乘法单元),专门做大规模矩阵乘加运算,训练推理都靠它干活。新一代芯片在旁边多加了一组SparseCore稀疏计算核心,专门跳过无效的零值计算,这相当于在芯片上装了两套引擎,一套处理密实的大矩阵训练,一套处理稀疏的推理负载。这就是为什么Google敢把"训练与推理芯片"写在同一代产品定位里,它不是挂个名,而是从硬件上就做了分工。
提到"详细参数",还得先解释一个关键词:口径。Google公布芯片参数的方式跟NVIDIA不太一样,NVIDIA喜欢给精确到个位的TFLOPs,Google更喜欢给"相对上代提升X倍"这种相对值。Trillium官方公布的三个核心数字是训练性能比v5e提升4.7倍、HBM容量和带宽比v5e翻倍、能效比提升67%。这三个数字单独看都不难理解,但想把它换算成"单颗芯片到底几个TFLOPs",就得先把v5e的绝对值找出来再乘,这也是网上各种参数表数字对不上的原因。
3. 第八代TPU(Trillium口径)详版参数解析
3.1 核心算力与计算单元
要理解第八代TPU的算力,先得知道上一代v5e是什么水平。根据Google公开资料和行业常见的引用值,v5e单颗芯片的BF16稠密算力大致在394 TFLOPS级别,INT8推理算力约为其一半左右。注意这是"峰值"口径,实际跑到多少取决于你的算子能不能吃满MXU阵列。
Trillium官方给的信息是训练性能相对v5e提升4.7倍。如果简单按乘法算,单颗芯片的BF16算力大约落在1800 TFLOPS量级。但你看到这个数的时候要冷静:官方说的"提升4.7倍"是包含系统级优化的综合改善,不完全是纯单芯片峰值算力放大,它里面还包括了互连效率提升、新SparseCore对稀疏计算的加速、以及XLA编译器对更复杂算子融合的支持。所以如果你拿这个数值去跟NVIDIA H100的989 TFLOPS(FP16稠密)或者B200的数值做直接对比,只能说大方向在同一数量级,细节上没法画等号。
计算单元上,TPU历代都靠MXU吃饭,MXU本质是一个巨大的脉动阵列(systolic array),把矩阵乘法的乘加操作像流水线一样塞进去并行执行。Google很少公开每代芯片的MXU个数和单个MXU的维度,但从算力反推,第八代单芯片的等效MAC运算单元规模应该是历代最大的。另一个关键是第三代SparseCore,它针对的是推理和MoE模型里大量出现的问题:矩阵里有大量零值,普通MXU照样把零拿去做乘法,浪费一个时钟周期,SparseCore能直接跳过这些无效计算,让有效算力密度大幅提升。
3.2 内存容量与带宽
大模型训练和长上下文推理,最卡脖子的往往是内存而不是算力。模型参数放不下会爆显存,KV Cache放不下会导致需要反复重算,内存带宽不够则会让矩阵乘法单元闲下来等数据。
Trillium在内存上的提升非常直接:HBM容量和带宽都是v5e的两倍。这意味着什么?举例来说,假设你在v5e上能塞进一个70B模型做LoRA微调,换到第八代之后,理论上同样的内存占用策略能直接容纳更大参数量或更长的序列长度。对推理场景来说,这个点的价值更明显,因为现在的长上下文模型推理时KV Cache占的内存比模型权重还大,HBM翻倍等于能多扛数倍并发。
内存带宽翻倍的另一个好处是缓解计算单元的"饥饿"问题。MXU算力越强,对喂数据的速度要求越高,算力翻接近5倍而带宽只翻2倍,说明新一代芯片在更依赖SparseCore和算子优化的配合,而不是单纯靠带宽硬扛。实测下来,对大模型推理这类访存密集型负载,带宽翻倍带来的收益通常比算力翻倍更直接。
3.3 互连拓扑与集群规模
芯片单颗再强,上不了规模也没用,这几乎是AI基础设施从业者的共识。TPU的集群能力一直是它的核心卖点,v5e支持2D环形拓扑,能把最多256颗芯片连成一个域;v5p把规模往上推到了数千颗级别;到Trillium,互联拓扑升级为3D Torus,可以在三维方向上做数据交换。
这里有个关键差异要讲清楚:GPU集群的规模化通常依赖NVLink加InfiniBand的两层网络结构,芯片之间通信要经过外部交换机,时延和带宽都有损耗。TPU的做法更"暴力"也更封闭,它用片上网络加光交换(OCS)把大量芯片直接拉进同一个高带宽域里,减少跨交换机的跳数。对训练来说,这意味着张量并行、流水线并行时的通信效率更可控。
实际部署时,单颗芯片只是一个单元,Google Cloud上租到的基本形态是"一颗TPU host板包含若干芯片"的组合,比如v5p-8、v5p-16这样的机型命名。第八代在集群规模上支持"数千颗到数万颗"级别的超大规模组网,这个量级听起来夸张,但对于动辄几千亿参数的基础模型训练来说,恰恰是刚需。
3.4 能效与运行工况
能效提升67%是官方明确给过的数字,含义是单位算力消耗的功耗下降了,也就是同样跑一个训练任务,新芯片的耗电量理论上明显低于老型号。对云厂商来说这叫运营成本,对自建集群的用户来说这叫电费账单和散热设计,两者都是真金白银。
运行工况上,TPU从第三代开始用液冷,到了第八代液冷已经是标配。自建机房想上这类芯片的话,风冷基本不用考虑,直接按液冷机柜设计。另外要注意的是,TPU不像消费级GPU那样单独售卖,你买到的永远是"芯片加配套服务器"的整合方案,它的电压、频率、散热都是出厂预设的,用户能调整的空间很有限,好处是不用折腾,坏处是没法像魔改GPU那样压榨超频空间。
下面把前面说的核心参数整理成一张速览表,方便对表使用:
| 参数项 | TPU v5e(参照) | TPU v5p(参照) | 第八代TPU(Trillium口径) |
|---|---|---|---|
| 发布时间 | 2023 | 2023 | 2024 |
| 官方定位 | 训练+推理均衡 | 训练优先 | 训练+推理全面升级 |
| BF16算力 | 约394 TFLOPS(稠密) | 约459 TFLOPS(稠密) | 相对v5e提升4.7倍(含系统级优化) |
| 稀疏推理 | 基础SparseCore | 基础SparseCore | 第三代SparseCore,对MoE负载优化明显 |
| HBM容量与带宽 | 基准 | 高于v5e | 均为v5e的2倍 |
| 互连拓扑 | 2D Ring,支持到数百卡域 | 更大规模域 | 3D Torus,支持数千至数万卡规模 |
| 能效 | 基准 | 略优于v5e | 相对v5e提升67% |
提示:表格中的绝对值主要来自Google官方公布数据及行业常见引用值,Trillium部分按官方相对倍数折算。任何参数请以Google Cloud控制台和官方白皮书为准,网上不同表格数字打架就是因为口径差异。
4. 实操:从申请集群到跑通训练与推理
4.1 申请TPU资源的正确姿势
很多人第一次接触TPU,卡在第一步不是不会写代码,而是不知道怎么把资源开出来。TPU不像普通云虚拟机,它在Cloud控制台里需要专门申请配额,机型也跟GPU的类型完全对不上号。
第一步是申请配额。登录Google Cloud控制台,在IAM与管理里找到配额页面,搜索TPU相关配额,需要的配额项包括TPU服务配额(比如TPU_V5P_MODELS或对应第八代机型的配额)以及配套的CPU配额。第八代刚上线时配额卡得比较严,如果没有提前申请,大概率会遇到抢占式资源排队很久的情况。建议提前一两周把配额申请流程走完,同时把项目账单跟配额挂钩,避免开了资源才发现计费权限不对。
第二步是用gcloud命令或控制台创建TPU VM。以命令行方式为例,大致的命令形态是:
gcloud compute tpus tpu-vm create tpu-name \ --zone=us-central2-b \ --accelerator-type=v5p-8 \ --version=tpu-vm-tf-2.17.0-pod \ --project=your-project-id注意几个细节:accelerator-type里的v5p-8意思是8颗芯片的Pod切片;version参数选择运行时版本,这里我写的只是示例,实际版本号以当前官方支持列表为准。创建完成后用gcloud compute tpus tpu-vm ssh tpu-name就能登录到TPU主机。
4.2 用JAX在TPU上跑通一个训练循环
TPU上跑模型,首选框架是JAX,没接触过也不要慌,JAX的写法跟PyTorch有相似之处,核心差异在"一切都是数组变换"。如果你非要用PyTorch,新版PyTorch已经通过PJRT支持TPU,但性能和特性覆盖上不如JAX那么贴合。
先装依赖:
pip install "jax[tpu]" -f https://storage.googleapis.com/jax-releases/lts/lts_jax_tpu.html装完之后你可以写一个最简的训练循环,目标是用TPU的MXU单元训练一个线性模型。核心逻辑是先定义损失函数,再用jit编译成XLA图,最后用grad做自动微分:
import jax import jax.numpy as jnp from jax import grad, jit # 假设一个简单的回归任务:y = wx + b x = jnp.linspace(-1, 1, 1024) y = 3.0 * x + 0.5 def loss(params): w, b = params pred = w * x + b return jnp.mean((pred - y) ** 2) @jit def train_step(params, lr=0.01): g = grad(loss)(params) return [p - lr * g_i for p, g_i in zip(params, g)] w = jnp.array(0.0) b = jnp.array(0.0) params = [w, b] for i in range(1000): params = train_step(params) if i % 100 == 0: print(i, loss(params))这段代码的真实价值不是教你写线性回归,而是让你体会TPU上训练的两个关键感受:第一,第一次执行的时候会很慢,因为XLA要把整个计算图编译成TPU能跑的机器码,这个编译过程可能花几十秒甚至几分钟,之后重复执行就快了;第二,如果你把jit去掉,每一步都在CPU上解释执行,你可能会误判芯片性能很弱,实际上TPU没有即时解释执行的能力,几乎所有高效计算都必须先过编译。
如果想看你的张量到底跑在哪台设备上,可以用jax.devices()打印设备列表,如果显示的是多个TPU核心设备,说明PJRT已经正确识别了第八代芯片。
4.3 用vLLM类工具做推理
推理侧现在的生态比训练侧成熟不少,vLLM这类推理引擎已经能直接跑在TPU上。相比GPU环境,TPU上跑推理服务有几个明显的好处,主要在于长上下文场景下KV Cache的存放不再那么捉襟见肘,HBM容量翻倍让单机并发数上得去。
以vLLM为例,一个最小可用启动命令形态大概是这样:
vllm serve meta-llama/Llama-3-8B \ --device tpu \ --tpu-chip-type trillium \ --max-model-len 8192这里--device tpu让vLLM走TPU后端,--tpu-chip-type按你的芯片代际填写。实际跑的时候需要注意几个点:vLLM对TPU的支持版本更新很快,先查清楚你的vLLM版本跟TPU运行时版本是否匹配;量化方面,建议优先用BF16而非INT8,虽然INT8推理更快,但部分算子在小模型上的精度损失会被放大;并发参数不要一上来就拉满,先按HBM容量的一半估一下能放多少条并发流,再逐步压测。
我做推理压测时习惯先跑一个长序列请求,比如让模型生成2048个token,观察P99时延。如果发现时延曲线在某个并发数附近突然陡增,多半是KV Cache把HBM占满了,触发了重新计算或排队,这种时候下调max-num-seqs比调大max-model-len更有效。
4.4 多芯片并行时的三个关键设置
单颗芯片训练小模型没意思,上多卡才是日常。TPU的多芯片并行有两种基础姿势:数据并行和模型并行。数据并行是每颗芯片拿着完整模型的副本处理不同批次数据,最后同步梯度;模型并行则是把模型的参数和计算拆到多颗芯片上,各自算一部分。
实际操作中最难的不是理解概念,而是让XLA知道你的设备拓扑。JAX里有一个模块叫jax.sharding,你可以显式地把模型参数切到多颗芯片上:
from jax.sharding import Mesh, PartitionSpec, NamedSharding devices = jax.devices() mesh = Mesh(devices, ('data',)) sharding = NamedSharding(mesh, PartitionSpec('data', None))这段代码的核心作用是告诉XLA编译时按数据并行方式排布阵列。如果你忘了做这个设置,即使TPU host上插了8颗芯片,计算也可能被XLA自动复制到单颗设备上,结果就是"看起来用了8颗,实际上只用了1颗"。
第二个关键设置是环境变量XLA_USE_BF16=1。这个变量让XLA在允许的时候把FP32计算自动降级成BF16,可以显著减少内存带宽消耗和提升计算速度,但代价是数值精度可能不够。对训练用建议只在预热或日志打印场景开,正式训练还是用混合精度策略手动控制。
第三个关键设置是选好编译缓存路径。XLA编译图非常耗时,如果不缓存,每次重启进程都要重新编译。把环境变量XLA_FLAGS=--xla_force_host_platform_device_count=8和JAX_COMPILATION_CACHE_DIR=/tmp/jax_cache配合使用,二次启动能省掉大量编译时间。
5. 常见问题与排查技巧实录
5.1 配额和资源排队问题
最常见的问题是创建TPU时提示配额不足,或者任务一直处于queued状态。这种时候先在控制台的配额页面确认TPU配额和CPU配额都申请了,配额区域也要对应上,比如us-central2的TPU配额不能用于us-east1。
还有一个很容易忽略的点:抢占式的TPU资源虽然便宜,但会被其他更高优先级任务抢走,训练跑到一半实例被释放是常有的事。如果跑的是长训练任务,建议用随预定的(on-demand)而非抢占式,多花点钱买稳定,实测下来能少很多心理折磨。
5.2 内存溢出与OOM
TPU的HBM是芯片上的固定资源,不像GPU那样可以有官方内存管理工具手动清理。碰到OOM时,第一反应不是加内存,而是减batch size,这是最立竿见影的手段。其次检查是不是张量在设备之间发生了不必要的复制,比如某些操作被XLA放在了CPU上,CPUGPU之间来回拷贝会放大内存占用。
日志里出现"XlaRuntimeError: RESOURCE_EXHAUSTED"时,除了减batch,还可以查看xla_memory_info或者用profile工具抓内存分配曲线。很多时候内存根本不是被模型参数吃掉的,而是被梯度缓冲和优化器状态吃掉的,两者的累加量通常是参数量的数倍。
5.3 算力利用率上不去的真实原因
这是TPU上最隐蔽的问题。你看到脚本在跑,监控显示设备利用率只有30%,动手排查时又不知道从哪下手。据我经验,绝大多数利用率低的原因不在芯片,而在计算图的形状不匹配。
MXU喜欢大而规整的矩阵乘法,如果你的模型里到处都是小矩阵、动态形状、或者频繁的padding,XLA把计算图上排布到MXU时会产生大量空转。解决办法有三个方向:尽量用静态shape,避免在训练循环里出现Python原生的动态分支;把小的矩阵拼接成大的批次再做乘法;检查算子的布局,比如把NCHW换成NHWC有时能明显提升MXU的利用效率。
还有一个容易踩的坑是:TPU上跑PyTorch的某些自定义算子时,如果这个算子没有对应的XLA实现,它会退回CPU执行,一次前向传播在TPU和CPU之间来回穿梭,利用率自然上不去。排查办法是在日志里搜"fallback"或者"CPU"关键字,发现了就换成内置算子或者用JAX重写这一层。
5.4 首次运行慢与编译缓存
很多人在TPU上的第一个反应是"这也太慢了",然后就开始怀疑芯片性能。实际上第一次跑慢几乎都是XLA编译在做功。我在前面提到的JAX_COMPILATION_CACHE_DIR就是用来解决这个问题的。
另外,训练脚本里如果每次迭代都重新执行jax.jit的函数定义,等于每次都在触发重新编译。正确做法是把训练step函数定义在循环外面,编译一次,循环里只调用。一个完整的step函数在第八代TPU上的编译时间通常在几十秒到几分钟不等,如果你把它放进循环里,训练节奏就完全废了。
下面把几个高频问题的处理思路整理成速查表:
| 现象 | 可能原因 | 处理路径 |
|---|---|---|
| 创建实例报配额不足 | 区域配额不匹配或未申请 | 控制台核对区域和配额项,补充申请 |
| 训练到一半实例消失 | 使用的是抢占式资源 | 改随预定资源,或做定时checkpoint |
| 日志报RESOURCE_EXHAUSTED | HBM被参数和优化器状态占满 | 减batch size,检查张量是否被复制到CPU |
| 设备利用率长期低于50% | 形状不规整或算子回退CPU | 静态shape、整批矩阵乘、查fallback日志 |
| 每轮迭代执行都很慢 | XLA在反复编译 | step函数定义移出循环,启用编译缓存目录 |
| 多个设备只有1个在算 | 未配置sharding策略 | 用jax.sharding.Mesh显式配置数据并行 |
5.5 与GPU习惯的迁移差异
如果你是从GPU生态迁移过来的,有几个习惯要改。第一,GPU时代你有nvidia-smi可以看显存和算力占用,TPU上没有完全对应的工具,建议用Cloud的监控面板或tpu命令行工具查看设备指标。第二,PyTorch代码在CUDA上未必能直接跑在TPU上,尤其是用了torch.cuda专属API的代码,迁移时需要改成PJRT抽象的设备接口。网上看到有人直接把YOLO、Mask2Former这些CUDA原生的训练脚本往TPU上搬,几乎没有不改算子直接跑成的,通常都要过一层XLA适配或改写Dataloader。
6. 选型建议和掏心窝的几句话
文章写到这里,参数、实操、排查都说完了,最后聊几句选型上的真实判断。
如果你问我TPU和GPU怎么选,我个人经验是:如果你的计算负载能被XLA良好编译,且训练规模需要数百张卡级的集群同步,TPU的性价比确实很夸张。第八代把互连和内存带宽都抬了一档,大规模训练时的通信瓶颈被压得很低,这在GPU集群上往往要花额外精力调NVLink和IB拓扑才能达到类似效果。反过来,如果你的场景偏CV、多模态、或者依赖大量自定义算子,GPU的生态还是明显更舒服,TPU上移植算子的成本可能会抵消算力优势。
另一条经验是关于"迁移成本"的。TPU的学习曲线不在硬件,而在软件模式。JAX的函数式编程风格、XLA的编译思维、sharding的显式编排,这三样东西熟练之后,TPU的表现比多数人预期的更惊喜;不熟练的时候,你会在各种"为什么这么简单的事情在TPU上这么绕"的疑问里消磨掉耐心。建议不要直接拿生产任务练手,先在JAX里把一个中等模型完整跑通,再谈上生产。
最后说一个我自己踩过几次坑后的体会:Google发布的数据永远是相对值优先,绝对值需要你自己折算,折算时务必带上场景。同一个"提升4.7倍",在纯稠密大矩阵训练里和稀疏长文本推理里,体验是完全不一样的。如果你是做长上下文推理的,HBM带宽翻倍那一条带来的实际收益,可能比算力提升4.7倍还要明显;如果你是做极致密计算的大模型预训练,那算力和集群规模才是真正的决定因素。把参数放到自己的场景里去解读,而不是对着纸质数字空谈强弱,这才是这份"详版参数"最大的价值。