跳转至

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. 没有需要从零写的 kernel。 KDA 的两半都已经有人做完了:解码侧,torchtpu-vllm 有一个融合的逐通道单步 Pallas kernel(kernels/gdn/v2/gdn_decode_kernel.py,state 索引、L2norm、门控全部内联);prefill 侧,openxla/tokamax#1103 刚把完整的 KDA 分块前向 + 反向 + varlen + CP 加进了 OpenXLA 的公共算子库,逐通道、Apache-2.0、331 项测试全通过,且 K3 的配置(K=V=128、chunk 64、gate_lower_bound=-5.0)逐条落在其支持范围内。剩下的是接线工作:布局转置、beta sigmoid 外提、状态布局对齐。估 5–9 天。

  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 —— 全是固定开销。v7x-16 实测 per-call 9.88 µs(层间串行,不可重叠),184 次折算 1818 µs = batch=1 HBM 下界的 331%:a2a 式 EP 在 batch=1 下会让解码从带宽瓶颈变成通信瓶颈。 所幸推荐的实现主线(torchtpu-vllm)的 MoE 路径一次 a2a 都不发(EP 靠本地专家 + 偏移,combine 走 all-reduce),这条风险是结构性避开的 —— 但未知量随之转移到「真实解码路径一步到底有多少次 collective」,那个数还没有。见 §12.1 ⑩。

技术栈选择上,建议以 torchtpu-vllm 为主线(工作量 27–47 天,对比纯 JAX 路线的 56–87 天)。这里要澄清一个容易搞错的前提:torch_tpu 本身只是 PyTorch 后端,不是服务栈;服务栈是 torchtpu-vllm,而它的 kernel 全部是 JAX/Pallas 写的,通过 torch_tpu._internal.pallas.jax_op 桥接成 torch 算子。所以「JAX 路线 vs PyTorch 路线」是个伪命题 —— 两条路线不仅共用 Pallas kernel 层,再往下还共用 tokamax(公共算子库,KDA 正在进入)与 tpu-raiden(KV 传输,JAX/torch 双绑定)。分歧只在上层的模型定义与调度框架。详见第 3、7 章。

方法论与局限

本报告主体为静态分析:依据 K3 官方 config.json 与 HF modeling 源码,逐一对照 maxtext / pallas-kernel(tops) / tokamax / torch_tpu / torchtpu-vllm / tpu-raiden 六个仓库的实际代码(含 openxla/tokamax#1103 的完整 diff),配合手工 roofline 推算。第 2–10 章里未特别标注的性能数字均为分析值并标注了假设。

参数模型已做交叉校验:按 config 逐层重算得总参 2.7795 T、激活 104.19 B,与官方公布的 2.8T / 104B 吻合,说明结构建模正确,后续显存与算力推算建立其上。

已进入真机验证阶段——部分结论已被实测修正

报告发出后启动了配套的实现项目,在 Cloud TPU v6e-4(单机)v6e-16(4 台 host) 上做了系统性的真机验证。其中九条改变了上面的判断,最要紧的五条:整模型第一次真编出来,v7x-16 上 TP=32 一开始装不下(实测 temp 347.85 GiB vs HBM 94.74 GiB)—— 而原因是个膝点不是线性累积:同样的图在 32 层以下 temp 平在 7.7 GiB(正好等于单层峰值,XLA 复用得干干净净),36 层之后调度塌掉;所以修法是把 93 层折成 scan(图里只剩 9 层)、而不是去省那份缓冲,改完复测 temp 降到 11.00 GiB,合计 77.59 / 94.74,装下了;再修掉一个白吃 16.33 GiB 的分片 bug 之后是 60.71,单次 prefill chunk 顶到 52K token(53248 装得下,57344 差 82 MiB),把 KV cache(实测 55296 B/token)叠回去之后"52K chunk"和"长上下文"是同一份余量的两种花法 —— 52K chunk 只剩 96K 上下文,压到 32K chunk 才有 178K;拓扑下限要按 device 而不是芯片来数(第 1 条的"16 chips"= TP=32);torchtpu-vllm 的多机拓扑表里根本没有 32-device 的条目,不打补丁 v7x-16 上第一步就死;"权重 ÷ TP"这个显存公式假设了全量张量并行,而在自建栈上这个假设一度让结论整个反号;decode 的成本几乎全在循环怎么写上,KDA kernel 只占 5.8%。

2026-07-31 起 v7x 真机也点亮了(GKE 上一片 16 芯 tpu7x-standard-4t,2x2x4,排队 19h51m)。最要紧的一条:报告里风险最高的那条假设通过了 —— tokamax 的 KDA Pallas kernel 不改一个参数就在 TPU7x 上编译、运行、算得准;一芯两 device 和整片 3031.6 GiB HBM 也都是量出来的了。

2026-08-06:真权重在生产服务线上端到端跑通了。 v7x-16 / TP=32 + EP,真 checkpoint(1453.7 GiB)从 GCS 直读,9 分 14 秒加载完,常驻 45.43 GiB/device,35 条 prompt 全部生成完成,输出的事实内容是对的(巴黎、彩虹七色、足球 11 人、尼罗河最长……)。这一条同时验证了 mxfp4 打包态常驻 + 前向解量化、MLA、KDA、MoE 路由、TP=32 分片 —— 任何一处切错,出来的会是通顺但事实错乱的文本。见 §12.6。

但仍然没有稳态吞吐/延迟数字:那次 21 分 36 秒里绝大部分是第一次运行的编译(前 2 条 prompt 占 14 分 31 秒,后 33 条只用 7 分 05 秒)。它证明"能跑通",不证明"跑多快"。

逐条见文末 第 12 章 · 实测补记。那一章写明了哪些还是分析值。


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
openxla/tokamax OpenXLA 的公共 Pallas 算子库 KDA 正在上游化到这里,见 3.2
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 tokamax —— KDA 正在上游化到公共算子库

openxla/tokamax 是 OpenXLA 的公共 Pallas 算子库,MaxText 已经把它列为依赖,torch_tpu 也通过 pallas.custom_jax_kernel 直接调用它的 ragged_dot它在两条技术路线的下方,是共享的。

2026-07-28 提交的 openxla/tokamax#1103(蚂蚁集团贡献,+10,724 行)把 KDA 完整地加进了这个库:

  • XLA 递推参考实现(base.py,逐 token 的 fori_loop,作为数值契约)
  • Pallas TPU 分块前向pallas_tpu_fwd.py,2240 行)
  • Pallas TPU custom-VJP 反向(pallas_tpu_bwd.py,1840 行)
  • packed varlen(segment_ids)、序列维 context parallelism(CP2/CP4)
  • autotuning 注册、331 项测试全通过、前后向各一份设计文档

公开 API 与 K3 的调用点几乎逐字对应:

def kimi_delta_attention(
    q, k, v,                       # [H,B,T,K] / [H,B,T,K] / [H,B,T,V]
    g: Float[Array, "H B T K"],    # ← 逐通道,docstring 写明 "Per-channel gate tensor in log space"
    beta: Float[Array, "H B T"], *,
    A_log=None,                    # [H]        ← K3: nn.Parameter(num_heads)
    dt_bias=None,                  # [H*K]      ← K3: nn.Parameter(projection_size=12288)
    initial_state=None,            # [B,N,H,K,V]
    output_final_state=False,
    use_qk_l2norm_in_kernel=False, use_gate_in_kernel=False,
    segment_ids=None, safe_gate=True, lower_bound=None,
    cp_context=None, chunk_size=64, N_max=None,
    implementation=None,           # "pallas_tpu" 优先,"xla" 兜底
)

XLA 参考实现里的递推与 K3 的 modeling_kimi_linear.py 完全一致 —— 注意 [:, None] 而非标量广播,这就是逐通道:

state = previous_state * jnp.exp(g_h[h, b, t])[:, None]
residual  = v_h[h, b, t] - k_h[h, b, t] @ state
new_state = state + (beta_h[h, b, t] * k_h[h, b, t])[:, None] * residual[None, :]
out_t = q_h[h, b, t] @ new_state

K3 的实际配置全部落在支持范围内(逐条对照 pallas_tpu.py_check_pallas_inputs_support):

约束 tokamax 要求 K3 实际值
dtype bf16 / fp32 bf16
K, V ≤ 256;CP 下需 %128==0 128 / 128
chunk_size 仅 64 64
lower_bound -5 <= lb < 0(sigmoid 门变体) gate_lower_bound: -5.0 ✅ 恰在边界内
A_log [H] [96]
dt_bias [H*K] [12288]

四处需要适配的差异

  1. 没有单步解码 kernel。 这个 PR 只有分块前向 + 反向 + XLA 参考,decode 路径仍要用 torchtpu-vllm 的 gdn/v2
  2. CP 模式下禁用 initial_stateoutput_final_state 意味着 CP 与"带状态接续的分块 prefill"目前不能组合 —— 对 1M context 的分段 prefill 是实际限制。
  3. beta 需在 kernel 外先做 sigmoid。 K3 用的是 use_beta_sigmoid_in_kernel=True,tokamax 无此开关(测试里由调用方 jax.nn.sigmoid(beta))。开销可忽略([H,B,T] 而已)。
  4. 布局是 head-first [H,B,T,K],而 K3/FLA 与 tops[B,T,H,K];状态是 [B,N,H,K,V],与 gdn/v2[num_states,H_v,K,V] 也需对齐。

另外:PR 目前是 OPEN 未合入reviewDecision: REVIEW_REQUIRED,copybara import PENDING),代码位于 tokamax._src.ops.experimental 实验命名空间。作为规划依据可以,但落地前要跟进它是否合入以及 API 是否变动。

3.3 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.4 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.5。

单看这个仓库没有服务化能力(无 paged attention、无 EP、无 MLA、最大推理示例是 Qwen3-0.6B)。这容易让人低估 PyTorch 侧的成熟度 —— 实际上服务栈不在这个仓库里,torch_tpu 是底座,不是终点。

3.5 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.6
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.6 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.7 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 | 本项目的核心缺口 | 5–9 天

这是把 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 表示 → 状态传递),要改的只是把衰减的秩从标量提到向量。而逐通道版本已经有两份现成实现,都是 JAX/Pallas:

来源 可得性 附带能力
tops.chunk_kda primatrix/pallas-kernel,MaxText 依赖 前向优化充分,有 v6e roofline 文档
tokamax.kimi_delta_attention openxla/tokamax#1103公开、Apache-2.0 + custom-VJP 反向、varlen、CP2/CP4、autotuning、331 项测试、XLA 参考实现

优先走 tokamax。 理由是它在依赖链的更下层:MaxText 已经依赖 tokamax,torch_tpu 也已经用 pallas.custom_jax_kernel 调用 tokamax 的 ragged_dot —— 接入机制是现成且验证过的,不需要把 kernel 源码搬进 torchtpu-vllm 再自己维护一份。

所以这项工作从"移植逐通道衰减的数学"降级为"接线":把 tokamax.kimi_delta_attentioncustom_jax_kernel 包成 torch op,处理 [H,B,T,K][B,T,H,K] 的布局转置、beta 的 sigmoid 外提、以及状态布局 [B,N,H,K,V]gdn/v2[num_states,H_v,K,V] 对齐。5–9 天。

