Kimi K3 移植到 Google TPU 生态的可行性分析¶
日期:2026-07-28 | 目标硬件:TPU v7x (Ironwood) 主选,v6e (Trillium) 次选 | 范围:推理移植 + 服务化成本
1. 执行摘要¶
Kimi K3 是 2.8T 总参 / 104B 激活的混合注意力 MoE 模型,93 层中 69 层是 KDA(Kimi Delta Attention,线性注意力),24 层是 NoPE 版 Gated MLA,MoE 部分为 896 experts top-16 的 LatentMoE 结构,权重以 MXFP4 做过 QAT。
结论:技术上可行,且 K3 的架构特征相当适配 TPU。核心工作量比预想的小 —— torchtpu-vllm 已经把大部分基础设施做完了,真正要写的只有一处。
五条主要判断:
-
显存不是瓶颈,拓扑下限是 16 chips。 权重实测算下来 1.553 TB。注意 MXFP4 只覆盖 routed experts(量化配置的
ignore列表排除了 attention、shared experts、dense MLP、lm_head、vision),所以是「1.449 TB MXFP4 + 104 GB BF16」的混合账。v7x 单 chip 192 GB,2×2×2 拓扑(8 chips,1.536 TB)装不下,最小可行是 16 chips。 -
长上下文是 K3 在 TPU 上最大的优势。 只有 24 层全注意力,采用 absorbed MLA 后 1M context 的 KV cache 仅 14.5 GB/序列。作为对照,DeepSeek-V3 的 61 层 MLA 在同等上下文下需要约 36 GB。K3 更长的上下文反而更省。
-
短上下文下瓶颈会反转。 KDA 的循环状态是 434 MB/序列且与序列长度无关。在约 32.5K token 以下,它比 KV cache 还占显存,直接压制并发批量。这是 TPU 侧容量规划最容易被忽略的一项。
-
KDA 解码不再是阻塞项,真正的缺口只剩一处衰减门粒度。 MaxText 的
attention_kda.py在自回归模式下确实直接抛NotImplementedError,但torchtpu-vllm里已经有一个融合的逐通道单步解码 Pallas kernel(kernels/gdn/v2/gdn_decode_kernel.py,带 state 索引、L2norm、门控全部内联)。剩下的缺口是它的分块 prefill 路径只做到 per-head 标量衰减,而 K3 的f_b_proj输出 reshape 成(h d)是逐通道的。把tops.chunk_kda已实现的逐通道衰减移植过去,是本项目唯一必须自己写的一块,估 8–14 天。 -
EP 的通信开销不在带宽,在延迟。 v7x 每 chip 1200 GB/s 双向 ICI,92 层 MoE 的 all-to-all 在 decode 时只占 HBM 时间的 2.1%、prefill 时只占 ICI 容量的 5.8%,余量一个数量级。但每步要发 184 次 collective、每次只搬 1.79 KB —— 全是固定开销。风险从带宽转移到了小消息延迟。
技术栈选择上,建议以 torchtpu-vllm 为主线(工作量 30–52 天,对比纯 JAX 路线的 56–87 天)。这里要澄清一个容易搞错的前提:torch_tpu 本身只是 PyTorch 后端,不是服务栈;服务栈是 torchtpu-vllm,而它的 kernel 全部是 JAX/Pallas 写的,通过 torch_tpu._internal.pallas.jax_op 桥接成 torch 算子。所以「JAX 路线 vs PyTorch 路线」是个伪命题 —— 两条路线共用同一个 kernel 层,分歧只在上层的模型定义与调度框架。tpu-raiden 的 KV cache 传输层同时提供 JAX 和 torch 两套 binding,同样是共享的。详见第 3、7 章。
方法论与局限
本报告为静态分析:依据 K3 官方 config.json 与 HF modeling 源码,逐一对照 maxtext / pallas-kernel(tops) / torch_tpu / torchtpu-vllm / tpu-raiden 五个仓库的实际代码,配合手工 roofline 推算。未在真机验证,所有性能数字均为分析值并标注了假设。
参数模型已做交叉校验:按 config 逐层重算得总参 2.7795 T、激活 104.19 B,与官方公布的 2.8T / 104B 吻合,说明结构建模正确,后续显存与算力推算建立其上。
2. K3 架构解剖¶
以下均从 config.json 与 modeling_kimi_linear.py 直接读出,不引用 model card 的概述值。
2.1 层结构:69 KDA + 24 MLA 的严格交错¶
linear_attn_config.full_attn_layers 明确列出全注意力层为 [4, 8, 12, ..., 88, 92, 93],其余 69 层为 KDA。规律是每 4 层插入 1 层全注意力,末尾 92、93 连续两层。第 1 层是 dense MLP(first_k_dense_replace: 1),其余 92 层为 MoE。
这个 3:1 的比例是 K3 长上下文经济性的根源 —— 只有 24 层需要维护随长度增长的 KV cache。
2.2 KDA:低秩衰减门 + 全秩输出门¶
单层结构(KimiDeltaAttention):
| 组件 | 形状 | 说明 |
|---|---|---|
q_proj/k_proj/v_proj |
7168 → 12288 | 96 heads × 128 |
q/k/v_conv1d |
kernel=4, depthwise | SiLU 激活的短因果卷积 |
f_a_proj → f_b_proj |
7168 → 128 → 12288 | 低秩衰减门(rank=128) |
b_proj |
7168 → 96 | delta rule 的 β(每头标量) |
g_proj |
7168 → 12288 | 全秩输出门(use_full_rank_gate: true) |
o_norm |
FusedRMSNormGated | sigmoid 门控 |
o_proj |
12288 → 7168 |
合计 443.7 M 参数/层,69 层共 30.62 B。
递推式本身是标准的 gated delta rule:
与 MaxText 现有实现的结构性差异
MaxText 的 attention_kda.py 用全秩 g_proj(7168→12288)生成衰减门,而 K3 用的是低秩 f_a_proj(7168→128) → f_b_proj(128→12288)。两者参数量差 86 M/层,权重不能直接映射,需要改写投影层结构。这是一个容易在权重转换阶段才暴露的坑。
2.3 MLA:NoPE 变体,带输出门¶
modeling_kimi_linear.py 里 KimiMLAAttention.__init__ 有一句硬断言 assert self.use_nope,且 self.rotary_emb = None —— K3 的全注意力层完全不用 RoPE。
qk_rope_head_dim: 64 这个字段仍在,但它退化成一段不做旋转的 MQA 共享尾维:k_rot 从 kv_a_proj_with_mqa 切出后直接 expand 到所有头。位置信息全部由 KDA 层的短卷积与递推结构隐式承载。
其余维度沿用 DeepSeek-V3 风格:q_lora_rank=1536、kv_lora_rank=512、qk_nope_head_dim=128、v_head_dim=128。额外多了 mla_use_output_gate: true 带来的 g_proj(7168→12288)。单层 232.2 M 参数,24 层共 5.57 B。
HF 参考实现的 KV cache 不可用于生产
KimiDynamicCache.update() 缓存的是物化后的 key_states [B, 96, T, 192] 与 value_states [B, 96, T, 128],合计 737,280 元素/token。1M context 下这是 1.41 TB/序列,完全不可行。
必须实现 absorbed MLA:把 kv_b_proj 吸收进 q 侧与 o 侧,只缓存压缩 latent kv_lora_rank + qk_rope_head_dim = 576 元素/token。差距 1280 倍。这不是优化,是可行性前提。
2.4 LatentMoE:把专家搬到低维空间¶
KimiSparseMoeBlock 的三段式结构是 K3 相对 K2 的主要改动:
hidden 7168
→ routed_expert_down_proj → latent 3584
→ [896 experts, 每个 3584 → 3072 → 3584] top-16
→ routed_expert_norm (RMSNorm)
→ routed_expert_up_proj → hidden 7168
+ shared_experts(hidden 7168 → 6144 → 7168) # 旁路,走原始 hidden
单个 routed expert = 3 × 3584 × 3072 = 33.03 M 参数,896 个共 29.6 B/层,92 层共 2.72 T —— 占全模型 98%。
路由是 noaux_tc + sigmoid + grouped topk(num_expert_group: 1、topk_group: 1,即实际未分组),routed_scaling_factor: 1.0,moe_renormalize: true。
LatentMoE 对专家并行是利好
专家在 3584 维而非 7168 维上工作,意味着 EP 场景下 all-to-all 的载荷直接减半。考虑到 92 层 MoE 每层两次 a2a,这个设计对分布式推理的通信开销影响很大。第 6.3 节会量化。
2.5 量化:MXFP4 的实际覆盖面¶
quantization_config 关键字段:
{ "format": "mxfp4-pack-quantized", "quant_method": "compressed-tensors",
"weights": { "num_bits": 4, "group_size": 32, "symmetric": true,
"strategy": "group", "scale_dtype": "torch.uint8" },
"ignore": ["re:.*self_attn.*", "re:.*shared_experts.*",
"re:.*mlp\\.(gate|up|gate_up|down)_proj.*",
"re:.*lm_head.*", "re:.*vision_tower.*", "re:.*mm_projector.*"] }
ignore 列表是本报告最关键的单条发现。 它意味着 MXFP4 只作用于 routed experts(以及未被 regex 命中的 latent 投影),attention、shared experts、dense MLP、lm_head、视觉塔全部保持 BF16。
scale_dtype: uint8 + group_size: 32 是标准 MX 规格:每 32 个 4-bit 元素共享一个 E8M0 指数,等效 4.25 bit/元素。
2.6 其余组件¶
- SiTU 激活:
beta·tanh(gate/beta)·sigmoid(gate)·up,beta=4.0。纯 element-wise,TPU 上无障碍,但需注意tanh与sigmoid走 EUP 流水线。 - AttnRes:
attn_res_block_size: 12,以 12 层为一块做跨块残差累积(_apply_attn_res用 prefix_sum + 投影 + norm)。改变了残差图的拓扑,会影响 remat/scan 的切分策略。 - MoonViT-V2:27 层、hidden 1024、12 头、patch 14、
merge_kernel_size [2,2]、patchmergerv2投影到 7168。约 400 M 参数,BF16 不量化。
3. TPU 侧现有能力盘点¶
本章逐一核对了以下仓库的实际代码(均为 2026-07 的状态):
| 仓库 | 内容 | 备注 |
|---|---|---|
primatrix/pallas-kernel |
KDA/MLA 的 Pallas kernel,包名 tops |
MaxText 通过 pyproject.toml pin 在 v0.5.4-rc2 |
primatrix/maxtext |
JAX 侧层实现(attention_kda.py / attention_mla.py / moe.py) |
与上一行是配套的一对 |
google-pytorch/torch_tpu |
PyTorch 的 TPU 后端(ATen + XLA + torch.compile) |
底座,不是服务栈 |
google-pytorch/torchtpu-vllm |
vLLM 的 TPU platform plugin,包名 vllm_torchtpu |
PyTorch 侧真正的服务栈 |
google/tpu-raiden |
C++ 实现的 KV cache 传输与持久化层 | JAX / torch 双绑定 |
google-pytorch/torchtitan |
PyTorch 侧训练栈 | 本报告范围外 |
内部 k25-jax |
K2.5 在 v7x 上的移植记录 | 实战约束与踩坑来源 |
google-pytorch/sglang-torchtpu 已于 2026-05 迁出到 sgl-project-dev,不建议作为路径。
3.1 pallas-kernel (tops) —— KDA kernel 的真正所在¶
MaxText 只是薄封装,真正的 Pallas 实现在 primatrix/pallas-kernel。包名是 tops,MaxText 的 pyproject.toml 里以 tops[tpu] @ git+https://github.com/primatrix/[email protected] 的形式引入。
tops/ops/kda/ 的算子清单:
| 文件 | 用途 | 对推理的价值 |
|---|---|---|
chunk_fwd.py / chunk_intra_fwd_fused.py |
分块前向 | prefill 可直接复用 |
chunk_bwd.py / chunk_intra.py |
分块反向 | 训练用,推理不需要 |
fused_recurrent.py |
逐 token 递推 | decode 路径,但只是 lax.scan |
wy_fast.py / gate.py |
WY 表示、门变换 | 支撑算子 |
naive.py + tops/cpu/ops/kda/ |
CPU 高精度参考 | 数值验证基线,很有用 |
前向已高度优化并有完整的 v6e roofline 分析文档。以 chunk_kda_fwd_pipeline_perf.md(bf16 varlen, H=16, T=8192, K=V=128, BT=64)为例:
- 判定 HBM 带宽瓶颈:DMA 约 590 MB → 369 µs,而 MXU 仅需 45 µs
- 列出的一项主要优化杠杆是
vand指令做的 bf16→f32 转换(235,648 条)
这个"KDA 前向是带宽瓶颈而非算力瓶颈"的结论对 v7x 是好消息:v7x 相对 v6e 带宽提升 4.5 倍(7.38 vs 1.638 TB/s),算力提升 2.5 倍,带宽提升幅度更大。
MLA 方面,tops 里只有 CPU 参考实现(tops/cpu/ops/mla/mla.py),没有 TPU Pallas kernel —— MLA 在 MaxText 侧走的是 splash attention / paged attention 通路。
chunk_size 硬编码为 64
attention_kda.py 里的 _TOPS_CHUNK_SIZE = 64 附注"tops chunk_kda kernel only supports chunk_size=64"。序列长度不是 64 的倍数时需 padding。v7x 的 VMEM 与 v6e 不同(hardware.py 记 32 MB/core vs v6e 64 MB/core),chunk/sub-chunk 的分块参数很可能需要重调。
3.2 MaxText —— 层实现,训练完备、推理缺口明确¶
src/maxtext/layers/ 相关文件:
attention_kda.py(634 行):KimiDeltaAttention完整实现,含 CP(all-gather 与 a2a/Ulysses 两种策略)、varlen packing、shard_map 包装。attention_mla.py(1299 行):MLA类推理支持完备 ——MlaKVCache、prefill/autoregressive 双模式、paged attention、mla_naive_kvcache开关区分物化与压缩缓存。moe.py(2773 行):ragged_all_to_all专家分发、megablox分组矩阵乘、ring-of-experts、attn_dp_expert轴(DP attention + EP 组合)。configs/models/kimi-k2-1t.yml:K2 的完整配置已存在,可作为 K3 配置的起点。
关键缺口(attention_kda.py:369):
if model_mode == MODEL_MODE_AUTOREGRESSIVE:
raise NotImplementedError("KDA autoregressive mode not yet implemented.")
即 KDA 层完全没有解码路径。同时该层也没有循环状态的 cache 管理(无 init_kv_caches 的对应物)。
3.3 torch_tpu —— 后端,不是服务栈¶
google-pytorch/torch_tpu 是 PyTorch 的 TPU 后端(ATen kernel + XLA 编译缓存 + torch.compile 集成)。
已具备: 完整 ATen 算子覆盖;FP8 全系 dtype 含 torch.float8_e8m0fnu(MX 格式的 scale 类型);ragged_dot(经 tokamax);allgather/allreduce/alltoall/reduce_scatter;Splash/Flash attention;DP/FSDP/TP 示例。
关键的一项:torch_tpu._internal.pallas.jax_op —— 把 JAX/Pallas kernel 包装成 torch custom op 的桥。这是理解整个技术路径的钥匙,见 3.4。
单看这个仓库没有服务化能力(无 paged attention、无 EP、无 MLA、最大推理示例是 Qwen3-0.6B)。这容易让人低估 PyTorch 侧的成熟度 —— 实际上服务栈不在这个仓库里,torch_tpu 是底座,不是终点。
3.4 torchtpu-vllm —— 真正的 PyTorch 侧服务栈¶
google-pytorch/torchtpu-vllm(包名 vllm_torchtpu)是 vLLM 的 TPU platform plugin,基于 vLLM v0.22.1,开发非常活跃(最近 commit 在 2026-07)。它的能力覆盖是本报告 Gap 判断的主要依据。
现有 kernel 与层:
| 路径 | 内容 | 对 K3 的意义 |
|---|---|---|
kernels/gdn/v2/gdn_decode_kernel.py |
GDN 单步解码 Pallas kernel | KDA 解码路径的核心,已现成 |
kernels/gdn/v3/ |
conv1d + GDN 融合的分块 kernel | prefill 路径 |
kernels/causal_conv1d/ |
因果短卷积 | KDA 的 k=4 短卷积 |
kernels/mla/v1/kernel.py |
MLA ragged paged attention(50 KB) | absorbed MLA + 分页 + prefill/decode 混批 |
kernels/fused_moe/v1/、megablox/gmm_v2.py |
分组矩阵乘 MoE | LatentMoE |
layers/vllm/quantization/mxfp4.py |
MXFP4 量化方法 | 直接对应 K3 的权重格式 |
layers/common/ragged_gated_delta_rule_*.py |
gated delta rule 的 chunked / ref / wrapper 三件套 | KDA 算法主体 |
runner/mamba_apc.py、custom_ops/mamba_state_copy_op.py |
循环状态的 prefix caching 与 slot 搬移 | KDA 状态的 cache 管理 |
distributed/kv_transfer/、offload/cpu_tpu.py |
KV 传输、host offload、P/D 分离 | 见 3.5 |
spec_decode/、sample/rejection_sampler.py |
投机解码 | 缓解 KDA 解码延迟 |
决定性的一点:这些 kernel 是 JAX/Pallas 写的,通过 pallas.jax_op 暴露成 torch op。
# layers/vllm/custom_ops/gdn_attention_op.py
import jax
from torch_tpu._internal import pallas
...
gdn_jax_op = pallas.jax_op(op_name, ...)
layers/common/gdn_attention.py 开头就是 import jax.numpy as jnp 和 from jax.sharding import PartitionSpec as P。所以「JAX 路线 vs PyTorch 路线」是个伪命题 —— torchtpu-vllm 的模型定义来自 vLLM 上游的 PyTorch 实现,kernel 却是 JAX/Pallas,两者在 jax_op 这一层缝合。第 7 章会展开这一点。
模型定义本身不在这个仓库(models/ 只有一个 wrapper context),走的是 vLLM 上游的 ModelRegistry。意味着只要 vLLM 上游支持某个模型、且所需 kernel 在这里有覆盖,就能跑 —— 不需要像 JAX 路线那样重写模型。
3.5 tpu-raiden —— KV cache 传输层¶
google/tpu-raiden(公开仓库,活跃开发中,README 明确标注"not yet recommended for general use")是 C++/Bazel 实现的 KV cache 传输与持久化库,同时提供 JAX 和 PyTorch 绑定(tpu_raiden/api/{jax,torch})。
能力:kv_cache/(含 global_registry 与 pool reshard)、transport/(block transport、H2H strided)、weight_sync/(权重热同步,面向 RL rollout)、POSIX 共享内存持久化(/dev/shm,serving 进程重启后 KV cache 不冷启动)。
torchtpu-vllm 已经把它接进来了 —— distributed/kv_transfer/v2/raiden_pool_manifest.py、offload/cpu_tpu.py、platforms/tpu_platform.py 都引用了 raiden,examples/disagg/ 下有完整的 P/D 分离多机脚本。
对 K3 的价值:1M context 下每序列 14.5 GB KV,把冷序列 offload 到 host DRAM(v7x 每 host 约 944 GB)能大幅提升并发上限;P/D 分离则让 prefill(算力瓶颈)和 decode(带宽瓶颈)各自用最优拓扑 —— 对 K3 这种两个阶段瓶颈截然不同的模型收益尤其大。
3.6 tpu-inference / vLLM-TPU(JAX 路线)—— 已有实战积累¶
内部 K2.5 移植项目(k25-jax)在 v7x 上跑 DeepSeek-R1 与 K2-Thinking 已积累了成体系的经验,其中直接适用于 K3 的约束:
| 约束 | 内容 |
|---|---|
| 量化格式 | tpu-inference 只支持 block-wise FP8(weight_block_size=[128,128]);per-tensor FP8 会 KeyError,INT4 compressed-tensors 不支持 |
| 权重加载 | 必须用 runai_streamer;默认 safetensors loader 与 tpu_streaming_loader 都会在 weight processing 阶段 OOM |
| 多机后端 | Ray(TPU_MULTIHOST_BACKEND=ray);Pathways 因 ALTS handshake 失败不可用;必须 --no-async-scheduling |
| GCS FUSE | 必须 file-cache:max-size-mb:0,否则 boot disk 被填满触发 eviction |
K3 的 MXFP4 与 tpu-inference 的 FP8 要求正面冲突
上表第一条意味着 K3 的 MXFP4 权重无法直接喂给当前 tpu-inference。要么先离线转换成 block-wise FP8(放弃 MXFP4 的显存优势,权重从 1.553 TB 涨到约 2.83 TB),要么在 tpu-inference 中新增 MXFP4 反量化通路。这是一个需要尽早决策的岔路口。
同时该项目记录了一个尚未解决的阻塞:GKE v7x 节点上 backbone XLA 编译会 crash,疑似 Docker 镜像 AOT 编译假设的 CPU feature(+prefer-no-scatter)与 GKE 节点 CPU 不匹配。K3 上线前需确认此问题已修复。
4. Gap 分析¶
按严重度排序。工作量为单人天估算,假设熟悉 JAX/Pallas。
每个 Gap 都分别给出两条技术路线下的严重度 —— 因为同一个缺口在两条路线上的代价差别很大,这也是第 7 章做路线选型的依据。
G1 — KDA 解码路径 | JAX 路线:阻塞 | torchtpu-vllm 路线:低¶
JAX 路线(maxtext + tops):三层都要补 —— tops 的 fused_recurrent_kda 只是 lax.scan(每步把 434 MB/序列状态读出写回,无法融合门变换与 delta 修正);没有状态 cache 抽象;attention_kda.py:369 直接抛 NotImplementedError。20–30 天。
torchtpu-vllm 路线:kernels/gdn/v2/gdn_decode_kernel.py 已经是一个完整的融合单步解码 Pallas kernel,而且参数契约与 K3 的 KDA 高度吻合:
def ...(q, k, v, g, # g: Per-key gating [T, H_v, K] ← 逐通道
initial_state, # [num_states, H_v, K, V] 状态 cache
state_indices, # i32[max_num_req] slot 索引
distribution, # i32[3] (decode_end, prefill_end, mixed_end)
b, *, scale,
use_qk_l2norm_in_kernel=False, # ← K3 的 Q/K L2 norm
use_gate_in_kernel=False, # ← 门变换融合
A_log=None, dt_bias=None, lower_bound=None, apply_silu=False)
这几个 flag(use_qk_l2norm_in_kernel / use_gate_in_kernel / lower_bound)与 tops.chunk_kda 的签名几乎逐字对应,两边显然是相互知情地开发的。配套的状态管理也齐了:mamba_state_copy_op.py 做 slot 搬移,runner/mamba_apc.py 做状态的 prefix caching,distribution 参数原生支持 prefill/decode 混批。
剩下的缺口只有一处,见 G2。
G2 — 分块 prefill 路径的衰减门只到 per-head | 本项目的核心缺口 | 8–14 天¶
这是把 K3 跑起来真正要解决的问题。对比四个实现的衰减门粒度:
| 实现 | 衰减门形状 | 粒度 |
|---|---|---|
tops chunk_kda(prefill) |
gk: [B, T, H, K] |
逐通道 ✅ |
torchtpu-vllm gdn/v2 decode |
g: [T, H_v, K] |
逐通道 ✅ |
| torchtpu-vllm chunked prefill | g_cumsum 后 jnp.exp(g_cumsum)[..., None] |
逐头标量 ❌ |
torchtpu-vllm gdn/v3 |
gating_log: [1, 1, num_v_heads] |
逐头标量 ❌ |
K3 的 KDA 是 g = f_b_proj(f_a_proj(h)) 再 rearrange('... (h d) -> ... h d', d=head_dim),即每个头内部逐通道的对角衰减矩阵(GLA/DeltaNet 风格)。而 torchtpu-vllm 的分块 prefill 路径是为 Qwen3-Next 的 GDN 设计的,衰减是每头一个标量(Mamba2 风格)—— 代码里 k_beta * jnp.exp(g_cumsum)[..., None] 的那个 [..., None] 广播就是证据。
好消息是:分块算法的骨架完全一样(log 空间 cumsum → chunk 内衰减差矩阵 → WY 表示 → 状态传递),要改的只是把衰减的秩从标量提到向量。而且 tops.chunk_kda 已经把逐通道版本做出来并优化过了,两边都是 JAX/Pallas,数学可以直接移植。
所以两个仓库的能力恰好互补,落点很清楚:tops 的逐通道分块 prefill + torchtpu-vllm 的逐通道解码 kernel + torchtpu-vllm 的服务栈。
G3 — MLA absorbed + 分页 | JAX:中 | torchtpu-vllm:低¶
kernels/mla/v1/kernel.py(50 KB)已实现 mla_ragged_paged_attention,文件头写明"TPU-Friendly and Data-Movement-Friendly MLA Ragged Paged Attention kernel",支持 prefill/decode 混批,配套 get_kv_cache_shape / update_kv_cache / ref_mla_ragged_paged_attention(参考实现,可做数值对拍)。
剩余工作只是确认它在 K3 的 NoPE + output gate 变体下正确 —— qk_rope_head_dim=64 那段不做旋转的 MQA 共享尾维需要核对。3–5 天。
G4 — MXFP4 | JAX:高 | torchtpu-vllm:低¶
layers/vllm/quantization/mxfp4.py 提供 VllmMxfp4Config / VllmMxfp4MoEMethod,注册进 vLLM 的量化系统,含 process_weights_after_loading 与 monolithic TPU forward 路径。底层 torch_tpu 有 float8_e8m0fnu dtype,kernels/quantized_matmul/blockwise_kernel.py 提供 block-wise 量化矩阵乘。
在这条路线上,剩下的工作只是验证现有实现能否吃 K3 的 compressed-tensors / mxfp4-pack-quantized 打包格式,权重按 1.553 TB 规划即可。5–8 天。
JAX 路线上这一项是硬冲突
第 3.6 节的约束 —— tpu-inference 只支持 block-wise FP8(weight_block_size=[128,128]),per-tensor FP8 会 KeyError,INT4 compressed-tensors 不支持 —— 意味着 K3 的 MXFP4 权重无法直接喂进去。要么离线转成 block-wise FP8(权重从 1.553 TB 涨到约 2.83 TB,拓扑下限从 16 chips 变成 32),要么自己在 tpu-inference 里新增 MXFP4 反量化通路(10–15 天)。
这是两条路线之间最大的一处实质差异,且它影响的不只是工期,还有硬件配置的下限。
G5 — KDA 投影层结构 | JAX:中 | torchtpu-vllm:低¶
MaxText 的 attention_kda.py 用全秩 g_proj(7168→12288) 生成衰减门,K3 用低秩 f_a_proj(7168→128)→f_b_proj(128→12288),差 86 M 参数/层,权重不能直接映射,JAX 路线要改写投影层结构(3–5 天)。
torchtpu-vllm 路线下这一项基本不存在 —— 模型定义走 vLLM 上游的 PyTorch 实现,低秩投影就是两个 nn.Linear,不需要动 kernel。这是"复用 vLLM 模型定义"的直接收益。
G6 — 896 experts 的路由与 EP | 中 | 6–12 天¶
noaux_tc + sigmoid + renormalize 的路由需对齐 —— torchtpu-vllm 已有 moe_routing.py 和 grouped top-k 的测试(test_moe_routing_select_experts_grouped_topk)。LatentMoE 的 down/up 投影需插到 MoE 前后。896 experts 在 92 层上的 a2a 调度开销需实测(见 6.4)。
G7 — MoonViT-V2 | 中 | 8–13 天¶
K2.5 阶段已有约 600 行 JAX 初版(MoonViTPatchEmbed / Attention / MLP / Block / Encoder / PatchMerger),但 V2 的 merge_type: sd2_tpool、pos_emb_type: divided_fixed、视频时序 patch 需要重做。torchtpu-vllm 路线下有 vision_attention.py 与 multi_modal_inference.py 示例可依托。
G8 — SiTU 与 AttnRes | 低 | 2–4 天 / ~2 天¶
SiTU 是纯 element-wise。AttnRes 的 12 层分块残差会改变 scan/remat 的切分边界 —— JAX 路线要在 decoders.py 里调整;PyTorch 路线只是模型定义里的加法,几乎无成本。
工作量汇总¶
| Gap | JAX 路线 | torchtpu-vllm 路线 |
|---|---|---|
| G1 KDA 解码 | 20–30 天(阻塞) | 已有,0 |
| G2 逐通道分块 prefill | 已有(tops),0 | 8–14 天(核心) |
| G3 MLA absorbed + 分页 | 5–8 天 | 3–5 天 |
| G4 MXFP4 | 10–15 天,或退回 block-FP8 | 5–8 天 |
| G5 KDA 投影层结构 | 3–5 天 | ~0 |
| G6 MoE 路由 / EP | 8–12 天 | 6–10 天 |
| G7 MoonViT-V2 | 8–13 天 | 8–13 天 |
| G8 SiTU / AttnRes | 2–4 天 | ~2 天 |
| 模型定义 | 需 JAX 重写(已含在上表) | 复用 vLLM 上游 |
| 合计 | 56–87 天 | 30–52 天 |
单人估算,不含真机调试与性能调优。
两条路线的工期差距集中在两处:G1(JAX 要从零写解码 kernel + 状态 cache + 层接口三层,torchtpu-vllm 已有)和模型定义(JAX 要整套重写,torchtpu-vllm 复用 vLLM 上游)。除此之外还有一个不体现在工期里但同样重要的差异 —— G4 决定权重是 1.553 TB 还是 2.83 TB,进而决定拓扑下限是 16 chips 还是 32。
torchtpu-vllm 路线的关键路径是 G2:边界清晰、有现成参考实现(tops.chunk_kda)可移植,相比"从零写解码 kernel"这种三层缺口,风险低得多。
5. 数值估算¶
硬件参数取自 pallas-kernel/tops/hardware.py(per-device 值,每 chip 2 device):
| v7x (Ironwood) | v6e (Trillium) | |
|---|---|---|
| HBM 容量 | 192 GB/chip | 32 GB/chip |
| HBM 带宽 | 7.38 TB/s | 1.638 TB/s |
| BF16 峰值 | 2307 TFLOPS | 918 TFLOPS |
| FP8 峰值 | 4614 TFLOPS | 918 TFLOPS |
| ICI(双向/chip) | 1200 GB/s | 800 GB/s |
| DCN(双向/chip) | 100 GB/s | 100 GB/s |
ICI/DCN 数字为实测值,tops/hardware.py 未记录,第 6.3 节的通信分析基于这组数字。
Note
tops/docs/references/tpu-hardware.md 中 v7 一行写 "HBM BW ~4.0 TB/s (Approximate)",与 hardware.py 的 3690 GB/s × 2 device = 7.38 TB/s 矛盾。后者与 Ironwood 公开规格(7.37 TB/s)一致,本报告采用后者,文档那行应为过时条目。
5.1 权重显存¶
| 部分 | 参数量 | 精度 | 显存 |
|---|---|---|---|
| routed experts + latent 投影 | 2.727 T | MXFP4 (4.25 bit) | 1.449 TB |
| attention / shared / dense / embed / head | 52.0 B | BF16 | 104.0 GB |
| 合计 | 2.7795 T | 1.553 TB | |
| 对照:全 BF16 | 2.7795 T | BF16 | 5.559 TB |
| 对照:experts 转 block-FP8 | 2.7795 T | FP8+BF16 | 约 2.83 TB |
5.2 KV cache 与 KDA 状态¶
- MLA absorbed:24 层 × 576 = 13,824 元素/token → FP8 下 13.5 KB/token
- MLA 物化(HF 参考实现):737,280 元素/token → BF16 下 1.41 MB/token,1280 倍差距
- KDA 循环状态:69 层 × 96 头 × 128 × 128 = 108.5 M 元素 → FP32 下 434 MB/序列,恒定
- KDA conv 状态:7.63 M 元素 → BF16 下 15.3 MB
每序列总占用:
| 上下文 | MLA KV (FP8) | KDA 状态 (FP32) | 合计 |
|---|---|---|---|
| 4,096 | 0.057 GB | 0.449 GB | 0.51 GB |
| 32,768 | 0.453 GB | 0.449 GB | 0.90 GB |
| 131,072 | 1.812 GB | 0.449 GB | 2.26 GB |
| 524,288 | 7.248 GB | 0.449 GB | 7.70 GB |
| 1,048,576 | 14.496 GB | 0.449 GB | 14.94 GB |
交叉点在约 32,507 token。 短于此,KDA 状态主导显存;长于此,KV cache 主导。
这个结论的实践含义:K3 在 TPU 上服务短请求时,并发上限由 KDA 状态决定,且无法通过缩短上下文来提升并发。把 KDA 状态存成 BF16 可以砍半(217 MB/序列),但需评估 delta rule 递推的数值稳定性 —— tops 的 fused_recurrent_kda 内部用 FP32 累加器,暗示精度敏感。
5.3 拓扑下限¶
按权重 + 15% 系统开销(XLA 编译缓冲、激活、碎片):
| 拓扑 | chips | HBM 总量 | 权重占比 | 余量 | 1M ctx 并发 |
|---|---|---|---|---|---|
| 2×2×2 | 8 | 1.536 TB | 101.1% | −17 GB | 装不下 |
| 2×2×4 | 16 | 3.072 TB | 50.6% | 1519 GB | ~86 |
| 2×4×4 | 32 | 6.144 TB | 25.3% | 4591 GB | ~261 |
| 4×4×4 | 64 | 12.288 TB | 12.6% | 10735 GB | ~611 |
v6e(32 GB/chip)需要约 56 chips 才装得下权重,且带宽只有 v7x 的 22%。v6e 不适合 K3,仅可用于单层/子模块的算子开发验证。
5.4 Decode Roofline¶
Decode 是带宽瓶颈。每步权重搬运量:
- batch=1(只读命中的 16 个专家):130.0 GB
- 大 batch(896 专家全被触达):1553 GB(即全模型)
| 配置 | 聚合带宽 | batch=1 理论 | 按 50% 带宽效率 | 大 batch (B=256) |
|---|---|---|---|---|
| v7x × 16 | 118.1 TB/s | 1.10 ms → 908 tok/s | ~454 tok/s | 13.15 ms/step → 51 µs/token/seq |
| v7x × 32 | 236.2 TB/s | 0.55 ms → 1816 tok/s | ~908 tok/s | 6.58 ms/step → 26 µs/token/seq |
一个反直觉的发现
batch=1 的 130 GB 搬运中,attention 权重(BF16,72.4 GB)比命中的 routed experts(MXFP4,25.8 GB)还多。原因是 attention 层没有被量化,而 KDA 每层 443.7 M 参数本身就重。
这意味着:对 attention 做量化的收益,在小 batch 解码场景下超过对专家做量化。 如果要进一步优化低延迟路径,attention 权重的 FP8 化优先级应高于 MXFP4 的精细调优。
5.5 Prefill Roofline¶
每 token 算力 = 2 × 104.19 B = 208 GFLOPs(不含注意力打分)。
| 配置 | MFU 30% | MFU 45% |
|---|---|---|
| v7x × 16, BF16 | 53,144 tok/s | 79,716 tok/s |
| v7x × 16, FP8 | 106,288 tok/s | 159,431 tok/s |
| v7x × 32, BF16 | 106,288 tok/s | 159,431 tok/s |
| v7x × 32, FP8 | 212,575 tok/s | 318,863 tok/s |
5.6 长上下文下注意力算力反超¶
Absorbed MLA 在解码时每 token 的注意力算力 = 2 × 2 × 24 层 × 96 头 × 512 × ctx:
| 上下文 | 注意力算力/token | 相对 MoE 部分 |
|---|---|---|
| 32,768 | 0.155 TFLOPs | 0.7× |
| 262,144 | 1.237 TFLOPs | 5.9× |
| 1,048,576 | 4.948 TFLOPs | 23.7× |
约 45K token 之后,注意力算力超过 MoE。 1M context 下解码从带宽瓶颈彻底转为注意力算力瓶颈。
作为对照,DeepSeek-V3 有 61 层全注意力,同等上下文下这项开销是 K3 的 2.54 倍。K3 的混合架构在超长上下文推理上对 TPU 是结构性利好。
6. Sharding 策略¶
6.1 整除性约束(先看这个)¶
| 维度 | 值 | /8 | /16 | /32 | /64 |
|---|---|---|---|---|---|
| KDA/MLA heads | 96 | 12 | 6 | 3 | ✗ 1.5 |
| routed experts | 896 | 112 | 56 | 28 | 14 |
| latent dim | 3584 | 448 | 224 | 112 | 56 |
| moe intermediate | 3072 | 384 | 192 | 96 | 48 |
| hidden | 7168 | 896 | 448 | 224 | 112 |
96 头无法被 64 整除。 这直接限定了纯头并行的 TP 上限为 32。要用 64 chips,必须让 TP ≤ 32 并把剩余并行度分给 EP 或 DP。
6.2 推荐方案:DP-attention + EP¶
这也是 k25-jax 在 DeepSeek-R1 上验证过的组合(enable_dp_attention: true + --enable-expert-parallel)。MaxText 侧对应 attn_dp_expert 网格轴。
v7x × 16(2×2×4):
- 专家权重 1.449 TB / 16 = 90.6 GB/chip,均匀无冗余
- Attention 权重 104 GB,DP=2 下每组复制一份 → 每 chip 13 GB
- KV cache 按 DP 组切分,每组独立管理,避免 TP 复制
- KDA 状态按头切分:434 MB / 8 = 54 MB/chip/序列
v7x × 32(2×4×4):
推荐 32 chips 作为生产配置:显存余量充裕(权重占 25%),并发可到 261 路 1M context,且 MXFP4 若暂不可用、退回 block-FP8(2.83 TB)时仍装得下。
6.3 通信量估算:ICI 带宽不是瓶颈¶
EP 的 all-to-all 是主要通信开销。每 token 每 MoE 层:
- dispatch:latent 向量 3584 元素,发往 top-16 专家所在的 chip
- FP8 载荷:3584 B × 16 = 57.3 KB,combine 同量 → 114.7 KB/层
- 92 层 → 10.55 MB/token 往返
对照:传统 MoE 在 7168 维上做 a2a 就是 21.1 MB/token。LatentMoE 让通信量正好减半,这是 2.4 节那个架构观察的量化形式。
按 v7x 每 chip 1200 GB/s 双向(单向可用 600 GB/s)、EP 铺满全部 chip 的 per-chip egress 模型:
Decode(每步)
| 配置 | B=1 | B=32 | B=256 |
|---|---|---|---|
| v7x × 16 — ICI 时间 | 1.1 µs | 35.2 µs | 281.4 µs |
| v7x × 16 — 占 HBM 时间 | 0.1% | 0.3% | 2.1% |
| v7x × 32 — ICI 时间 | 0.5 µs | 17.6 µs | 140.7 µs |
| v7x × 32 — 占 HBM 时间 | 0.1% | 0.3% | 2.1% |
| v6e × 64 — 占 HBM 时间 | 0.0% | 0.1% | 0.7% |
Prefill(稳态)
| 配置 | 吞吐 | 每 chip ICI | 占 ICI 容量 |
|---|---|---|---|
| v7x × 16 @ MFU 30% | 53,141 tok/s | 35.0 GB/s | 5.8% |
| v7x × 32 @ MFU 30% | 106,283 tok/s | 35.0 GB/s | 5.8% |
| v7x × 32 @ MFU 45% | 159,424 tok/s | 52.6 GB/s | 8.8% |
结论:ICI 带宽有一个数量级的余量,EP 方案在带宽维度上完全成立。 即便按 3D torus 的 bisection 约束保守打 2× 折扣,decode 也只到 4%、prefill 到 18%,结论不变。
这个结论值得强调,因为 92 层 MoE + top-16 的 a2a 直觉上像是个大问题。真正的问题在下一节。
6.4 真正的约束是 collective 延迟,不是带宽¶
上表有个细节值得盯:B=1 时每次 collective 只搬 1.79 KB/chip,传输耗时 3 ns。 这个数量级下带宽完全无关,全部时间都是 collective 的固定开销。
每步要做 2 × 92 = 184 次 all-to-all。按 v7x × 32、batch=1 的 HBM 下界 550 µs 对照:
| 单次 collective 开销 | 184 次合计 | 占 HBM 时间 |
|---|---|---|
| 1 µs | 184 µs | 33% |
| 2 µs | 368 µs | 67% |
| 5 µs | 920 µs | 167% |
低延迟解码路径的真实风险
如果单次 a2a 的固定开销落在 2–5 µs,batch=1 解码会从带宽瓶颈变成通信延迟瓶颈,延迟翻倍甚至更多。92 层的深度在这里是放大器 —— 同样的 per-collective 开销,层数少一半风险就减半。
所以真正需要实测的量不是 ICI 带宽,而是「小消息 all-to-all 的 per-call 开销」。 这是个便宜很多的测量:不需要 K3,拿一个 92 层的 dummy MoE 就能测。建议提到 P0 阶段做掉。
缓解手段(按优先级):
1. 小 batch 时降低 EP 度数,专家改走 TP —— 用带宽换延迟,而带宽正好有 10× 余量
2. 让 shared experts 的计算与 routed experts 的 dispatch 重叠(shared 走 hidden 维旁路,不依赖 a2a 结果)
3. MaxText 的 use_ring_of_experts,把 a2a 换成流水化的 ring
4. 跨层融合 dispatch —— 受限于层间数据依赖,收益有限
6.5 跨 slice(DCN):EP 绝不能跨 slice¶
DCN 100 GB/s 双向(单向 50 GB/s)vs ICI 1200 GB/s,带宽差 12×。把 EP 组拆到跨 slice(最坏情形:32 chips 的 a2a 全部经 DCN):
| 每 chip 载荷 | ICI | DCN | |
|---|---|---|---|
| B=1 | 0.33 MB | 0.5 µs (0.1%) | 6.6 µs (1.2%) |
| B=32 | 10.55 MB | 17.6 µs (0.3%) | 211.0 µs (3.2%) |
| B=256 | 84.41 MB | 140.7 µs (2.1%) | 1688 µs (25.7%) |
括号内为占该 batch 下 HBM 时间的比例。Decode 在大 batch 下已经吃掉四分之一的步时间;prefill 则直接超容量:
| Prefill 稳态 | 需要 | DCN 单向容量 | |
|---|---|---|---|
| v7x × 32 @ MFU 30% | 35.0 GB/s/chip | 50 GB/s | 70% |
| v7x × 32 @ MFU 45% | 52.6 GB/s/chip | 50 GB/s | 105% ← 不可行 |
再叠加 6.4 的延迟结论 —— DCN 的 per-call 延迟比 ICI 高一个数量级以上,而 K3 的 a2a 恰好是延迟敏感型 —— 跨 slice EP 在带宽和延迟两个维度上同时失败。
硬性约束:整个 EP 组必须落在单个 slice 内。 好在 32 chips 远小于 v7x 单 slice 规模,这不构成实际限制,但它把「先按 32 chips 规划」从建议变成了要求 —— 如果显存压力迫使你扩到需要跨 slice 的规模,EP 方案就得整体重新设计。
DCN 唯一合理的用途是跨 slice 做 DP:复制整个模型实例扩展吞吐,slice 之间只有请求级交互,每步 DCN 流量约等于零。
6.6 长序列 prefill 的 Context Parallelism¶
1M context 的 prefill 单靠 batch 维切不动,需要 CP。MaxText 的 context_parallel_strategy 提供两种,且明确支持 KDA:
all_gather:需tops >= v0.5.4-rc2的CPContext,attention_kda.py里有硬断言保护 —— 缺失时拒绝运行,因为"CP 会静默破坏跨 rank 的循环状态"a2a(Ulysses):序列↔头维度互换,attention_kda.py里把 q/k/v/g 融合成单次 all-to-all
KDA 在 CP 下的难点是递推状态的跨 rank 传递,tops 有专门的设计文档(docs/design-docs/ops/kda/global_cp.aligned.zh.md、kda-chunk-bwd-cp.zh.md)。短卷积还需要 halo exchange(cp_a2a.halo_exchange_for_conv)拉取前一个 rank 的 kernel_size-1 个 token —— 否则每个 CP 边界的因果卷积都是错的。
推理场景下这部分是现成的,因为 CP 的复杂度主要在反向传播,前向已实现。
6.7 不推荐 Pipeline Parallelism¶
93 层理论上可切,但推理场景下 PP 引入的 bubble 会直接抬高 TTFT,且 AttnRes 的 12 层分块残差会与 PP 切分边界冲突。单 slice 内用 EP + TP 已足够,除非要跨 slice 扩展。
7. 技术路线选型¶
习惯上会把这件事描述成"JAX 路线 vs PyTorch 路线"的二选一。核对代码之后,这个框架并不成立。
7.1 关键事实:两条路线共用同一批 Pallas kernel¶
torchtpu-vllm 的 GDN / MLA / MoE kernel 全部是 JAX + Pallas 写的,通过 torch_tpu._internal.pallas.jax_op 包装成 torch custom op:
# src/vllm_torchtpu/layers/common/gdn_attention.py
import jax
import jax.numpy as jnp
from jax.sharding import PartitionSpec as P
from vllm_torchtpu.kernels.gdn.v3 import wrapper as gdn_v3_wrapper
所以真正的分界线不在"JAX 还是 PyTorch",而在模型定义与服务层用哪一套:
| 层 | JAX 路线 | torchtpu-vllm 路线 |
|---|---|---|
| Kernel | Pallas(tops) | Pallas(vllm_torchtpu.kernels) |
| 中间层 | JAX / Flax | torch custom op ← pallas.jax_op |
| 模型定义 | 需用 JAX 重写 | 复用 vLLM 上游 PyTorch 实现 |
| 服务层 | vLLM + tpu-inference | vLLM + vllm_torchtpu |
| KV 传输 | — | tpu-raiden(JAX/torch 双绑定) |
注意最后一行:tpu-raiden 同时提供 JAX 和 PyTorch 绑定(tpu_raiden/api/jax 与 tpu_raiden/api/torch)。基础设施层已经是两边共用的了。
7.2 路线 A:JAX(tpu-inference + MaxText + tops)¶
优势:tops.chunk_kda 是唯一现成的逐通道 KDA 分块实现,且有完整的 v6e roofline 分析文档;MaxText 的 CP(含 KDA 支持)与 ragged_all_to_all 成熟;已有 K2.5/DeepSeek-R1 在 v7x 上的实战积累与踩坑记录。
劣势:KDA 解码要从零写(kernel + cache + 接口三层);模型定义要用 JAX 重写;MXFP4 走不通,必须退回 block-wise FP8,权重从 1.553 TB 涨到 2.83 TB。
56–87 天。
7.3 路线 B:torchtpu-vllm(+ torch_tpu + tpu-raiden)¶
优势:GDN 逐通道解码 kernel、MLA ragged paged attention、MXFP4、循环状态的 prefix caching 与 slot 管理、P/D 分离、host offload、投机解码 —— 全部现成;模型定义复用 vLLM 上游,K3 的低秩门、SiTU、AttnRes 都只是 PyTorch 层的事;开发活跃(2026-07 仍在合入 Mamba prefix caching、DP scheduler 等)。
劣势:分块 prefill 路径的衰减门只到 per-head,需要提到逐通道(G2);相比 JAX 路线缺少 CP,1M context 的 prefill 切分需要评估;v7x 上的实战记录不如 k25-jax 丰富。
30–52 天。
7.4 建议:以 torchtpu-vllm 为主线,从 tops 移植逐通道衰减¶
理由有三:
- 工作量少一半左右,且关键路径(G2)边界清晰、有现成参考实现可移植 —— 比"从零写解码 kernel"风险低得多。
- MXFP4 能用,权重 1.553 TB 而非 2.83 TB。这直接决定 16 chips 是否可行,也给 block-FP8 留了退路。
- 战略方向一致 —— torch_tpu 生态是后续重点投入的方向,torchtpu-vllm + torchtitan + tpu-raiden 是配套的推理/训练/传输三件套。押在收敛的方向上,长期维护成本更低。
tops 不会被浪费 —— 它在这条路线里的角色是逐通道 KDA 数学的参考实现与移植源,tops/cpu/ops/kda/ 的高精度 CPU 实现还是数值对拍的基线。
保留的两处风险:CP 缺失对 1M context prefill 的影响需要早验证;如果 G2 的移植意外受阻,JAX 路线仍是可回退的 plan B(届时接受 block-FP8 与 32 chips)。
8. 服务化与成本¶
8.1 平台对比¶
| v7x × 32 | v6e × 64 | H20 × 64 | |
|---|---|---|---|
| HBM 总量 | 6.144 TB | 2.048 TB | 6.0 TB |
| 装得下 K3 (MXFP4 1.553 TB) | ✅ 25% | ✅ 76% | ✅ |
| 装得下 K3 (BF16 5.559 TB) | ✅ 90% | ❌ | ⚠️ 93% |
| 聚合 HBM 带宽 | 236 TB/s | 105 TB/s | 256 TB/s |
| 互连(双向/芯片) | ICI 1200 GB/s | ICI 800 GB/s | NVLink(未核实) |
| a2a 占 decode 时间 (B=256) | 2.1% | 0.7% | — |
| 软件就绪度 | 需 30–52 天 | 同左 | 官方推荐 vLLM/SGLang,开箱 |
v6e 的问题不只是容量:64 chips 才凑够带宽,而 64 又不能整除 96 头,sharding 会很别扭。v6e 仅建议用于算子开发。
H20 是 Moonshot 官方评测用过的硬件(model card 脚注提到部分 benchmark 在 H20 上跑),vLLM/SGLang 开箱即用。这是 TPU 方案必须对标的基线 —— TPU 路线要在 2–3 个月的工程投入后,在单位成本或长上下文吞吐上显著胜出才成立。
TPU 的结构性优势在 5.6 节:1M context 解码时注意力算力占绝对主导,而 v7x 单 chip 4614 TFLOPS FP8 的算力密度加上 7.38 TB/s 带宽,在这个 regime 下有真实优势。K3 在 TPU 上的价值主张应当锚定超长上下文,而非通用短请求服务。
8.2 GKE 部署¶
k25-jax 的基建可直接复用,要点:
# 1. workload policy 必须先建
gcloud compute resource-policies create workload-policy <name> \
--accelerator-topology=2x4x4 --type=HIGH_THROUGHPUT
# 2. node pool 创建:--reservation-affinity=none 与 --flex-start 都必须
gcloud container node-pools create <pool> \
--machine-type=tpu7x-standard-4t --tpu-topology=2x4x4 \
--placement-policy=<name> --reservation-affinity=none --flex-start
其他约束:LWS (LeaderWorkerSet) 编排 + Ray 多机后端;GCS FUSE 必须 file-cache:max-size-mb:0;单 host 内存上限约 920 Gi,权重加载走 runai_streamer 可在 5 分钟内完成 641 GB(2.0–2.4 GB/s)—— K3 的 1.553 TB 按此速率约 11–13 分钟。
注意:v7x 节点上 backbone XLA 编译 crash 的问题在 K2.5 项目中尚未解决(疑似 Docker 镜像 AOT CPU feature 不匹配)。这是 K3 上线前的前置依赖。
9. 路线图¶
以 torchtpu-vllm 为主线的排期。
| 阶段 | 内容 | 工期 | 出口标准 |
|---|---|---|---|
| P0 前置 | 解决 v7x GKE 编译 crash;跑通 torchtpu-vllm 上现有的 GDN 模型(Qwen3-Next 系)确认 gdn/v2 解码 kernel 可用;用 dummy 92 层 MoE 实测小消息 a2a 的 per-call 开销 |
5–10 天 | 现有 GDN 模型在 v7x 上出正确 token;拿到 a2a 延迟曲线 |
| P1 数值对拍 | 用 tops/cpu/ops/ 参考实现搭建 KDA/MLA 逐层对拍框架,特别是逐通道衰减的 golden 数据 |
5–8 天 | 单层输出与 HF 实现误差在容差内 |
| P2 逐通道分块 prefill(G2,关键路径) | 把 tops.chunk_kda 的逐通道衰减移植进 ragged_gated_delta_rule_chunked.py:g_cumsum 从 [B,T,H] 升到 [B,T,H,K],去掉 [..., None] 广播,重调分块 |
8–14 天 | 分块 prefill 与 P1 golden 对齐;与 gdn/v2 解码 kernel 的 state 交接一致 |
| P3 全模型组装 | K3 模型定义(vLLM 上游或自建)、MLA absorbed 分页、LatentMoE 路由、SiTU/AttnRes、MXFP4 权重转换(G3–G6、G8) | 12–20 天 | 2.8T 全模型在 v7x-32 上前向正确 |
| P4 分布式与性能 | EP + DP-attention sharding、CP prefill、接 tpu-raiden 做 KV 卸载与 P/D 分离、性能调优 |
15–25 天 | 达到 5.4/5.5 节估算的 60% |
| P5 多模态 | MoonViT-V2(G7) | 8–13 天 | 图像/视频输入端到端 |
| P6 服务化 | GKE 部署、压测、SLO 调优 | 8–12 天 | 生产 SLO |
关键路径 P0→P1→P2→P3→P4,约 45–77 天。 P5 可与 P4 并行;P6 部分可与 P4 重叠。
对照:若走纯 JAX 路线,P2 换成"从零写 KDA 解码 Pallas kernel + 状态 cache + 层接口"(20–30 天,三层缺口),P3 需要全套 JAX 模型重写,关键路径拉长到约 75–115 天。差距主要来自 G1 和模型定义两项。
10. 风险¶
| 风险 | 影响 | 缓解 |
|---|---|---|
| 184 次小消息 a2a 的 per-call 开销过高 | batch=1 解码延迟翻倍(5 µs/次 → +167%) | 最高优先级实测,用 dummy 92 层 MoE 即可测;小 batch 降 EP 度数换延迟 |
| 逐通道衰减把分块 prefill 的 VMEM 占用抬高 | 需要缩小 chunk_size,prefill 吞吐下降 | g_cumsum 从 [B,T,H] 变 [B,T,H,K] 是 128× 放大;P2 阶段按 v7x 的 32 MB VMEM 重新扫 chunk_size |
| MXFP4 在 TPU 上无高效通路 | 退回 block-FP8,权重 2.83 TB,必须 32 chips | torchtpu-vllm 已有 quantization/mxfp4.py,风险主要落在纯 JAX 路线;先按 32 chips 规划兜底 |
| KDA 状态 FP32→BF16 数值不稳 | 并发上限砍半 | P1 阶段用 CPU 参考实现做精度扫描 |
| 显存压力迫使 EP 组跨 slice | DCN 带宽只有 ICI 的 1/12,prefill 直接超容量(105%) | 硬性约束 EP 在单 slice 内;跨 slice 只做 DP |
| v7x GKE 编译 crash 未解 | 项目无法启动 | P0 前置,必要时联系 TPU 团队 |
torchtpu-vllm / tpu-raiden 仍在快速迭代 |
上游 API 变动导致返工;tpu-raiden README 明确标注"尚不建议一般使用" | 锁 commit(raiden 用 lkg.version);改动集中在少数适配层,不散布到模型定义 |
| torch 侧的 raiden wheel 成熟度低于 JAX 侧 | KV 卸载 / P/D 分离落地推迟 | 该能力属 P4 优化项而非功能路径,可先不接 |
| K3 权重格式或结构后续变更 | 转换脚本返工 | 权重映射与模型定义解耦 |
11. 结论¶
可行,建议推进,但要把预期锚定在正确的场景上。
K3 的三个架构选择恰好都对 TPU 友好:混合注意力把 KV cache 压到只有 24 层的量级,LatentMoE 把 EP 通信砍半,MXFP4 QAT 让 2.8T 模型能装进 16 chips。1M context 解码时注意力算力占 23.7× 主导,正是 v7x 算力密度能发挥的 regime。
技术栈这一侧有两点容易被误判,值得单独说明。
第一,PyTorch 侧的成熟度远高于只看 torch_tpu 得到的印象。 torch_tpu 本身只是后端,最大的推理示例是 Qwen3-0.6B;但服务栈在 torchtpu-vllm 里,那里已经有一个融合的逐通道 GDN 单步解码 Pallas kernel(state 索引、prefill/decode/mixed 分区、L2norm、门控全部内联),以及 MLA ragged paged attention、MXFP4 MoE、循环状态的 prefix caching、投机解码、P/D 分离。
第二,「JAX 路线 vs PyTorch 路线」是个伪命题。 这些 kernel 本身就是 JAX/Pallas 写的,只是通过 torch_tpu._internal.pallas.jax_op 暴露成 torch 算子;tpu-raiden 也同时提供 JAX 和 torch 两套 binding。两条路线共用同一个 kernel 层和同一个传输层,分歧只在上层的模型定义与调度框架。
于是真正的缺口收敛到一处:torchtpu-vllm 的分块 prefill 路径衰减门只做到 per-head 标量(代码里 jnp.exp(g_cumsum)[..., None] 这个广播就是证据),而 K3 的 f_b_proj 输出 reshape 成 (h d) 要求逐通道。tops.chunk_kda 已经实现了逐通道版本,两边都是 JAX/Pallas、标志位命名几乎一致 —— 这是一个边界清晰、有现成参考实现可移植的任务,不是从零开始。
通信侧的结论同样与直觉相反。92 层 MoE 的 EP all-to-all 看起来像个大问题,但 ICI 带宽有一个数量级余量:v7x 每 chip 1200 GB/s 双向,decode 时 a2a 只占 HBM 时间的 2.1%,prefill 只占 ICI 容量的 5.8%。真正的风险在延迟而非带宽 —— 每步 184 次 all-to-all,每次只搬 1.79 KB,全是固定开销;若单次开销落在 2–5 µs,batch=1 解码会直接翻倍。这个量必须实测,好在成本很低(一个 92 层 dummy MoE 就够)。
三个必须尽早验证的未知量:小消息 a2a 的 per-call 开销、v7x GKE 编译 crash 是否已修复、以及 gdn/v2 解码 kernel 在现有 GDN 模型上是否真的跑得通。第一个决定小 batch 场景下 EP 度数怎么选,第二个决定项目能否启动,第三个决定 30–52 天这个估算是否成立 —— 它是整个 torchtpu-vllm 路线优势的来源,值得在 P0 花几天做实证而不是停留在读代码。
最后一条判断:不要用 K3 在 TPU 上做通用短请求服务 —— 那个 regime 下 KDA 的 434 MB/序列常量状态会压制并发,且 H20 + vLLM 开箱即用。TPU 方案的价值主张是超长上下文。
本报告基于静态代码分析与手工 roofline 推算,未经真机验证。参数模型已用官方公布的 2.8T/104B 交叉校验。所有性能数字为分析值,实测可能有显著偏差。涉及的仓库均在快速迭代中,代码引用对应 2026-07 的状态。