ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

第八代TPU(Trillium)参数详解与训练推理实践

第八代TPU(Trillium)参数详解与训练推理实践 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喜欢给精确到个位的TFLOPsGoogle更喜欢给相对上代提升X倍这种相对值。Trillium官方公布的三个核心数字是训练性能比v5e提升4.7倍、HBM容量和带宽比v5e翻倍、能效比提升67%。这三个数字单独看都不难理解但想把它换算成单颗芯片到底几个TFLOPs就得先把v5e的绝对值找出来再乘这也是网上各种参数表数字对不上的原因。3. 第八代TPUTrillium口径详版参数解析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 TFLOPSFP16稠密或者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参照第八代TPUTrillium口径发布时间202320232024官方定位训练推理均衡训练优先训练推理全面升级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 \ --zoneus-central2-b \ --accelerator-typev5p-8 \ --versiontpu-vm-tf-2.17.0-pod \ --projectyour-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, lr0.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_BF161。这个变量让XLA在允许的时候把FP32计算自动降级成BF16可以显著减少内存带宽消耗和提升计算速度但代价是数值精度可能不够。对训练用建议只在预热或日志打印场景开正式训练还是用混合精度策略手动控制。第三个关键设置是选好编译缓存路径。XLA编译图非常耗时如果不缓存每次重启进程都要重新编译。把环境变量XLA_FLAGS--xla_force_host_platform_device_count8和JAX_COMPILATION_CACHE_DIR/tmp/jax_cache配合使用二次启动能省掉大量编译时间。5. 常见问题与排查技巧实录5.1 配额和资源排队问题最常见的问题是创建TPU时提示配额不足或者任务一直处于queued状态。这种时候先在控制台的配额页面确认TPU配额和CPU配额都申请了配额区域也要对应上比如us-central2的TPU配额不能用于us-east1。还有一个很容易忽略的点抢占式的TPU资源虽然便宜但会被其他更高优先级任务抢走训练跑到一半实例被释放是常有的事。如果跑的是长训练任务建议用随预定的on-demand而非抢占式多花点钱买稳定实测下来能少很多心理折磨。5.2 内存溢出与OOMTPU的HBM是芯片上的固定资源不像GPU那样可以有官方内存管理工具手动清理。碰到OOM时第一反应不是加内存而是减batch size这是最立竿见影的手段。其次检查是不是张量在设备之间发生了不必要的复制比如某些操作被XLA放在了CPU上CPUGPU之间来回拷贝会放大内存占用。日志里出现XlaRuntimeError: RESOURCE_EXHAUSTED时除了减batch还可以查看xla_memory_info或者用profile工具抓内存分配曲线。很多时候内存根本不是被模型参数吃掉的而是被梯度缓冲和优化器状态吃掉的两者的累加量通常是参数量的数倍。5.3 算力利用率上不去的真实原因这是TPU上最隐蔽的问题。你看到脚本在跑监控显示设备利用率只有30%动手排查时又不知道从哪下手。据我经验绝大多数利用率低的原因不在芯片而在计算图的形状不匹配。MXU喜欢大而规整的矩阵乘法如果你的模型里到处都是小矩阵、动态形状、或者频繁的paddingXLA把计算图上排布到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_EXHAUSTEDHBM被参数和优化器状态占满减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倍还要明显如果你是做极致密计算的大模型预训练那算力和集群规模才是真正的决定因素。把参数放到自己的场景里去解读而不是对着纸质数字空谈强弱这才是这份详版参数最大的价值。
返回列表