落点变成:tokamax 的逐通道分块 prefill + torchtpu-vllm 的逐通道解码 kernel + torchtpu-vllm 的服务栈tops 退居为第二来源与数值对拍基线。

两处必须自己解决的残留

  • CP 与带状态接续的分块 prefill 不能组合 —— tokamax 在 CP 模式下禁用 initial_state / output_final_state。1M context 的分段 prefill 要么不开 CP,要么自己补状态跨 rank 的接续逻辑。这是 P4 阶段要早验证的一项。
  • PR 尚未合入。 如果它最终没进 main,退回 tops.chunk_kda 移植(回到 8–14 天)。

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.7 节的约束 —— 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 / tokamax),0 5–9 天(核心)
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 天 27–47 天

单人估算,不含真机调试与性能调优。

两条路线的工期差距集中在两处:G1(JAX 要从零写解码 kernel + 状态 cache + 层接口三层,torchtpu-vllm 已有)和模型定义(JAX 要整套重写,torchtpu-vllm 复用 vLLM 上游)。除此之外还有一个不体现在工期里但同样重要的差异 —— G4 决定权重是 1.553 TB 还是 2.83 TB,进而决定拓扑下限是 16 chips 还是 32。

torchtpu-vllm 路线的关键路径是 G2,而它现在的性质是"接线"而非"写 kernel":逐通道分块前向在 tokamax(公开)和 tops(内部)各有一份现成实现,要做的是布局适配与状态对齐。相比"从零写解码 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,仅可用于单层/子模块的算子开发验证。

2026-07-31 实测修正:这张表偏保守约 6%

在真机上量到的是 94.7 GiB/device × 2 device/chip = 189.4 GiB/chip,一片 2×2×4 共 3031.6 GiB ≈ 3.26 TB("192 GB" 这个规格数看来是 GiB)。 也就是 16 芯那一行的实际 HBM 比表里多 ~6%,权重占比 50.6% → ~47.7%表的结论方向不变:8 芯装不下、16 芯是最小可行档,这两条现在有真机 HBM 数字撑着。详见 §12.3。

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 阶段做掉。

实测补记:做掉了,9.88 µs(串行)—— 比上表最差的一档还差一倍,184 次折算 1818 µs = HBM 下界的 331%。但主线 torchtpu-vllm 的 MoE 路径里一次 a2a 都没有,下面缓解手段第 1 条正是它已有的形态,所以这条风险是结构性避开的。完整数据和四条方法限制见 §12.1 ⑩。

