跳转至

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 已经把大部分基础设施做完了,真正要写的只有一处。

五条主要判断:

  1. 显存不是瓶颈,拓扑下限是 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。

  2. 长上下文是 K3 在 TPU 上最大的优势。 只有 24 层全注意力,采用 absorbed MLA 后 1M context 的 KV cache 仅 14.5 GB/序列。作为对照,DeepSeek-V3 的 61 层 MLA 在同等上下文下需要约 36 GB。K3 更长的上下文反而更省。

  3. 短上下文下瓶颈会反转。 KDA 的循环状态是 434 MB/序列且与序列长度无关。在约 32.5K token 以下,它比 KV cache 还占显存,直接压制并发批量。这是 TPU 侧容量规划最容易被忽略的一项。

  4. 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 天。

  5. 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.jsonmodeling_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_projf_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:

S' = S ⊙ exp(g_t)
residual = v_t − k_tᵀ S'
S = S' + β_t · k_t ⊗ residual
o_t = scale · q_tᵀ S

与 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.pyKimiMLAAttention.__init__ 有一句硬断言 assert self.use_nope,且 self.rotary_emb = None —— K3 的全注意力层完全不用 RoPE

qk_rope_head_dim: 64 这个字段仍在,但它退化成一段不做旋转的 MQA 共享尾维k_rotkv_a_proj_with_mqa 切出后直接 expand 到所有头。位置信息全部由 KDA 层的短卷积与递推结构隐式承载。

其余维度沿用 DeepSeek-V3 风格:q_lora_rank=1536kv_lora_rank=512qk_nope_head_dim=128v_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: 1topk_group: 1,即实际未分组),routed_scaling_factor: 1.0moe_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)·upbeta=4.0。纯 element-wise,TPU 上无障碍,但需注意 tanhsigmoid 走 EUP 流水线。
  • AttnResattn_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.pycustom_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 jnpfrom 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.pyoffload/cpu_tpu.pyplatforms/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):三层都要补 —— topsfused_recurrent_kda 只是 lax.scan(每步把 434 MB/序列状态读出写回,无法融合门变换与 delta 修正);没有状态 cache 抽象;attention_kda.py:369 直接抛 NotImplementedError20–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_cumsumjnp.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_tpufloat8_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_tpoolpos_emb_type: divided_fixed、视频时序 patch 需要重做。torchtpu-vllm 路线下有 vision_attention.pymulti_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 递推的数值稳定性 —— topsfused_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):

mesh = (data=2, tensor=8)          # attention: DP=2, TP=8 → 12 头/chip
     × expert=16                    # MoE: 56 experts/chip
  • 专家权重 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):

mesh = (data=4, tensor=8) × expert=32   # 28 experts/chip,12 头/chip

推荐 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-rc2CPContextattention_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.mdkda-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/jaxtpu_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 移植逐通道衰减

理由有三:

  1. 工作量少一半左右,且关键路径(G2)边界清晰、有现成参考实现可移植 —— 比"从零写解码 kernel"风险低得多。
  2. MXFP4 能用,权重 1.553 TB 而非 2.83 TB。这直接决定 16 chips 是否可行,也给 block-FP8 留了退路。
  3. 战略方向一致 —— 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.pyg_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 的状态。