缓解手段(按优先级): 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 路线 共用?
公共算子库 tokamax tokamax(经 pallas.custom_jax_kernel
专用 Kernel Pallas(tops) Pallas(vllm_torchtpu.kernels) 同为 Pallas,可互相移植
中间层 JAX / Flax torch custom op ← pallas.jax_op
模型定义 需用 JAX 重写 复用 vLLM 上游 PyTorch 实现
服务层 vLLM + tpu-inference vLLM + vllm_torchtpu
KV 传输 tpu-raiden(api/jax tpu-raiden(api/torch

注意首尾两行:tokamax 和 tpu-raiden 都同时服务两条路线 —— MaxText 依赖 tokamax,torch_tpu 也调用 tokamax 的 ragged_dot;tpu-raiden 提供 JAX 与 torch 双绑定。算子库与传输层这两个基础设施层已经是共用的,openxla/tokamax#1103 把 KDA 放进前者,等于让 K3 最关键的那个 kernel 也落到共用层上。

7.2 路线 A:JAX(tpu-inference + MaxText + tops)

优势:逐通道 KDA 分块前向天然可用(tops.chunk_kda,另有 tokamax 版本),且有完整的 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,需要接入 tokamax 的逐通道实现(G2);CP 支持不如 MaxText 成熟,1M context 的 prefill 切分需要评估;v7x 上的实战记录不如 k25-jax 丰富。

27–47 天。

7.4 建议:以 torchtpu-vllm 为主线,KDA 分块前向接 tokamax

理由有三:

  1. 工作量约为一半,且关键路径(G2)已经从"写 kernel"降级为"接线" —— 逐通道分块前向在 tokamax 和 tops 各有一份现成实现,torch_tpu 调用 tokamax 的机制也已验证过。
  2. MXFP4 能用,权重 1.553 TB 而非 2.83 TB。这直接决定 16 chips 是否可行,也给 block-FP8 留了退路。
  3. 战略方向一致 —— torch_tpu 生态是后续重点投入的方向,torchtpu-vllm + torchtitan + tpu-raiden 是配套的推理/训练/传输三件套,下方还有共用的 tokamax。押在收敛的方向上,长期维护成本更低。

tops 在这条路线里的角色是第二来源与数值对拍基线 —— 如果 tokamax 的 PR 没能合入,它就是 G2 的回退实现;tops/cpu/ops/kda/ 的高精度 CPU 实现无论如何都是精度扫描的基线。

保留的三处风险:tokamax#1103 未合入;CP 与带状态接续的分块 prefill 不能组合,对 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%
软件就绪度 需 27–47 天 同左 官方推荐 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 可用;跟进 tokamax#1103 合入状态并在 v7x 上跑其 pallas_tpu_test.py;用 dummy 92 层 MoE 实测小消息 a2a 的 per-call 开销 5–10 天 现有 GDN 模型在 v7x 上出正确 token;tokamax KDA 测试在 v7x 通过;拿到 a2a 延迟曲线
P1 数值对拍 tops/cpu/ops/ 与 tokamax 的 XLA 参考实现搭建 KDA/MLA 逐层对拍框架,产出逐通道衰减的 golden 数据 5–8 天 单层输出与 HF 实现误差在容差内
P2 逐通道分块 prefill(G2,关键路径 tokamax.kimi_delta_attentionpallas.custom_jax_kernel 接入 torchtpu-vllm:[H,B,T,K][B,T,H,K] 转置、beta sigmoid 外提、状态 [B,N,H,K,V][num_states,H_v,K,V] 对齐 5–9 天 分块 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,约 42–72 天。 P5 可与 P4 并行;P6 部分可与 P4 重叠。

P0 需要增加一项跟进事项:确认 tokamax#1103 的合入状态。它目前是 OPEN、REVIEW_REQUIRED、copybara import PENDING。如果推进缓慢,可以先按 PR 的 head ref 做本地集成验证(Apache-2.0,允许),但生产前要等它落到 main。

对照:若走纯 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 度数换延迟
tokamax#1103 未能合入或 API 变动 G2 从 5–9 天回到 8–14 天(改为移植 tops.chunk_kda P0 阶段跟进合入状态;接入时把适配层与 kernel 调用解耦,便于换源
tokamax 的 CP 模式禁用 initial_state / output_final_state CP 与带状态接续的分块 prefill 不能组合,1M context 分段 prefill 受限 P4 早验证;退路是长序列 prefill 不开 CP,或自行补状态跨 rank 接续
逐通道衰减把分块 prefill 的 VMEM 占用抬高 需要缩小 chunk_size,prefill 吞吐下降 衰减从 [B,T,H][B,T,H,K] 是 128× 放大,而 tokamax 只支持 chunk_size=64;P2 阶段核对 v7x 的 32 MB VMEM 是否够用
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 算子;再往下,两条路线还共用 tokamax(公共算子库)与 tpu-raiden(KV 传输)。分歧只在上层的模型定义与调度框架。

于是缺口收敛到 KDA 的分块 prefill 一处 —— torchtpu-vllm 自带的那条路径衰减门只做到 per-head 标量(jnp.exp(g_cumsum)[..., None] 这个广播就是证据),而 K3 的 f_b_proj 输出 reshape 成 (h d) 要求逐通道。而这一处也已经有现成实现:tops.chunk_kda 是一份,2026-07-28 提交的 openxla/tokamax#1103 是另一份 —— 后者公开、Apache-2.0、附带反向与 CP、331 项测试全通过,且落在两条路线共用的那一层上。K3 的配置逐条核对下来全部在其支持范围内。

所以本项目没有需要从零写的 kernel。 剩下的是布局转置、状态对齐、模型组装、分布式调优这类工程工作 —— 难度和不确定性都低得多。

通信侧的结论同样与直觉相反。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 模型上是否真的跑得通tokamax#1103 能否合入且其测试在 v7x 上通过。第一个决定小 batch 场景下 EP 度数怎么选,第二个决定项目能否启动,后两个决定 27–47 天这个估算是否成立 —— 它们是整条路线优势的来源,值得在 P0 花几天做实证而不是停留在读代码。

最后一条判断:不要用 K3 在 TPU 上做通用短请求服务 —— 那个 regime 下 KDA 的 434 MB/序列常量状态会压制并发,且 H20 + vLLM 开箱即用。TPU 方案的价值主张是超长上下文。


12. 实测补记(v6e / v7x 真机,持续更新)

报告发出后启动了配套实现项目,机器是 Cloud TPU v6e-4(单机 4 芯)和 v6e-16(4 台 host × 4 芯)。选 v6e-16 是因为它的 host 数跟 v7x-16 完全一样,凡是"多机才暴露得出来"的问题都能提前跑掉。

这一章只写实测的东西,并明确区分「实测修正了报告」「实测支持了报告」「仍是分析值」三类。

12.1 改变了报告判断的十条

① 拓扑下限要按 device 数,不是芯片数 —— 第 1 条判断的"16 chips"等于 TP=32。

v7(Ironwood)一颗芯片是 2 个 chiplet,各有独占的 96 GB HBM,JAX/vLLM 把它们暴露成 2 个独立 device。"每芯片 192 GB"是两块 96 GB 的标称和,不是共享池。所以要拿「权重 ÷ TP」跟 96 GiB 比:

配置 每 device 权重 结论
TP=8(4 芯 = 1 host) 181.60 GiB 超上限 86 GiB,装不下
TP=16(8 芯 = 2 host) 90.80 GiB 只剩 5.17 GiB,KV cache 上限 ~181k token-slot
TP=32(16 芯 = 4 host,即 v7x-16 满配) 45.40 GiB 剩 50.59 GiB,~1.77M token-slot

TP=16 那行还有个附注:上表按 HBM 100% 可用算,而 vLLM 的 gpu_memory_utilization 默认 0.9 先砍掉 9.6 GiB/device —— 5.17 GiB 的余量在默认参数下是负的。对 TP=32 这只是把余量从 50.6 砍到 ~41 GiB。

报告第 1 条说"最小可行 16 chips",方向是对的,但它在并行度上的含义是 TP=32 而不是 TP=16 —— 而这个数恰好落在下面第 ② 条的坑里。

(这张表算的是"总权重 ÷ TP"。那个公式本身还藏着一个更根本的问题:它假设了全量张量并行,而全不全量取决于实现,不是数学性质。见第 ⑤ 条 —— 在自建 JAX 栈上,这个假设一度让结论整个反号。)

torchtpu-vllm 的多机拓扑表里没有 32-device 条目,不补丁 v7x-16 第一步就死。

执行器用 device 数当 key 查一张硬编码表,查不到直接 Cannot find topology for 32 chips。而这张表没有代次维度,同一个 key 在 v6e 和 v7 上语义不同:v6e 的拓扑串是 3 分量 X,Y,Z,v7 是 4 分量 X,Y,Z,D(第四位是每芯片 device 数)。

在 v6e 上喂 4 分量串会当场 CHECK-fail,不是静默算错:

F0730 tpu_topology.cc:1620] Check failed: false ghostlite does not have TensorNode
    ... tpu::TpuTopology::TensorNodesPerChip()

这条实测有两个用处:它确认了第四分量在 libtpu 里的语义就是 "TensorNodes per chip",从而正面支持我们给 v7x-16 推的 2,2,4,2;也说明 v7x 上串写错是响的不是哑的,拓扑这一环要么过、要么崩在一个看得懂的地方。

顺带一条上游观察:Ray 自己的节点 label 就带着 ray.io/tpu-topology(实测 4x4)和 ray.io/accelerator-typeTPU-V6E),两者加起来足以唯一确定拓扑串,而执行器把它们全无视了。比往硬编码表里加行更根本的修法在这里。

③ MoE 峰值激活曾经是"能不能跑 prefill"的开关,不是性能问题。

原实现(jnp.take 按 token 物化专家权重)的临时张量是 top_k × latent × inter × 3 × 2 byte ≈ 1.06 GB/token/层 —— 8 个 token 同时在飞 7.88 GiB/层,64 个 token 63 GiB/层,连 TP=32 剩下的 50 GiB 都装不下一层的临时量。

这个公式本身是算出来的,但决定性的那一半是实测的:在 v6e-4 上编译分片后的 MoE、读 HLO 确认 GSPMD 对"沿分片轴 gather"的处理是每个设备算出一份完整形状(本地没有的专家位置填 0)再拼回去 —— 也就是说这个数不随 TP 变小,TP=32 也救不了。换成 shard_map + 本地专家循环后同规模退回到几 MB/层(这一版的峰值是真机实测 + 保守包络上界,不是纯公式)。

补记(2026-07-30):上面那句"几 MB/层"量的是当时的 lax.scan 逐本地专家循环,实现此后又换过一次,数字要跟着改。 现在的实现是 group_sort + jax.lax.ragged_dot(换它是为了 decode:真实路由稀疏度下单 token 快 9.29×,见下一节)。代价是临时缓冲换了形状 —— 从 [N, inter] 变成五个 [N × top_k, *]多乘了一个 top_k = 16。真 K3 形状下的结构上界(五个缓冲全物化、一个都不复用):8 token 12 MiB/层,512 token 516 MiB/层,比 lax.scan 版的同规模 19 MiB/层大 27 倍。v6e-4 上按同等行数实测到 305 MiB/层,落在上界的 59%。

结论不变,但理由要换个说法:512 MiB/层对 TP=32 剩下的 50 GiB 仍然是约 1%,而且层是串行的、不累加(「不累加」这半句在 32 层以内成立、到 93 层就不成立了 —— 见 ⑨ 的膝点)—— 所以"真权重能跑 prefill"这个判断没有变,变的是它的余量从"三个数量级"缩到"两个数量级"。原来那句"几 MB/层"如果照抄进任何后续规划,会把余量高估约 30 倍。

同一轮真机数据还推翻了一条旧结论:lax.scan 版下本地专家数 n_local 不影响临时缓冲(4 和 28 量出来一位不差都是 195,904 byte),换成 ragged_dot 之后 group_sizes 的长度就是 n_local,同样 512 token 下 n_local=4 是 290,080 byte、n_local=28(真 K3 TP=32 的本地专家数)是 13,002,944 byte,差 44.8 倍。这条不改变上面的账(结构上界里没有 n_local 项,13 MB 仍在 68 MiB 上界之内),但它是个提醒:这类"某个维度不影响显存"的结论是当时那版实现的性质,不是这份计算的数学性质,换实现就要重测。

报告第 5 章的显存账里没有这一项,因为静态分析看不见 GSPMD 的这个行为。这是"读代码读不出来、必须上机"的典型例子。

④ KDA 的 T=1 decode 不是"pad 到 64 多算一点",是整个掉出 kernel。

报告第 3 章(_TOPS_CHUNK_SIZE = 64 那条)和本章早先的说法都是「序列长度不是 64 的倍数时需 padding」。读 tokamax 源码发现不对:pallas_tpu.pycheck_inputs_support 在任何预处理之前就硬性检查 segment_ids is None and seq_len % chunk_size != 0 并直接 NotImplementedErrorcommon.py 里那段 pad-and-slice 只服务变长(segment_ids)路径。而 kimi_delta_attention 的默认 implementation=("pallas_tpu", "xla") 会把这个异常 catch 住、静默退到 xla 的逐 token 顺序递推

所以真实的代价模型是二值的:"能不能进 chunk kernel",不是"多算多少倍"。v6e-4 实测(B=4, H=8, K=V=128,中位数):

T 每 token 耗时 走的路径
1 244.79 µs xla 顺序递推(含主机分发开销)
32 129.77 µs xla 顺序递推
64 6.70 µs pallas_tpu chunk kernel
65 127.74 µs xla 顺序递推
128 5.64 µs chunk kernel
256 5.11 µs chunk kernel

T=65T=6419.1× —— 一个 token 的差别,代价差近 20 倍。取 xla 的渐近每 token 成本(~128 µs,从 T=32/65 两点读出)比 chunk kernel(6.7 µs),decode 每步的 kernel 层面折扣是 ~19×。(T=1 那格的 244.79 µs 里有一半以上是每次调用的主机往返开销,不宜直接当分子。)

含义:KDA "O(1) 状态 decode" 这个卖点,在当前 tokamax 上拿不到。 要么去上游补一个 fused_recurrent 变体(HF 侧 modeling_kimi_linear.py:561 正是这么做的:mode = 'fused_recurrent' if use_cache and q_len == 1 else 'chunk'),要么在 Pallas 外面自己写单步递推。这条对 v7x 同样成立 —— 它是 kernel 的能力缺口,不是硬件代次问题。

补记(2026-07-31):两条"看起来能绕开"的路都实测否掉了,缺口只剩一个填法。

  • 不能靠"显式强制 pallas"逼它别退回。 tokamax api.py:118 把单个字符串包成 1 元组,:150-152 明写 except NotImplementedError: if len(implementation) == 1: raise显式传单个后端串就是没有回退,这是上游的既定设计(docstring 自己说"A sequence tries implementations in order, falling back when...")。所以强制 implementation='pallas_tpu' 做 decode 的结果不是"跑得更快",是当场 NotImplementedError。反过来说,这也意味着默认那个二元组是会骗人的:你以为在跑 pallas,实际在跑 xla,没有任何提示。
  • 不能借变长(segment_ids)入口把 T=1 塞进 chunk kernel。 那条路径的做法是 pad 到 chunk 边界再切(common.py:160-183),T=1 会被 pad 成一个 64 的 chunk —— 算 64 个 token 的活,出 1 个 token。按 T=64 实测的 6.70 µs/token 折算,一整个 chunk 调用约 429 µs(这是从每 token 中位数折算的,不是直接测到的调用耗时),而 xla 顺序递推的渐近每 token 成本约 128 µs429 vs 128 —— 变长路径给 decode 用会慢约 3.3×。 它不是"以后有空再做",是方向就不对。

唯一的填法是写一个专用的单步 kernelinitial_state 进、final_state 出、不 pad。这条工作量在报告里没有计入。

⑤ "权重 ÷ TP" 是对实现的一个假设,不是数学性质 —— 自建栈上它一度让结论反号。

报告第 5 章的显存账把总权重除以 TP。这隐含假设每一个张量都被切开了。真去查自建 JAX 栈实际切了什么(tools/memory_budget.py --stack jax,分类是从 k3/loader.py::_collect_specs 现场生成的真实 ShardSpec,不是手抄的,有测试钉住),发现最初只切了 embedding / lm_head / MoE 专家本体,其余全部复制。复制的那部分是一个跟 TP 无关的常数地板

复制项地板:101.34 GiB/device  →  34.71(切了 KDA/MLA 头轴投影)
                              →   0.81(又切了 shared_experts/dense_mlp/
                                        LatentMoE 投影/router gate.weight)

第一个数是致命的:101.34 GiB 本身就超过 96 GiB 的单 device HBM,意味着 TP 开到无穷也装不下 —— 跟"总权重 ÷ TP"给出的"TP=32 装得下"是反号的判断。两轮切分之后地板降到 0.81 GiB,账才第一次跟报告的公式吻合(差不到 1 GiB):

TP 权重/device(jax 栈实测口径) 结论
8(4 芯 = 1 host) 182.31 GiB 超上限 86 GiB,装不下
16(8 芯 = 2 host) 91.56 GiB 剩 4.41 GiB,装得下(很薄,还没算 KV cache)
32(16 芯 = 4 host,v7x-16 满配) 46.18 GiB 49.79 GiB

头条结论因此从"任何 TP 都装不下"变成"TP=16 就装得下真 K3",v7x-16 满配从"勉强"变成"有一半 HBM 是余量"。

这条的教训不在数字上,在账和实现的耦合关系上:报告那个公式没算错,它只是描述了一个"完全切分"的理想实现。移植过程中真正要盯的是"今天到底切了哪些",而这个答案每合入一个 PR 就会变一次 —— 所以这份账被做成了一个从 loader 现场取 ShardSpec 的脚本,而不是文档里的一张表。同一类事情在第 ③ 条里也发生过(n_local 从"不影响显存"变成"影响 44.8 倍"):"某某不影响某某"这种结论是当时那版实现的性质,不是这份计算的数学性质。

⑥ Mosaic kernel 在多设备上不会自动分区 —— 而且它在 eager 下是"悄悄变慢",在 jit 下是"当场崩"。

报告第 4 条说 KDA kernel 拿来就能用、剩下的是接线工作。这条成立,但它捡到的是单设备的 kernel。只要 mesh 上不止一个设备,不显式用 jax.shard_map 包住 pallas_call,JAX 就直接拒绝:

NotImplementedError: Mosaic kernels cannot be automatically partitioned.
    Please wrap the call in a shard_map.       (jax/_src/tpu_custom_call.py)

真正难缠的是这个异常什么时候能被看见。tokamax 默认的 implementation=("pallas_tpu", "xla") 会 catch 住它:

  • eager(自建 JAX 栈整模型今天就是这么跑的):trace/lower/compile 同步发生在 tokamax 那个 try/except 的作用域里 → 异常被接住 → 静默退回 xla。也就是说,在写出 shard_map 之前,"pallas kernel 在多卡上跑过"这件事一次都没真发生过,而所有测试都是绿的。
  • jax.jit(vLLM 那条线的集成路径):trace 只把 pallas_call 记成一条 jaxpr equation,不报错,try/except 顺利退出;Mosaic 的检查发生在 .lower() 阶段,那时早已出了作用域 → 异常直接穿透,当场崩

同一个根因,两种完全不同的表现。对选型很要紧:自建栈那条线的风险是"你以为在跑 kernel,其实在跑参照实现,慢 31×";生产那条线的风险是"多卡第一次编译就死"。后者其实是好事 —— 崩在一个看得懂的地方,总比绿着算错强。

修法是无条件包 shard_map 沿头轴切(头之间没有交互,切开是恒等变换;头数不整除设备数时打印一行并退化成全复制,不静默)。附带撞到一个上游洞:tokamax 自己的 xla 参照实现shard_map 不兼容 —— kda/base.py::_fwdjnp.zeros 初始化 fori_loop 的 carry,这个初始值不依赖任何分片输入,被推导成 invariant,而循环体一碰到按头分片的 q/k/v/g 就变成 varying,撞 scan body carry input and carry output must have equal types。实测确认这是静态类型推导的问题、不是数值问题(关掉检查后与不分片直接调 kernel bit-exact 相同)。

⑦ decode 的成本几乎全在"循环怎么写"上,不在 KDA kernel 上。

报告把 KDA kernel 的质量当成 decode 性能的主要变量。在 v6e-4 上把同一个模型的 decode 循环写成三种形态,各测稳态 per-token(MINIMAL 配置、16 步、随机权重):

形态 稳态 per-token
A:eager 逐步(Python 循环,每步一次主机往返) 20,145,838 µs(= 20.1
B:jax.jit 单步 + Python for 22,200 µs(= 22.2 ms)
C:lax.scan 整体 jit(prefill + 全部 decode 一个 jit) 1,554 µs(= 1.55 ms)

B 比 A 快 907×,C 比 B 再快 14.3×。而在形态 C 里用假 kernel 量 KDA 的份额:

含 KDA           1534.57 µs
抽掉 kernel      1445.29 µs   → kernel-only 份额 5.8%,Amdahl 上界 1.06×
抽掉整个 block   1166.81 µs   → 整个 block 份额 24.0%,Amdahl 上界 1.32×

在一个已经 jit 好的 decode 循环里,就算 KDA kernel 变成完全免费,整步也只快 6%。 这不推翻"要移植 KDA kernel"(prefill 那 31× 是真的,见 §12.2),但它改变了优化顺序:第 ④ 条那个"专用单步 kernel"的收益上界是 1.06×,而把循环从 A 改成 C 是 13000×。先做循环形态,再谈 kernel。

两个必要的限定:这是 MINIMAL 形状(4 层)、随机权重、v6e-4 单机,份额会随层数/专家数/序列长度变,不外推到真 K3 和 v7x;形态 A 那个 20 秒/token 是 eager 每步一次主机往返的开销,它衡量的是"不 jit 有多贵",不是任何人会真的部署的配置。顺带一个交叉验证:形态 C 那行 1554.20 µs 和份额测试里独立测出的"含 KDA"1534.57 µs 差 1.3% —— 两条测试各自重建模型、各自计时,落在同一个数上,说明这个稳态量本身是稳的。

⑧ 多 host:SPMD 让"计算"白拿,但"数据摆放"那一层有三个两两不等的编号,弄混了不报错。

在 v6e-16(4 台 host × 4 芯,host 数跟 v7x-16 完全一样)上跑自建 JAX 栈。好消息先说:jax.distributed.initialize 之后,jax.devices() 自动返回全局 16 个设备,k3.sharding.make_mesh() 一个字不用改就是 16 宽,整除判断也自动按 16 走。切分代码零改动 —— 这是 JAX 这条线相对 vLLM 那条线(要补拓扑表、要处理 ray rank 排序、要打上游补丁,见第 ② 条)实打实省下来的。

坏消息是数据摆放那一层。有三个都叫"编号"的东西 —— 启动器给的 process_idjax.process_index()mesh 里的设备顺序 —— 它们两两不相等:

  1. 启动器给的 process_idjax.process_index()。实测:传 process_id=0 的那台机器,process_index()2。这是个非恒等置换,两轮独立运行完全一致(不是随机的)—— TPU 上进程编号由运行时按物理拓扑定,process_id 只用来跟 coordinator 握手。这跟 vLLM 那条线的"ray rank 与地址表排序不一致"是同一类陷阱的翻版。
  2. process_index() 也不足以定位数据jax.make_mesh() 会按物理拓扑重排设备以优化 collective:进程 0 的 4 个设备落在 mesh 的第 0、1、6、15 位,不连续。于是"进程 i 拿全局张量的第 i 段连续列"这个假设当场失效。jax.devices() 恰好是 process-major 的([0,0,0,0,1,1,1,1,…]),所以今天用它建 mesh 是对的 —— 但对的原因是巧合,而且是个随时会没的巧合:哪天有人为了 collective 性能换成 jax.make_mesh()(v7x 上大概率要做),loader 会在没有任何报错的情况下开始加载错位的权重。

弄混任何两个都不抛异常、不打警告,返回一个形状和 dtype 完全正确、内容整体错位的全局数组。真权重下的表现就是"模型加载成功、输出乱码",而所有形状类的守卫都拦不住。

这里还有一条方法论上的教训值得单记:第一轮探针我用"全局求和"核对,8128 == 8128,全绿。那条核对是废的 —— 求和对列置换不变,摆错位置它一样过。换成逐元素(第 j 列全填 j)之后错位立刻现形。在一个"错了也不报错"的地方,选错校验量等于没有校验。

唯一跟设备顺序无关的答案是问 sharding 要:sharding.addressable_devices_indices_map(global_shape) 逐设备告诉你"这台设备要全局张量的哪一块",两种 mesh 建法下都对。process_index() 只用来回答"这块数据归不归我读",不用来算它在全局的哪个位置。

对报告的意义:多 host 不是移植的阻塞项,但阻塞点的位置跟报告设想的不一样。 报告把多机的难点放在通信和拓扑上 —— SPMD 让这一层基本白给。真正要花力气的是数据摆放,而且它的失败模式是静默的,需要专门设计校验量才看得见。

⑨ 整模型第一次真的拿去编译:v7x-16 上 TP=32 一开始装不下 —— 不是因为 MXFP4 把省下的还了回去,是因为图大到某个点之后 XLA 的调度自己塌了。把 93 层折成 scan 之后装下了,实测 77.59 / 94.74 GiB。

到 ⑤ 为止,所有显存账都是按类目手算权重 + 单层 MoE 真机 temp 外推拼出来的:93 层的整模型图从来没有被构造+编译过。2026-07-31 补上了这一步(nnx.eval_shape 抽象构造整模型 → 按真实 ShardSpec 分片 → AOT lower/compile 一次 prefill 前向),在 v7x-16(32 device,真实目标规模)上跑 TP=32,4 个 rank 报同一个数:

RESOURCE_EXHAUSTED: the total memory required for HLO temporaries (347.85G)
exceeds available HBM (94.74G)
手算预测 实测(93 层)
权重(≈ jit argument) 46.18 GiB 50.26 GiB +8.8%
临时缓冲 temp 7.42 GiB 347.85 GiB +4590%
合计/device 53.61 GiB ≥ 398.11 GiB v7x 上限 94.74

这个数是实测,站得住:当前实现的 93 层整模型,在 v7x-16 TP=32 上编不出来。 但"为什么"这一栏我写错过两次,两次都写在这份报告里过,所以下面把它连同错法一起留着。

第一版归因(错的):只改层数、其余一字不动地扫一遍,得到一条严格的直线 —— temp(L) = 3.4469·L − 3.0472 GiB,最大残差 0.0005 GiB,外推 L=93 得 317.51、实测 317.5158。而斜率有确切出处:3.4469 正好是每层解量化之后那份 fp32 专家权重本体(每层专家参数量 ÷ TP × 4 byte = 3.4453,比值 1.0005)。结论看起来无可辩驳:93 层的解量化结果同时活着,一层都没释放

那一趟扫描跑在 CPU 后端上 —— 因为 XLA:CPU 没有 HBM 上限,编不出来的配置在那里编得出来,才量得到 memory_analysis()。理由是充分的,代价是它量的不是目标平台。

在 v7x 上重新扫一遍,形状完全不同(tokens=64,TP=32):

L 2 8 16 24 32 36 40 48 56 64 93
temp GiB 7.35 7.62 7.65 7.68 7.71 13.9 23.7 42.6 190* 222* 348*

* = OOM 报文里 XLA 自己说"需要多少"。

L ≤ 32 是平的。 平在 7.35~7.71 GiB —— 正好是手算预测的 7.42,也正好是 ⑤ 那次在同一台机器上实测的单层 _moe_infer 峰值 7.3450。32 层结构完全相同的层摞在一起,XLA 把那块解量化缓冲复用得干干净净,一点没涨。36 层之后才开始涨,然后一路失控。

这是一个膝点,不是线性累积。 三条侧证:

  • 玩具复现,CPU 和 TPU 方向相反。 一个 L 层、每层"解量化一份只依赖参数的 256 MiB fp32 权重 → 乘激活"的最小图:CPU 上斜率/本体 = 1.0000(每层都活着),TPU 上 = 0.0005(压根不涨)。同一份代码,同一个 JAX,两个后端给出的是相反的答案。
  • 阴性对照:拿掉解量化,增长还在。 把专家权重直接以 fp32 存着(图里根本没有解量化这个操作),v6e 上斜率 1.1714,对照 mxfp4 的 1.3579 —— 86% 的增长跟解量化无关
  • 手算预测从来没错。 它是个单层口径的数,而在能验证的整个区间(L ≤ 32)里,它逐点吻合。为此一度立项去"把它改成整模型口径",那个改动会让它在每一个实测点上错 4~32 倍,已撤。

顺带排除掉三个听起来最像解药的修法,都是量过的,不是推的:

  • jax.checkpoint / nnx.remat:在纯推理前向里是空操作。 remat_p.bind 硬编码 differentiated=False,只有 jax.grad/jax.vjp 的 partial-eval 会把它翻成 True;而防 CSE 的 optimization_barrier 只在 differentiated 为真时才插入。没有反向传播,remat 什么都不做。
  • jax.lax.optimization_barrier:不是那把锁。 XLA:CPU 在 CSE 之后、调度之前用 cse_barrier_expander 把它整个剥掉;TPU 上在真实 MoE 形状里手工插进 shard_map 内部,斜率 1.3579 → 1.3602,没降反微涨。
  • xla_memory_fitting_level=EFFORT_O3:有效果,但差两个数量级。 L=64 从 221.85 降到 197.74 GiB。(顺带一个坑:这个 flag 走 LIBTPU_INIT_ARGS,放进 XLA_FLAGS 会让 parse_flags_from_env 直接 abort。)

所以修法方向也换了:不是去发明一种让 XLA 复用缓冲的办法 —— 它已经会了 —— 而是把图缩回它已经会复用的那个尺寸。K3 的层结构给了一个干净的切法:full_attn_layers = [4, 8, …, 92, 93] 意味着第 1~92 层恰好是 23 个结构完全相同的 4 层块(KDA, KDA, KDA, MLA),外加末尾孤零零一个 MLA。写成 nnx.scan,展开的图里只剩一个块,93 层塌回 4 层的规模,落在膝点左边很远的地方。

这条后来验证了,成立。 scan 改造合入后(93 层折成前导 4 层 + 22 个 4 层块的 scan + 末尾 1 层 = 图里只剩 9 层),在同一台 v7x-16、同一个 TP=32、同一个 K3_FULL、同一个探针脚本上复测,只换了源码:

L 8 32 64 93
temp GiB(改造前) 7.62 7.71 222* 348*
temp GiB(scan 后) 7.62 9.03 9.02 11.00
编译耗时 s(前 / 后) — / 16.4 49.8 / 15.6 70.6* / 14.8 — / 15.3

* = OOM 报文里 XLA 自己说"需要多少",编译没成功。

L=93 每 device 合计 77.59 GiB(权重 66.59 + temp 11.00),上限 94.74。膝点没了:拟合斜率 0.0349 GiB/层,是每层 fp32 本体(3.4453)的 1.0%

L=8 那一列是这趟测量里最要紧的一个点,它是阳性对照。 层数太少时可 scan 的块不够两个,代码自动退化成逐层展开 —— 于是它改造前后逐位相同7.621867507696152,两趟一模一样。这一位不差的相等排除了"其实是环境/编译器/我在别处动了什么",剩下的唯一变量就是图折了没有。第二个独立证据是编译耗时跟层数脱钩了(15.6 / 14.8 / 15.3 秒,而改造前 48 层就要 72.8 秒)。

这句"装得下"的边界得一起记住:说的是 93 层整模型的 prefill 前向图编得出来、并且装进 94.74 GiB/device。不含 KV cache、KDA recurrent state(那些另算),没有灌过真权重(仍是抽象形状树),整模型尺度上没有验过数值scan 版和展开版的对拍只在缩小配置上做过)。

prefill 长度这一维当天补上了。 上面那趟整个只跑了 tokens=64,而 64 个 token 的图里激活小到可以忽略 —— 量到的 11.00 GiB 基本全是权重解量化和 MoE 的固定开销,所以"装得下"当时只对 64 token 成立。把层数定死在 93、只扫 token 数:

tokens 64 1024 4096 8192 16384 32768
temp GiB 11.00 13.77 9.21 11.04 18.68 22.54
合计 / 94.74 77.59 80.36 75.79 77.62 85.27 89.13

扫到 32K token 都装得下,最紧的一点余量 5.61 GiB。而且 4096 那点(9.21)比 64 那点(11.00)还低 —— 这条曲线在低段根本不随 token 数走,是个跟 N 无关的地板;真正随 N 涨的只有 MoE 里 [N × top_k, ...] 那一族缓冲,要到 N ≈ 8192 才追上地板,之后也是次线性(8192→16384 涨 7.65,16384→32768 只涨 3.86)。低段非单调,所以这条曲线只能取包络上界,不能拿两点连直线。

我开跑之前预期它会在某处撞墙,那个预期是算错的:把 x_sorted 的 896 MiB 写成了 896 GiB,top_k 用了 8(真值 16),宽度用了 hidden 7168(真值是 latent 3584)—— 三处错凑出一个"4096 就要 28 GiB"的假墙。回去核原始记录,那条旧结论本身全是对的,错在我引用它的时候顺手把数重算了一遍。引一条已经算对的结论时重算,比直接引用更容易出错。

顺带,复测本身又挖出一个 bug:arg 那一栏从 50.26 涨到 66.59,而权重一个字节都没变。日志里刷了一片 hidden=7 / 15 / 22 不能被 n_dev=32 整除——退化为全复制,7/15/22 正好是三个 L 的块数 —— scan 把每个参数摞了一根前导轴之后,一处"从参数形状第 0 维读 hidden_size"的代码读到的是块数,整除检查判否,两个本来切得好好的投影静默退化成每片存全量。逐位对得上:88 个 scan 层 × 189.9 MiB = 16.32 GiB,实测差 16.33。所以 77.59 是个上界,修完预计 61.3。

修完之后我复测了一趟,那个"预计 61.3"只对了一半 —— 这半条错,比整节其他内容都更值得写下来。

arg 那一栏确实回到了 50.27(预测 50.26,对到小数点后两位)。但总量没有跟着降

tokens 64 4096 16384 32768
temp 修前 11.00 9.21 18.68 22.54
temp 修后 10.44 8.55 18.10 36.91
合计修后 60.71 58.82 68.37 87.18

32768 那点省下的 16.33 被 temp 吃回去 14.37,净赚 1.95。我那张"整列下移 16.33"的预测表,做法本身是错的 —— 它默认了 temp 和 arg 互相独立,而它们不独立:arg 腾出来的地方,XLA 会拿去花。支持这个读法的是三个不同 N 的合计都贴着天花板(修前 32768 是 89.13、修后 87.18、53248 是 89.78 —— N 差 1.6 倍,合计都落在 87~90 / 94.74);反过来,如果是"多了一个跟 N 成正比的新缓冲",16384 也该涨一半,实测它反而降了 0.58。

顺带把墙找到了。 既然要重测,就一路往上推到编不出来为止:

tokens 32768 49152 53248 57344 61440 65536
合计 / 94.74 85.57 ✅ 86.81 ✅ 89.78 ✅ 94.82 ❌ 99.58 ❌ 106.12 ❌

L=93 / TP=32 下单次 prefill chunk 的上限是 52K token 量级:53248 装得下,57344 差 82 MiB。这是个能直接写进服务配置的数,也是这一节里唯一一个不带"预计"两个字的上限。(中途先踩到的 49152 报的是 vmem 不是 HBM,那不是真墙 —— --xla_tpu_scoped_vmem_limit_kib=98304 就过去了,但它得走 LIBTPU_INIT_ARGS,塞进 XLA_FLAGS 会当未知 flag 直接 abort。)

形状值得注意:32768 → 53248 这一段 temp 只从 35.30 爬到 39.51(约 0.21 GiB / 1K token),然后一头撞墙。是悬崖不是曲线,所以悬崖前那段的斜率不能外推。 还有一件必须说清楚的:上面全部是"只有权重"的账,KV cache 和 KDA recurrent state 一个字节都没进过这张表 —— 真实服务把那两项叠上去,52K 只会更小。

当天晚些时候把那两项也量了,"只会更小"具体是多小有答案了。 同样 L=93 / TP=32 / K3_FULL / v7x-16,这次编的是 decode_step(cache 作为 jit 实参真的进图,形状和分片都从编译器要,不手写 —— 手写等于把要检验的假设再抄一遍然后拿它验证自己):

  • KV cache 55296 B/token,精确等于 576 channel × 4 B(fp32) × 24 个 MLA 层。拟合截距 +96 B,正好是 24 个 length 写指针标量 —— 一个免费的自校验。
  • KDA state 13.5 MiB/device,恒定,对上下文长度完全不变(它没有 seq 轴)。
  • decode 图自己的 temp 是 29.8 GiB,同样对上下文长度基本不动(4096 → 32768 只走了 +0.03)。这个数比我预期的大得多:它已经是 prefill 32K 那一档(35.30)的 84%。"decode 很省"这个直觉在 K3 上不成立 —— MoE 权重解量化那块地板,两张图都得付。

于是"这台机器能服务多长上下文"这个问题,第一次有了形状 —— 它不是一个数,是一条 trade-off 曲线。权重 50.27 是死的,剩下的余量要在"prefill chunk 开多大"和"上下文留多长"之间分:

prefill chunk prefill temp 留给 cache 能服务的上下文
32768 35.30 GiB 9.17 GiB 178K token
49152 36.54 GiB 7.93 GiB ≈ 154K token
53248 39.51 GiB 4.96 GiB 96K token ← 上面那面墙

所以"单次 prefill chunk 顶到 52K"和"上下文能有多长"是同一份余量的两种花法,不能同时取最大。 上一段那个 52K 是在没有 cache 的前提下量出来的;真要用满它,上下文就只剩 96K。(这张表是把两组实测数直接相减得到的,没有重新编译验证过 —— 按上一段刚立的规矩,它只能当下一步的输入,不能当结论。写出来是因为上次我做了同样的相减却没标出来。)

顺带撞出两个账本 bug,都属于"公式声称的范围和调用方以为的范围不一致"这一类:解析式预测把 KV 低估了 1.800×,把 KDA 低估了 7%。KV 那个 1.8 是两个独立错误相乘 —— dtype 按 bf16 算(实际 fp32,2×),布局按另一条服务栈的 paged KV 算(那侧 kv_c 和 k_pe 各自对齐到 128 得 640 channel,自建栈这侧就是 512+64=576,1.111×)。两个数各自都对,错的是那个函数签名里有 stack 参数,而 KV 那条路径完全不看它。KDA 那 7% 是漏了 short-conv 的三个环形缓冲 —— 那个函数的 docstring 明写着自己只算"delta-rule 递归状态",没算错,是调用方把它当成了"KDA 的全部常量状态"。两个 bug 的形状一样:一个函数诚实地做了 A,调用方以为它做的是 A+B。

这个 bug 的现场证据只有 stdout 上刷的那几行提示 —— 编译成功、数值正确、没有任何测试会红,只是每片白背 16 GiB。⑧ 那条教训(在"错了也不报错"的地方,选错校验量等于没有校验)在这里又演了一遍,只不过这次连校验量都还没有,只有一行 print。

改造之前那两条推论,一条被 scan 绕过去了,一条仍然成立:

  • 加 TP 治不了 —— 这条当时对,也正因为它对,才只剩折图这一条路。 temp 几乎严格随 1/TP 缩(v6e TP=4 实测 2938.9 GiB ÷ 8 = 367.4,v7x TP=32 实测 347.85,差 5%),但膝点的位置跟 TP 无关 —— 93 层在哪个 TP 下都在膝点右边,TP 只是把那个 45 倍的惩罚除小一点,按线性外推要 TP≈117 才装得下,而 v7x-16 一共 32 个 device。scan 不是把惩罚除小,是把图挪回膝点左边。
  • 它跟 ③ 是同一个假设的两次失效。 ③ 修好了"每层"的峰值(1.06 GB/token/层 → 512 MiB/层),结语写的是"而且层是串行的、不累加"。这句话在 L ≤ 32 时是对的,在 L=93 时不对 —— 而分界线在哪,手算是看不出来的。

对报告的意义:§5.1「权重显存」和 §5.3「拓扑下限」的账本身没错,但它们回答的是"权重放不放得下",不是"这个模型跑不跑得起来" —— 这两个问题第一次被分开问,就是在这一条里。报告第 1 条"最小可行 16 chips"作为权重容量的结论一直成立(⑤ 和 §12.3 的 HBM 实测都撑着它);作为可运行配置的结论,一度被实测否定(编不出来),现在重新成立,而且口径已经推进到"带 cache 的服务配置"这一层(权重 50.27 / 94.74 GiB,prefill chunk 32K + 上下文 178K,或者 chunk 52K + 上下文 96K,二选一)—— 但代价是产品代码里多了一层 scan 结构,而不是"报告说得对"。中间那段"否定"不是走了弯路,它是这条结论现在唯一的立足点:手算从头到尾都说装得下,说得下的理由却是错的。

这一条真正的教训不在显存上。 第一版归因给出的是一条残差 0.0005 GiB、斜率跟物理量比值 1.0005 的直线 —— 漂亮到我没有想过要在目标平台上复核它。拟合得越好,越不代表你量的是对的东西。 用一个能跑的后端代替一个量不了的后端,省下的是机时,赌上的是结论的符号。

而"修完预计 61.3"那半条,是同一件事换了个更隐蔽的样子:这次连拟合都没有,只有一次减法。 数是实测的(16.33 确实省下来了),推理的每一步单看都对,唯一错的是那个没写出来的前提 —— temp ⊥ arg。一个没写出来的假设,是不会有人来检查的,包括我自己。所以这条曲线上第四次栽跟头,仍然栽在"为什么"那一栏,不在"是多少"那一栏。规矩因此多了一条,跟"取包络不取斜率"并列:改了图就重编重测,不要拿"这一档变了 X"去算总量。

⑩ 小消息 a2a 的 per-call 开销测出来了:9.88 µs,比 §6.4 那张表最差的一档还差一倍 —— 但这个风险落在报告推荐的 EP 设计上,不落在报告推荐的那条实现路线上。

§6.4 把「小消息 all-to-all 的 per-call 开销」列为本项目最高优先级的未知量,并给了 1/2/5 µs 三档推算。2026-08-05 在 v7x-16(4 pod / 32 device / 真 TPU7x)上实测(tools/a2a_latency_bench.py,字节数显式取 dispatch_bytes_per_call 算出来的 1792 B,减掉等价空转基线,N 扫到 1000 确认收敛,四个 pod 独立量):

per-call p50 ×184 次/解码步 占 550 µs HBM 下界
串行(chained,层间真实依赖) 9.88 µs 1818 µs 331%
重叠(unchained,需要有独立工作可调度) 2.35 µs 433 µs 79%

四个 pod 的折算值:串行 1818 / 1884 / 1802 / 1824 µs,重叠 433 / 429 / 431 / 430 µs。两端都在报告那张表之外 —— 乐观端 2.35 µs 略差于表里的 2 µs 档,悲观端 9.88 µs 比表里最差的 5 µs 档还差一倍。

该用串行那一行。 184 次调用在层间是真实依赖的(第 N+1 层的 dispatch 要等第 N 层的 combine 出结果,层内 dispatch→专家→combine 也是串的),batch=1 解码根本不存在一池子互相独立的 a2a 供 XLA 重叠。§6.4 那个 warning box 里写的"延迟翻倍甚至更多",实测是 3.3 倍

报告的一条前提被正面实证了all_to_all 的成本对字节数完全不敏感 —— 64 B 到 16384 B 全部落在 9~10 µs,16 KB(9 倍的量)跟 64 B 一个价。§6.4 说"这个数量级下带宽完全无关,全部时间都是 collective 的固定开销",这句话现在有实测支撑,不再是推理。

但风险的落点跟报告设想的不一样。 §6.2 推荐的 sharding 是「DP-attention + EP」,那条路上确实有 184 次 a2a;而 §7.4 推荐的实现主线是 torchtpu-vllm,它的 MoE 路径里一次 all-to-all 都没有

  • vllm_torchtpu/layers/vllm/fused_moe.pyfused_moe_gmm 里零个 collective;
  • EP 是「每个设备只算本地专家 + 一个 experts_start 标量偏移」(mxfp4.py::_forward_monolithic_tpu 的注释原文:EP global->local remap happens inside fused_moe_gmm via an elementwise subtract);
  • combine 走 vLLM FusedMoEreduce_results=True,也就是一次 all-reduce
  • 整个 torchtpu-vllm 树里 jax.lax.all_to_all 只出现在 gdn_attention.py 的 PCP 路径和一个 experimental kernel 里,跟 MoE 无关。

也就是说 §6.4 列的缓解手段第 1 条「小 batch 时降低 EP 度数、专家改走 TP —— 用带宽换延迟」,正是主线已经在做的形态,只不过它不是被当成"缓解手段"选的,是上游本来就这么实现的。这条风险因此是结构性避开的,而不是需要我们去做的事。

这条同时把未知量换了个位置,没有消掉它。 换成 all-reduce 之后要问的是"92 层 combine + 每层 attention 输出的 all-reduce,一步总共多少"。同一次 sweep 里 psum 的 chained 数是 16.3~18.8 µs(比 a2a 贵,符合 all-reduce ≈ reduce-scatter + all-gather 的直觉),但那是一维 32 设备 mesh 上的合成 psum,不是真实解码路径里的 all-reduce —— 数量、尺寸、有没有跟计算重叠,三件事一件都没查。不要拿 18.8 × 186 去外推,那个乘法看着诱人,但它正是这个报告栽过四次的那类推理(见 ⑨ 的结语)。真正该做的是把一步 decode 编译出来数 HLO,已经排进队列。

方法上有四条限制,抄这些数的人必须一起抄走

  1. N < 1000 量不出重叠。 unchained N=100 是 9.4 µs,跟串行一样;到 N=1000 才掉到 2.35 µs。第一轮我只跑了 N=100,据此得出过"XLA 拿不到重叠收益"的结论 —— 那是错的,独立调用不够多时调度器摊不开。
  2. 这是代理指标。 量的是一维 32 设备 mesh 上的 jax.lax.all_to_all,K3 真实的 dispatch 是 ragged 的、走 shard_map。它量准的是"这个 collective 的固定开销有多大",不是端到端的 MoE 通信时间。
  3. ep_degree=64 那个假设没验。 1792 B 是按 EP=64 反推的,而我们在 32 设备上量。因为成本对字节数不敏感,EP=64 时字节减半不会让它变快;参与设备翻倍反而大概率更慢。没测,别推。
  4. 1 µs 以下这套测量没有分辨率。 psum 的 unchained 有几格 p90 是小负数(−0.1 ~ −4.6 µs),那是减基线的噪声。

工具本身还挡下了两个假数字:72 格里有 2 格(all_gather / 64 B / unchained)被逐格 HLO 核实标了 ⚠️ —— 编译器在小消息上把 all_gather 整个换成了 all_reduce,不核实的话那两个"很漂亮的小数字"会以 all_gather 的名义进报告。同一个 op 在同一次 sweep 里被换掉与不被换掉并存,这件事本身值得记一笔:collective 的实现选择是按尺寸触发的,阈值在 (64, 256] 之间某处,没有精确测。

12.2 实测支持了报告的几条

  • KDA kernel 不用从零写,这条成立。 tokamax#1103 在 v6e-4 上跑最小 K3 形状(H=8, B=2, T=256, K=V=128, chunk=64, lower_bound=-5.0, use_qk_l2norm_in_kernel=True):pallas_tpu vs 它自带的 xla 参考实现 max abs err 1.54e-7(均值 3.6e-9),耗时 0.72 ms vs 22.65 ms(31.6×),20 次取中位数,jax 0.11.0。那个 implementation='xla' 参考实现是白捡的对拍基准 —— 移植期间它比性能数字有用得多。

但报告第 4 条说的"接线工作"里有一处当时没看出来:use_qk_l2norm_in_kernel=True 必须传。不传的话 delta-rule 递推无界增长,实测输出量级到 1e14 —— 这不是数值误差是语义错误,而且它不报错、只是给你一个形状正确的垃圾张量。 - KDA 端到端能进 vLLM 的连续批处理。 2026-07-29 在 v6e-4 上,KDA 经 torch_tpu._internal.pallas.jax_op 接进 torchtpu-vllm,随机权重端到端出 token。报告第 7 章"分歧只在上层模型定义与调度框架"这条判断,到这一步是走通了的。 - 多机 Ray 路径跑通了 —— 而且是完整 K3 结构。 2026-07-30 在 v6e-16(4 台 host、共 16 芯)上,TPU_MULTIHOST_BACKEND=ray、TP=16 满片,两轮都 exit=0、35/35 出 token:一轮纯 KDA,一轮是真 K3 结构(3 KDA + 1 MLA、SiTU、MLA 输出门、LatentMoE、AttnRes 跨层残差)。这条路径在此之前从未被真机执行过,上游至今没有任何测试覆盖它,是 v7x 之前风险最大的一段。 - MoE 的 ragged_dot 改造做完了,而且实测确认了「稀疏才是重点」。 报告第 5 章说现有 dense 实现在真实路由下浪费巨大 —— 把玩具配置(4 专家)换成真实的「896 里抽 top-k、每设备 28 个本地专家」之后,每个 group 的占用率从 44% 掉到 1.8%。换成 group_sort + ragged_dot(三次 ragged matmul + scatter-add 合并)后在 v6e-4 上实测:真实稀疏度下 decode(1 token)9.29×,prefill(512 token)2.36×。但同一份探针也量出了反面:把 num_experts_global 退化成 28(每个本地专家都有 token、no-EP)时新实现反而是 0.67×~0.86×,比 lax.scan 慢。这个改造的收益完全来自路由稀疏度,不是「ragged matmul 更快」——真实 K3 恒为 896 专家所以站得住,但换个专家数就不成立。

  • 自建栈端到端生成跟 HF 逐 token 全等了。 2026-07-31,k3/ 这套自建 JAX 栈跑贪心生成,跟 HF generate(do_sample=False) 在同一份随机权重、同一个输入下比 token 序列:两组配置(全程 xla、以及 prefill 强制走 pallas_tpu chunk kernel)在 v6e-4 上都逐位相同。这是第一次有判据回答"它会不会生成跟 HF 一样的文本",而不只是"某一层的输出张量对得上"。

但要连着看下一个数:最小 top1-top2 logits 间距 3.5e-3。全等不是靠 argmax 的巨大间隔蒙对的,可这个余量也确实薄 —— 真权重、真词表下它会更薄,所以"逐 token 全等"这条判据在真权重上大概率要换成别的形式(这也正是 §12.5 里 TP 数值判据那条还挂着的原因:tiny 随机权重下 90.5% 的解码步是 bf16 平局)。 - 权重分发不需要额外设计。 多机加载时 32 个 worker 进程各读各的 checkpoint,看着像 32× 读放大;实测网络出口是 ~4×(每台 host 一份),因为 host 内靠 Linux 内核 page cache 共享 —— 在真实 gcsfuse 挂载点上让 4 个独立进程依次摸同一个 64 MiB safetensors 文件,后 3 个总共只带来第 1 个 2.4% 的 page cache 涨幅(本地盘对照组是 20.2%)。所以不需要 driver 广播这种设计。真正的约束是 host RAM 装不装得下工作集,而不是某个 gcsfuse 配置项。

12.3 v7x 真机第一次点亮(16 芯 / 4 host,2026-07-31)

这一节推翻了本报告此前"v7x 硬件一次都没摸到"的说明。 机器是 GKE 上一片 tpu7x-standard-4t × 4 台、拓扑 2x2x4、共 16 芯,走 queued provisioning 排了 19 小时 51 分才交付,租期 7 天。v7x 在 Cloud TPU API 里依然查不到, 只以 GKE 机型存在 —— 报告第 8 章那条判断成立。

先说结论:报告里风险最高的那条假设通过了。

  • KDA 的 Pallas mosaic kernel 在 TPU7x 上能编译、能跑、算得准。 tokamax 的 kernel 是照 v5e/v6e 写的,v7x 是新一代 MXU/VMEM,报告 §3.1 曾判断 "chunk/sub-chunk 的分块参数很可能需要重调"。实测:chunk_size=64 的 tokamax KDA kernel 不用改一个参数就在 TPU7x 上编译通过并跑出正确结果, 4 台 host 一致,tests/test_kda_kernel.py 4 passed / 1 failed。 这是整个 v7x 路线上单点风险最高的一环,现在它是通的。

  • 唯一那条红的,是我们自己的一个洞被 v7x 照出来了 —— 不是 kernel 的错。 pallas_tpu vs xla 对拍 max abs err = 1.576e-4(容差 1e-5)。分三种 默认 matmul 精度重跑,差距在 highest 下塌到 2.105e-7;再让两条路径各自 跟自己的 highest 版本比,4 台 host 结果完全一致:

xla       : max|default - highest| = 1.576e-04   ← 被默认精度降级
pallas_tpu: max|default - highest| = 0.000e+00   (逐位不变)

Pallas kernel 是逐位不动的那一侧,飘的是 tokamax 的 XLA 参照物。 根因是调 tokamax 时没钉 precision,跟着进程默认走 —— 而默认 matmul 精度是硬件相关的:v6e 上默认够准(对拍 1.5e-7,见 §12.2 第一条), v7x 上默认是一趟 bf16。对移植的普遍教训:任何"在 v6e 上对拍绿了"的 数值判据,只要参照物那侧没有显式钉精度,换到 v7x 就可能红,而红的原因 跟被测对象无关。这不是放宽容差能解决的,是要把参照物钉稳。

  • 一芯两 device,实测确认。 jax.device_count() 在这片 16 芯的机器上是 32,设备 coords 两两相同、core_on_chip 为 0/1。§12.1 第 1 条那个 "16 chips = TP=32" 此前是从 hardware.py 读出来的,现在是真机量出来的。 连带一条新坑:TPU_ACCELERATOR_TYPE=tpu7x-32 里的 32 是设备数,而 节点池的 --tpu-topology 2x2x4、ProvisioningRequest 的 count: 4、pod 的 google.com/tpu: "4" 说的都是芯片/机器。同一片机器上两个数都对。

  • HBM 实测 94.7 GiB/device → 189.4 GiB/chip → 整片 3031.6 GiB(≈3.26 TB)。 报告 §5.3 按 192 GB/chip 算的 16 芯 = 3.072 TB,实测略微保守("192 GB" 这个规格数看来是 GiB)。对照 §5.1 的 1.553 TB 权重:一片 16 芯 v7x 装下 权重后还剩约 1.7 TB 给 KV cache 和激活。报告第 1 条"最小可行 16 chips" 这个结论,现在有真机 HBM 数字撑着了。只限「权重装得下」这一层含义 —— 同一片机器上整模型编译是 RESOURCE_EXHAUSTED,见 ⑨。)

  • 跨 host collective 通了。 4 台机器、32 设备一个 mesh,all-reduce / matmul / all-gather 都对得上(sum(x²) 跨 host 规约 vs all-gather 回来在 host 上用 numpy 重算,相对误差 5.96e-08,是 fp32 累加顺序不同的正常量级)。

  • v6e 上那个 process_id ≠ process_index() 的坑,v7x 上原样复现 (0→1、1→2、2→0、3→3)。§12.1 里那条"多 host 代码不能认 rank"的修法 不是 v6e 的偶然。

  • jax.make_mesh() 的重排是拓扑相关的,不能当成 jax 的固定行为。 v6e-16(4x4)上它会把 process 顺序打散成 [0,0,1,1,1,1,0,2,3,3,3,3,2,2,2,0];同样的调用在这片 v7x(2x2x4 三维)上 完全不重排。含义是:一条"重排前后结果一样吗"的测试,搬到 v7x 上会 变成跑了两遍同一件事然后绿。

  • GKE 上多主机 TPU 还要 headless Service + spec.subdomain Indexed Job 给了正确的 hostname,但没有 subdomain 时 GKE 的 TPU webhook 注入的是 单主机那套 TPU_WORKER_HOSTNAMES=localhost,jax 直接起不来 (Expected 4 worker addresses, got 1)。

  • 这片机器不能拆开当 4 台单机用。 2x2x4 是一整片,libtpu 初始化时就 要求 4 个 worker 地址到齐。含义是 v7x 上跑任何东西 —— 哪怕是一条单机 测试 —— 都必须 4 进程一起起。

  • 整套 k3/ 测试在 v7x 上跑了一遍:154/252 绿,而红的没有一条是 v7x 算错。 98 条红按报错归类下来:~55 条是"32 除不尽"MINIMALH=8num_experts=8 是照 4/8/16 设备挑的,一芯两 device 之后是 32), ~31 条是测试假设自己跑在单进程里Fetching value for jax.Array that spans non-addressable devices),3 条 perf 探针除零,2 条数值(就是上面 那个精度问题),1 条是我推源码时漏了目录。换句话说:v7x 上目前没有 一条"算错了"的证据,红的全部是移植工程量,不是可行性问题。

  • 一个副作用值得单独记:单进程假设的测试会在"恰好持有 device 0"的那台机器上假绿。 4 台 host 里有一台多绿了 25 条,全是 MoE 那批。原因是那些测试把数组建在 jax.devices()[0] 上,而那块设备属于 process_index 0 —— 只有那台机器 float() 得出来,另外三台抛异常。它不报"测试写错了",它报的是另外三台 机器的错,多机调试时很容易顺着错误的线索走。

  • "钉精度"要成对做,只钉一侧比不钉更糟。 上面那个修复合入后在同一片机器 上复跑:原先红的两条数值测试绿了,但另外三条本来绿的翻红 —— 因为那几条 测试自己直接调 tokamax 当参照物,那些调用点没钉。被测侧钉了、参照侧没钉, 逐位相等的断言当然过不去。而这个错误在 v6e 上完全不可见(两侧默认都够准)。

这一节还没回答的:MoE 的 jax.lax.ragged_dot 在 TPU7x 上仍然没有一个 有效数据点 —— 上面那批 ragged_dot 测试要么被"非本地可寻址"挡在计算之前, 要么是靠"恰好持有 device 0"绿的,两种都没真正验到 v7x 的 mosaic 后端; 那组避免 VMEM OOM 的 tile 参数是照 v6e 的预算手调的,v7x 的 VMEM 未必一样。 整模型前向、真权重、以及任何吞吐/延迟数字在 v7x 上都还是零 —— 这片 机器是用来回答"能不能跑"的,不是"跑多快"。

12.4 真权重第一次上机(v7x-16,2026-08-04)

§12.3 那次点亮回答的是"能不能跑",用的全是 tiny 配置和 --load-format dummy。这一节是真 checkpointgs://k25-model-weights/moonshotai/Kimi-K3)第一次进 HBM 并真的算出 logits。

加载。 v7x-16、4 pod / 32 device / TP=32,直读 GCS:

RESULT 返回 636.2s  block 后 636.2s (10.60 分钟)
RESULT 叶子 282 个  全局 1558.5 GiB
RESULT 加载后 HBM ['50.3' ×8]  RSS 峰值 55.4 GiB
RESULT 本 host 读吞吐 ≥ 0.61 GiB/s

10.6 分钟 —— §12.2 提到的 docs/analysis/gcs-weight-loading.md §5 判据是"直读 > 5 分钟才值得建 Orbax 快照",这个数落在需要快照的那一侧。

两条独立的算术对上了,这是目前对"真权重到底多大"最强的交叉验证。 完全不同的两条代码路径:一条把字节真搬进 HBM 之后数叶子,一条只读 safetensors 头部做直方图(k3serve/tools/check_real_ckpt.py,服务机上跑):

dtype 张量数 字节
BF16 1,954 105.67 GiB
F32 506 0.04 GiB
U8(mxfp4 packed + scale) 494,592 1347.12 GiB
  • 1347.12 + 0.04 + 105.67 = 1452.83 GiB = checkpoint 的字节总量。
  • 1347.12 + 0.04 + 2×105.67 = 1558.50 GiB = JAX 线量到的叶子和,逐位对上。差的正好是 BF16 那一份被存成了 fp32。

第二条不只是对账,它在真权重上验证了 dtype 政策:自建栈"默认 fp32、唯一例外是 routed experts 的 mxfp4"这条规矩,在 1.45 TB 的真数据上确实是这么执行的 —— 494,592 个 U8 张量原样进 HBM 没被解开(解量化发生在前向里、shard_map 窄化到本地专家之后),而 1,954 个 BF16 张量各自翻了一倍。

静态校验四项全绿(服务机,check_real_ckpt.py,退出码 0):

index 里 497220 个键,mapper 之后剩 497052 个(剔除 168 个 vision_tower./mm_projector.)
✅ 专家 weight_packed 改名没有触发 ValueError
非专家参数:checkpoint 2367 个(折叠前 2460 个),模型 2367 个,缺失 0,多余 0,形状不一致 0
专家参数:checkpoint 端 494592 条 key(92 层 × 5376),模型端 368 个(stacked)张量(92 层 × 4)

5376 = 896 专家 × 3 个 w × 2(weight_packed + weight_scale)368 = 92 × 42460 − 93 = 2367(93 处 gate_proj+up_projgate_up_proj 融合)—— 每个数都能被结构解释。

这一项通不过的时间比通过的时间长得多,卡了两轮,两次都不是形状逻辑的问题:先是干跑没按 tpu_runner 的顺序装配 quant config,后是 K3CompressedTensorsConfig 借了 torchtpu-vllm 的 MoE 方法却没借它的 LinearBase 规则,导致真 config 的排除式量化声明(targets:["Linear"] + ignore 名单)把 routed_expert_down_proj 误判成量化层,撞上 TPU 平台的空 mxfp4 linear kernel 表。后者不是干跑专有的 —— 真加载会在同一行炸。 干跑校验的价值就在这儿:它在真权重进服务线之前把这块石头搬开了。

第一次真前向 + 生成。 93 层全部实际参与计算,产出 logits,贪心生成 4 步:

[K3_GEN] prompt_ids=[163584, 1, 2, 3, 4, 5, 6, 7]
[K3_GEN] step=0 min=-7.994 max=15.79 mean=-0.7347 std=2.214 top1-top2=1.88532
[K3_GEN] step=1 min=-8.433 max=18.45 mean=-0.2842 std=2.101 top1-top2=1.84963
[K3_GEN] step=2 min=-9.426 max=19.25 mean=-0.3158 std=2.056 top1-top2=2.16082
[K3_GEN] step=3 min=-9.025 max=18.54 mean=-0.4968 std=2.144 top1-top2=1.53508
[K3_GEN] median_top1_top2_gap=1.86748
[K3_GEN] tokens=[163584, 1, 2, 3, 4, 5, 6, 7, 2494, 9, 10, 9422]

四个 pod 的 tokens= 和中位间距逐字节相同

这组数解掉了 §12.5 里挂着的一条前提。 tiny 随机权重上 1120 个解码步里有 1014 个(90.5%)top-1 与 top-2 的 logits 完全相等,argmax 落在哪个 token 上跟切分对错无关;真权重上 4/4 步都不平局,间距中位数 1.867,而 bf16 在 ~19 这个量级上的分辨率约 0.06。logit 场塌成几个可表示档位是 make_tiny_ckpt.py 随机权重的产物,不是 K3 的性质 —— "逐 token 全等"这条 TP 判据在真权重上是有信息量的,那条路线没有白设计。

这里没有性能数字。 逐步耗时 170~192 s 不是解码延迟:k3/generate.py 的贪心循环是 Python for、每步各自展开编译,这个数里编译/dispatch 占多少没有拆开量过。不要引用它。

12.6 服务线真权重端到端跑通(v7x-16,2026-08-06)

§12.4 是自建 JAX 栈在真权重上算出 logits。这一节是生产服务线(vLLM + torch_tpu + torchtpu-vllm)在同一份 checkpoint 上跑完整的推理请求。两者是 不同的栈、不同的代码路径、不同的判据。

结果。 v7x-16(4 pod / 32 device),TP=32 + EP,--load-format k3_gcs

Processed prompts: 100%|██████████| 35/35 [21:36<00:00, 37.04s/it]
[v7x-serve] 退出码 0,生成 35 条 ✅
阶段 时间 数字
权重加载 13:54:00 → 14:03:14(9 分 14 秒 1453.7 GiB → 45.43 GiB/device 常驻
KV cache 14:03:14 26191 块 / 24.38 GiB per device
生成 35 条 × 32 token → 14:33:57(21 分 36 秒) 见下面「这不是性能数字」

判据不是「出了 token」,是输出的事实内容对不对。

§12.5 一直挂着一条:随机权重上判不了数值(90.5% 的解码步是 bf16 平局)。 真权重把这条前提解除了 —— 35 条输出里事实类的全对

prompt 输出
The capital of France is Paris. The population is approximately 67 million people…
The colors of the rainbow are red, orange, yellow, green, blue, indigo, and violet.
How many players are on a standard soccer team on the field at one time? 11 players (including the goalkeeper)
尼罗河最长 / 肾脏过滤废物 / 金刚石最硬 / 瑞士法郎 / 鳄梨是 guacamole 主料 全对

(有几条续写漂成 quiz 或代码格式,那是 base model 对无 chat template 短 prompt 的正常行为,不是数值问题。)

为什么这比「35/35」强得多:语义正确是一个跨越全部 93 层、同时覆盖 mxfp4 打包态常驻 + 前向解量化、MLA、KDA、MoE 路由、TP=32 分片的端到端数值 判据。任何一处切错或解错,出来的会是通顺但事实错乱的文本,或者干脆是 乱码 —— 这两种失败模式都不会被"能不能出 token"抓到。

显存账最终落点。 §12.1 的 mxfp4 那条到这里收口:

判断 结果
全 bf16(上游 process_weights_after_loading 的默认行为) 5.06 TiB ÷ 32 = 161.8 GiB/device ❌ 真机确认 OOM
打包态常驻、解量化挪进前向 预估 ~45 GiB/device 实测 45.43
KV cache 实得 24.38 GiB/device

45.43 是从 96 个 safetensors 的头里逐张量算出来的(1453.7 GiB ÷ 32), 不是估的。vLLM profile 给 KV 留了 24.38 GiB,说明权重+激活占约 60.9 GiB, 96 GiB 的 device 仍有余量。

这不是性能数字,别引用它。 21 分 36 秒 / 0.86 tok/s 绝大部分是第一次 运行的编译--enforce-eager 下 93 层的每个算子各自编译,前两条 prompt 花了 14 分 31 秒,后 33 条只用了 7 分 05 秒。这个数既不是延迟也不是吞吐, 它唯一说明的是"编译能编完、算完能出正确答案"。稳态解码延迟仍然是空白 —— 跟 §12.5 第一条说的一样。

这条路上清掉的八道坎。 每一道都不是调参数能过的,记在这里是因为它们 构成了"可行性"这个问题的实际成本:

# 卡点 解法
1 上游 mxfp4 在加载时解成 bf16,161.8 GiB/device 装不下 打包态常驻,解量化挪进前向(k3serve/ops/moe_mxfp4.py
2 不开 EP 时,加载途中 HBM 就满(常驻够、瞬态不够) --enable-expert-parallel
3 加载 48 分钟:每个 rank 都读了全部 896 个专家 接上上游的 EP 权重过滤,并显式加 --enable-ep-weight-filter(上游默认
4 容器 /dev/shm 只有 64 MB,Ray Compiled Graph 塞不下 remount 到 32G
5 --num-gpu-blocks-override 4096(给 tiny 调的)对真 K3 要 248 GB 留空,让 vLLM 自己 profile
6 Too many leaves for PyTreeDef:EP 分支里的一维常量被抬成额外输入 换成 jnp.zeros((1,),…)(走 broadcast_in_dim,不进 constvars)
7 MLA 块大小漏了 kernel 内部的 packing 对齐(96 头 ÷ TP32 = 3,kernel 眼里是 4) align_to(H, q_packing)
8 RAY_CGRAPH_get_timeout 默认 300 秒,第一次 execute 编不完 抬到 7200

这八条里有六条的共同形态:报错指向的地方不是出问题的地方。 #4 崩在 gloo 的 "Connection closed by peer"、#5 崩得像"权重太大"、#7 是一句没有消息体 的 AssertionError、#8 长得像网络故障。加上一次运维事故(节点 ephemeral-storage 只有 43.8 GiB,三轮不清 /tmp/ray 写到 67 GiB,pod 被 kubelet 驱逐、JobSet maxRestarts=0 不自愈,整个 v7x-16 得手工重建)—— 移植 K3 到 TPU 的实际工作量,诊断占的比重远大于写代码。 这一条对 "可行性"的判断比任何单项性能数字都重要。

12.5 仍然是分析值 / 仍未验证

诚实地列出来,这些不要当成已验证:

  • v7x 上仍然没有一个稳态的吞吐/延迟数字。 §12.6 那次真权重端到端跑完了(35/35,21 分 36 秒),但那个数绝大部分是第一次运行的编译 —— 前 2 条 prompt 用了 14 分 31 秒,后 33 条只用 7 分 05 秒。它证明"能跑通",不证明"跑多快"。除此之外有的还是两个点:⑩ 的 a2a per-call(9.88 µs),和加载耗时。稳态整步 decode 要多久,仍然不知道 —— §12.4 那次生成的逐步 170~192 s 是 Python for 循环每步各自编译的产物,不是延迟。MoE 的 ragged_dot tile 参数是照 v6e 的 VMEM 预算手调的、在 v7x 上没验,第 5 章那些 roofline 数字全部还是分析值。v7x 的性能不能从 v6e 外推(MXU 是另一代),也不能从"kernel 能跑"外推。
  • ~~真权重(1.553 TB)没有加载过。~~ 已验,见 §12.4 —— 1452.83 GiB 真进了 HBM,93 层算出 logits,静态校验四项全绿。第 5 章的显存账现在是称出来的(两条独立路径逐位对上)。~~仍然没做的是服务线上的真权重。~~ 已验,见 §12.6 —— 生产服务线(vLLM + torch_tpu)在 v7x-16 / TP=32 + EP 上跑完 35/35,输出事实正确,常驻 45.43 GiB/device。(k3serve-v6e-4 依然放不下 1.45 TB,服务线的真权重只能在 v7x-16 上跑,而它不在 verifier 的自动流程里。)
  • ~~「184 次小消息 a2a 的 per-call 开销」还是没测。~~ 已测,见 ⑩ —— 9.88 µs(串行)/ 2.35 µs(重叠)。但它把未知量挪了个位置,没有消掉:主线 torchtpu-vllm 的 MoE 不发 a2a,走的是 all-reduce combine,而真实解码路径一步到底有多少次 collective、各多大、能不能跟计算重叠,一件都没查。这是接替它的最高优先级未知量。同一次 sweep 里合成 psum 的串行数是 16.3~18.8 µs,那不是这个问题的答案,别拿它乘层数。
  • tokamax#1103 仍未合入上游,PyPI 发布版里没有 kda/。报告第 10 章那条风险依然挂着。
  • ~~自建 JAX 栈的多 host 计算层刚验到 KDA,整模型还没过。~~ 已解决,而且根因不在分片上。 当时的现象是整模型在 16 设备 mesh 上对拍单设备红着:0.43% 的 logits 超出 atol=1e-5, rtol=1e-3,最大绝对偏差 0.0172,比 fp32 重结合该有的量级大两个数量级;argmax 全程一致,所以不是路由分岔、不是切错。先怀疑的「参照物没钉精度」(default_matmul_precision("highest"))被重跑逐位证伪。逐层二分最后定位到测试自己的合成权重_synthetic_state_dict 对每个参数一律 N(0,1)、没按 fan-in 缩,hidden_size=1024 的权重配 N(0,1) 输入做矩阵乘,激活几层之内飘到 ~1e4(真 checkpoint 是 O(1)~O(10)),fp32 在那个量级只剩 1e-3 的绝对分辨率 —— 单设备与 16 设备的求和结合顺序差异就被放大成 1e-4 相对量级的「失败」。是数值噪声,不是切分 bug。 修法是把真正参与矩阵乘/卷积求和的权重按 1/sqrt(fan_in) 缩(embedding 是查表、1D 参数没有 fan-in,都不缩),判据一个字没动 —— 改的是不具代表性的输入,不是容差。现在 pytest tests/multihost 整目录在 v6e-16 上是绿的。

这一条跟 §12.4 里「tiny 随机权重上 90.5% 的解码步是 bf16 平局」是同一课上了两遍:随机权重不是「没有信息的权重」,它是有误导性的权重 —— 既能把数值噪声放大成看起来像 bug 的东西,也能把真 bug 藏进一片平局里。凡是拿合成权重当输入的判据,都要先问一句「这个量级像不像真模型」。 - "每台 host 只读自己那份"这件事,两条线的状态不一样。 服务线已经做到了(§12.6):接上上游的 EP 权重过滤之后每个 rank 只读 1/ep_size 份专家,实测每 host 读的字节数从 11,630 GiB 降到 1,080 GiB(1/10.8),墙钟 48 分钟 → 9 分 14 秒。自建 JAX 栈那条还没接 —— loader 今天走 jax.device_put(host_local_numpy, ...),多进程下不抛异常、数值也对,但要求每台 host 在内存里放下完整的那个张量,而且 4 台机器各做一遍同样的 I/O。 - ~~端到端生成只在随机权重的 MINIMAL 配置上做过。~~ 已验,见 §12.6 —— 真词表、真层数(93)、真权重、TP=32,35 条 prompt 全部生成完成且事实内容正确。判据形式确实跟着换了:从"逐 token 全等"换成"输出的事实内容对不对",理由见下一条。 - "切分对不对"这件事,在 tiny 随机权重上根本判不了。 vLLM 那条线上想用"TP=1 与 TP=4 逐 token 全等"当判据,实测发现 90.5% 的解码步是 bf16 平局(top1 与 top2 的 logits 在 bf16 下相等),argmax 落在哪个 token 上取决于比较顺序,跟切分对错无关。真权重上这条前提解除了(§12.6):35 条输出的事实内容全对,而语义正确本身就是一个跨越 93 层、覆盖 MLA + KDA + MoE + mxfp4 解量化 + TP=32 分片的端到端数值判据 —— 任何一处切错或解错,出来的会是通顺但事实错乱的文本。

不过原设计的那条判据(TP=1 vs TP=4 逐 token 比对)在真权重上做不了:45.43 GiB/device 的常驻只有 32 device 装得下,TP=1/4 根本起不来。替代方案是同一 TP 下开/关 EP 对拍(两者数学等价、形状和通信完全不同),以及固定 prompt 集的事实正确率 + 可复现性。


本报告主体(第 2–10 章)基于静态代码分析与手工 roofline 推算;第 12 章是真机实测补记(v6e-4 单机 + v6e-16 四机 + v7x 16 芯四机,最后更新 2026-08-06),逐条标注了哪些修正了主体结论、哪些仍是分析值。真权重(1453.7 GiB)已在 v7x-16 上跑通生产服务线的端到端推理:35/35,输出事实正确,常驻 45.43 GiB/device(§12.6);自建 JAX 栈那条也已加载并算出 logits(§12.4);a2a per-call 已实测(§12.1 ⑩)。仍然没有的是稳态 decode 延迟 —— 现有那个 0.86 tok/s 绝大部分是首次编译,不能当性能引用。 参数模型已用官方公布的 2.8T/104B 交叉校验。涉及的仓库均在快速迭代中,代码引用对应 2026-07 的状态。