Kimi K3 移植到 Google TPU 生态的可行性分析¶
日期:2026-07-28 | 目标硬件:TPU v7x (Ironwood) 主选,v6e (Trillium) 次选 | 范围:推理移植 + 服务化成本
1. 执行摘要¶
Kimi K3 是 2.8T 总参 / 104B 激活的混合注意力 MoE 模型,93 层中 69 层是 KDA(Kimi Delta Attention,线性注意力),24 层是 NoPE 版 Gated MLA,MoE 部分为 896 experts top-16 的 LatentMoE 结构,权重以 MXFP4 做过 QAT。
结论:技术上可行,且 K3 的架构特征相当适配 TPU。核心工作量比预想的小 —— torchtpu-vllm 已经把大部分基础设施做完了,真正要写的只有一处。
五条主要判断:
-
显存不是瓶颈,拓扑下限是 16 chips。 权重实测算下来 1.553 TB。注意 MXFP4 只覆盖 routed experts(量化配置的
ignore列表排除了 attention、shared experts、dense MLP、lm_head、vision),所以是「1.449 TB MXFP4 + 104 GB BF16」的混合账。v7x 单 chip 192 GB,2×2×2 拓扑(8 chips,1.536 TB)装不下,最小可行是 16 chips。 -
长上下文是 K3 在 TPU 上最大的优势。 只有 24 层全注意力,采用 absorbed MLA 后 1M context 的 KV cache 仅 14.5 GB/序列。作为对照,DeepSeek-V3 的 61 层 MLA 在同等上下文下需要约 36 GB。K3 更长的上下文反而更省。
-
短上下文下瓶颈会反转。 KDA 的循环状态是 434 MB/序列且与序列长度无关。在约 32.5K token 以下,它比 KV cache 还占显存,直接压制并发批量。这是 TPU 侧容量规划最容易被忽略的一项。
-
没有需要从零写的 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)逐条落在其支持范围内。剩下的是接线工作:布局转置、betasigmoid 外提、状态布局对齐。估 5–9 天。 -
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.json 与 modeling_kimi_linear.py 直接读出,不引用 model card 的概述值。
2.1 层结构:69 KDA + 24 MLA 的严格交错¶
linear_attn_config.full_attn_layers 明确列出全注意力层为 [4, 8, 12, ..., 88, 92, 93],其余 69 层为 KDA。规律是每 4 层插入 1 层全注意力,末尾 92、93 连续两层。第 1 层是 dense MLP(first_k_dense_replace: 1),其余 92 层为 MoE。
这个 3:1 的比例是 K3 长上下文经济性的根源 —— 只有 24 层需要维护随长度增长的 KV cache。
2.2 KDA:低秩衰减门 + 全秩输出门¶
单层结构(KimiDeltaAttention):
| 组件 | 形状 | 说明 |
|---|---|---|
q_proj/k_proj/v_proj |
7168 → 12288 | 96 heads × 128 |
q/k/v_conv1d |
kernel=4, depthwise | SiLU 激活的短因果卷积 |
f_a_proj → f_b_proj |
7168 → 128 → 12288 | 低秩衰减门(rank=128) |
b_proj |
7168 → 96 | delta rule 的 β(每头标量) |
g_proj |
7168 → 12288 | 全秩输出门(use_full_rank_gate: true) |
o_norm |
FusedRMSNormGated | sigmoid 门控 |
o_proj |
12288 → 7168 |
合计 443.7 M 参数/层,69 层共 30.62 B。
递推式本身是标准的 gated delta rule:
与 MaxText 现有实现的结构性差异
MaxText 的 attention_kda.py 用全秩 g_proj(7168→12288)生成衰减门,而 K3 用的是低秩 f_a_proj(7168→128) → f_b_proj(128→12288)。两者参数量差 86 M/层,权重不能直接映射,需要改写投影层结构。这是一个容易在权重转换阶段才暴露的坑。
2.3 MLA:NoPE 变体,带输出门¶
modeling_kimi_linear.py 里 KimiMLAAttention.__init__ 有一句硬断言 assert self.use_nope,且 self.rotary_emb = None —— K3 的全注意力层完全不用 RoPE。
qk_rope_head_dim: 64 这个字段仍在,但它退化成一段不做旋转的 MQA 共享尾维:k_rot 从 kv_a_proj_with_mqa 切出后直接 expand 到所有头。位置信息全部由 KDA 层的短卷积与递推结构隐式承载。
其余维度沿用 DeepSeek-V3 风格:q_lora_rank=1536、kv_lora_rank=512、qk_nope_head_dim=128、v_head_dim=128。额外多了 mla_use_output_gate: true 带来的 g_proj(7168→12288)。单层 232.2 M 参数,24 层共 5.57 B。
HF 参考实现的 KV cache 不可用于生产
KimiDynamicCache.update() 缓存的是物化后的 key_states [B, 96, T, 192] 与 value_states [B, 96, T, 128],合计 737,280 元素/token。1M context 下这是 1.41 TB/序列,完全不可行。
必须实现 absorbed MLA:把 kv_b_proj 吸收进 q 侧与 o 侧,只缓存压缩 latent kv_lora_rank + qk_rope_head_dim = 576 元素/token。差距 1280 倍。这不是优化,是可行性前提。
2.4 LatentMoE:把专家搬到低维空间¶
KimiSparseMoeBlock 的三段式结构是 K3 相对 K2 的主要改动:
hidden 7168
→ routed_expert_down_proj → latent 3584
→ [896 experts, 每个 3584 → 3072 → 3584] top-16
→ routed_expert_norm (RMSNorm)
→ routed_expert_up_proj → hidden 7168
+ shared_experts(hidden 7168 → 6144 → 7168) # 旁路,走原始 hidden
单个 routed expert = 3 × 3584 × 3072 = 33.03 M 参数,896 个共 29.6 B/层,92 层共 2.72 T —— 占全模型 98%。
路由是 noaux_tc + sigmoid + grouped topk(num_expert_group: 1、topk_group: 1,即实际未分组),routed_scaling_factor: 1.0,moe_renormalize: true。
LatentMoE 对专家并行是利好
专家在 3584 维而非 7168 维上工作,意味着 EP 场景下 all-to-all 的载荷直接减半。考虑到 92 层 MoE 每层两次 a2a,这个设计对分布式推理的通信开销影响很大。第 6.3 节会量化。
2.5 量化:MXFP4 的实际覆盖面¶
quantization_config 关键字段:
{ "format": "mxfp4-pack-quantized", "quant_method": "compressed-tensors",
"weights": { "num_bits": 4, "group_size": 32, "symmetric": true,
"strategy": "group", "scale_dtype": "torch.uint8" },
"ignore": ["re:.*self_attn.*", "re:.*shared_experts.*",
"re:.*mlp\\.(gate|up|gate_up|down)_proj.*",
"re:.*lm_head.*", "re:.*vision_tower.*", "re:.*mm_projector.*"] }
ignore 列表是本报告最关键的单条发现。 它意味着 MXFP4 只作用于 routed experts(以及未被 regex 命中的 latent 投影),attention、shared experts、dense MLP、lm_head、视觉塔全部保持 BF16。
scale_dtype: uint8 + group_size: 32 是标准 MX 规格:每 32 个 4-bit 元素共享一个 E8M0 指数,等效 4.25 bit/元素。
2.6 其余组件¶
- SiTU 激活:
beta·tanh(gate/beta)·sigmoid(gate)·up,beta=4.0。纯 element-wise,TPU 上无障碍,但需注意tanh与sigmoid走 EUP 流水线。 - AttnRes:
attn_res_block_size: 12,以 12 层为一块做跨块残差累积(_apply_attn_res用 prefix_sum + 投影 + norm)。改变了残差图的拓扑,会影响 remat/scan 的切分策略。 - MoonViT-V2:27 层、hidden 1024、12 头、patch 14、
merge_kernel_size [2,2]、patchmergerv2投影到 7168。约 400 M 参数,BF16 不量化。
3. TPU 侧现有能力盘点¶
本章逐一核对了以下仓库的实际代码(均为 2026-07 的状态):
| 仓库 | 内容 | 备注 |
|---|---|---|
primatrix/pallas-kernel |
KDA/MLA 的 Pallas kernel,包名 tops |
MaxText 通过 pyproject.toml pin 在 v0.5.4-rc2 |
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] |
✅ |
四处需要适配的差异
- 没有单步解码 kernel。 这个 PR 只有分块前向 + 反向 + XLA 参考,decode 路径仍要用 torchtpu-vllm 的
gdn/v2。 - CP 模式下禁用
initial_state与output_final_state。 意味着 CP 与"带状态接续的分块 prefill"目前不能组合 —— 对 1M context 的分段 prefill 是实际限制。 beta需在 kernel 外先做 sigmoid。 K3 用的是use_beta_sigmoid_in_kernel=True,tokamax 无此开关(测试里由调用方jax.nn.sigmoid(beta))。开销可忽略([H,B,T]而已)。- 布局是 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.py、custom_ops/mamba_state_copy_op.py |
循环状态的 prefix caching 与 slot 搬移 | KDA 状态的 cache 管理 |
distributed/kv_transfer/、offload/cpu_tpu.py |
KV 传输、host offload、P/D 分离 | 见 3.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 jnp 和 from jax.sharding import PartitionSpec as P。所以「JAX 路线 vs PyTorch 路线」是个伪命题 —— torchtpu-vllm 的模型定义来自 vLLM 上游的 PyTorch 实现,kernel 却是 JAX/Pallas,两者在 jax_op 这一层缝合。第 7 章会展开这一点。
模型定义本身不在这个仓库(models/ 只有一个 wrapper context),走的是 vLLM 上游的 ModelRegistry。意味着只要 vLLM 上游支持某个模型、且所需 kernel 在这里有覆盖,就能跑 —— 不需要像 JAX 路线那样重写模型。
3.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.py、offload/cpu_tpu.py、platforms/tpu_platform.py 都引用了 raiden,examples/disagg/ 下有完整的 P/D 分离多机脚本。
对 K3 的价值:1M context 下每序列 14.5 GB KV,把冷序列 offload 到 host DRAM(v7x 每 host 约 944 GB)能大幅提升并发上限;P/D 分离则让 prefill(算力瓶颈)和 decode(带宽瓶颈)各自用最优拓扑 —— 对 K3 这种两个阶段瓶颈截然不同的模型收益尤其大。
3.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):三层都要补 —— tops 的 fused_recurrent_kda 只是 lax.scan(每步把 434 MB/序列状态读出写回,无法融合门变换与 delta 修正);没有状态 cache 抽象;attention_kda.py:369 直接抛 NotImplementedError。20–30 天。
torchtpu-vllm 路线:kernels/gdn/v2/gdn_decode_kernel.py 已经是一个完整的融合单步解码 Pallas kernel,而且参数契约与 K3 的 KDA 高度吻合:
def ...(q, k, v, g, # g: Per-key gating [T, H_v, K] ← 逐通道
initial_state, # [num_states, H_v, K, V] 状态 cache
state_indices, # i32[max_num_req] slot 索引
distribution, # i32[3] (decode_end, prefill_end, mixed_end)
b, *, scale,
use_qk_l2norm_in_kernel=False, # ← K3 的 Q/K L2 norm
use_gate_in_kernel=False, # ← 门变换融合
A_log=None, dt_bias=None, lower_bound=None, apply_silu=False)
这几个 flag(use_qk_l2norm_in_kernel / use_gate_in_kernel / lower_bound)与 tops.chunk_kda 的签名几乎逐字对应,两边显然是相互知情地开发的。配套的状态管理也齐了:mamba_state_copy_op.py 做 slot 搬移,runner/mamba_apc.py 做状态的 prefix caching,distribution 参数原生支持 prefill/decode 混批。
剩下的缺口只有一处,见 G2。
G2 — 分块 prefill 路径的衰减门只到 per-head | 本项目的核心缺口 | 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_cumsum 后 jnp.exp(g_cumsum)[..., None] |
逐头标量 ❌ |
torchtpu-vllm gdn/v3 |
gating_log: [1, 1, num_v_heads] |
逐头标量 ❌ |
K3 的 KDA 是 g = f_b_proj(f_a_proj(h)) 再 rearrange('... (h d) -> ... h d', d=head_dim),即每个头内部逐通道的对角衰减矩阵(GLA/DeltaNet 风格)。而 torchtpu-vllm 的分块 prefill 路径是为 Qwen3-Next 的 GDN 设计的,衰减是每头一个标量(Mamba2 风格)—— 代码里 k_beta * jnp.exp(g_cumsum)[..., None] 的那个 [..., None] 广播就是证据。
好消息是:分块算法的骨架完全一样(log 空间 cumsum → chunk 内衰减差矩阵 → WY 表示 → 状态传递),要改的只是把衰减的秩从标量提到向量。而逐通道版本已经有两份现成实现,都是 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_attention 经 custom_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_tpu 有 float8_e8m0fnu dtype,kernels/quantized_matmul/blockwise_kernel.py 提供 block-wise 量化矩阵乘。
在这条路线上,剩下的工作只是验证现有实现能否吃 K3 的 compressed-tensors / mxfp4-pack-quantized 打包格式,权重按 1.553 TB 规划即可。5–8 天。
JAX 路线上这一项是硬冲突
第 3.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_tpool、pos_emb_type: divided_fixed、视频时序 patch 需要重做。torchtpu-vllm 路线下有 vision_attention.py 与 multi_modal_inference.py 示例可依托。
G8 — SiTU 与 AttnRes | 低 | 2–4 天 / ~2 天¶
SiTU 是纯 element-wise。AttnRes 的 12 层分块残差会改变 scan/remat 的切分边界 —— JAX 路线要在 decoders.py 里调整;PyTorch 路线只是模型定义里的加法,几乎无成本。
工作量汇总¶
| Gap | JAX 路线 | torchtpu-vllm 路线 |
|---|---|---|
| G1 KDA 解码 | 20–30 天(阻塞) | 已有,0 |
| G2 逐通道分块 prefill | 已有(tops / 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 递推的数值稳定性 —— tops 的 fused_recurrent_kda 内部用 FP32 累加器,暗示精度敏感。
5.3 拓扑下限¶
按权重 + 15% 系统开销(XLA 编译缓冲、激活、碎片):
| 拓扑 | chips | HBM 总量 | 权重占比 | 余量 | 1M ctx 并发 |
|---|---|---|---|---|---|
| 2×2×2 | 8 | 1.536 TB | 101.1% | −17 GB | 装不下 |
| 2×2×4 | 16 | 3.072 TB | 50.6% | 1519 GB | ~86 |
| 2×4×4 | 32 | 6.144 TB | 25.3% | 4591 GB | ~261 |
| 4×4×4 | 64 | 12.288 TB | 12.6% | 10735 GB | ~611 |
v6e(32 GB/chip)需要约 56 chips 才装得下权重,且带宽只有 v7x 的 22%。v6e 不适合 K3,仅可用于单层/子模块的算子开发验证。
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):
- 专家权重 1.449 TB / 16 = 90.6 GB/chip,均匀无冗余
- Attention 权重 104 GB,DP=2 下每组复制一份 → 每 chip 13 GB
- KV cache 按 DP 组切分,每组独立管理,避免 TP 复制
- KDA 状态按头切分:434 MB / 8 = 54 MB/chip/序列
v7x × 32(2×4×4):
推荐 32 chips 作为生产配置:显存余量充裕(权重占 25%),并发可到 261 路 1M context,且 MXFP4 若暂不可用、退回 block-FP8(2.83 TB)时仍装得下。
6.3 通信量估算:ICI 带宽不是瓶颈¶
EP 的 all-to-all 是主要通信开销。每 token 每 MoE 层:
- dispatch:latent 向量 3584 元素,发往 top-16 专家所在的 chip
- FP8 载荷:3584 B × 16 = 57.3 KB,combine 同量 → 114.7 KB/层
- 92 层 → 10.55 MB/token 往返
对照:传统 MoE 在 7168 维上做 a2a 就是 21.1 MB/token。LatentMoE 让通信量正好减半,这是 2.4 节那个架构观察的量化形式。
按 v7x 每 chip 1200 GB/s 双向(单向可用 600 GB/s)、EP 铺满全部 chip 的 per-chip egress 模型:
Decode(每步)
| 配置 | B=1 | B=32 | B=256 |
|---|---|---|---|
| v7x × 16 — ICI 时间 | 1.1 µs | 35.2 µs | 281.4 µs |
| v7x × 16 — 占 HBM 时间 | 0.1% | 0.3% | 2.1% |
| v7x × 32 — ICI 时间 | 0.5 µs | 17.6 µs | 140.7 µs |
| v7x × 32 — 占 HBM 时间 | 0.1% | 0.3% | 2.1% |
| v6e × 64 — 占 HBM 时间 | 0.0% | 0.1% | 0.7% |
Prefill(稳态)
| 配置 | 吞吐 | 每 chip ICI | 占 ICI 容量 |
|---|---|---|---|
| v7x × 16 @ MFU 30% | 53,141 tok/s | 35.0 GB/s | 5.8% |
| v7x × 32 @ MFU 30% | 106,283 tok/s | 35.0 GB/s | 5.8% |
| v7x × 32 @ MFU 45% | 159,424 tok/s | 52.6 GB/s | 8.8% |
结论:ICI 带宽有一个数量级的余量,EP 方案在带宽维度上完全成立。 即便按 3D torus 的 bisection 约束保守打 2× 折扣,decode 也只到 4%、prefill 到 18%,结论不变。
这个结论值得强调,因为 92 层 MoE + top-16 的 a2a 直觉上像是个大问题。真正的问题在下一节。
6.4 真正的约束是 collective 延迟,不是带宽¶
上表有个细节值得盯:B=1 时每次 collective 只搬 1.79 KB/chip,传输耗时 3 ns。 这个数量级下带宽完全无关,全部时间都是 collective 的固定开销。
每步要做 2 × 92 = 184 次 all-to-all。按 v7x × 32、batch=1 的 HBM 下界 550 µs 对照:
| 单次 collective 开销 | 184 次合计 | 占 HBM 时间 |
|---|---|---|
| 1 µs | 184 µs | 33% |
| 2 µs | 368 µs | 67% |
| 5 µs | 920 µs | 167% |
低延迟解码路径的真实风险
如果单次 a2a 的固定开销落在 2–5 µs,batch=1 解码会从带宽瓶颈变成通信延迟瓶颈,延迟翻倍甚至更多。92 层的深度在这里是放大器 —— 同样的 per-collective 开销,层数少一半风险就减半。
所以真正需要实测的量不是 ICI 带宽,而是「小消息 all-to-all 的 per-call 开销」。 这是个便宜很多的测量:不需要 K3,拿一个 92 层的 dummy MoE 就能测。建议提到 P0 阶段做掉。
实测补记:做掉了,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-rc2的CPContext,attention_kda.py里有硬断言保护 —— 缺失时拒绝运行,因为"CP 会静默破坏跨 rank 的循环状态"a2a(Ulysses):序列↔头维度互换,attention_kda.py里把 q/k/v/g 融合成单次 all-to-all
KDA 在 CP 下的难点是递推状态的跨 rank 传递,tops 有专门的设计文档(docs/design-docs/ops/kda/global_cp.aligned.zh.md、kda-chunk-bwd-cp.zh.md)。短卷积还需要 halo exchange(cp_a2a.halo_exchange_for_conv)拉取前一个 rank 的 kernel_size-1 个 token —— 否则每个 CP 边界的因果卷积都是错的。
推理场景下这部分是现成的,因为 CP 的复杂度主要在反向传播,前向已实现。
6.7 不推荐 Pipeline Parallelism¶
93 层理论上可切,但推理场景下 PP 引入的 bubble 会直接抬高 TTFT,且 AttnRes 的 12 层分块残差会与 PP 切分边界冲突。单 slice 内用 EP + TP 已足够,除非要跨 slice 扩展。
7. 技术路线选型¶
习惯上会把这件事描述成"JAX 路线 vs PyTorch 路线"的二选一。核对代码之后,这个框架并不成立。
7.1 关键事实:两条路线共用同一批 Pallas kernel¶
torchtpu-vllm 的 GDN / MLA / MoE kernel 全部是 JAX + Pallas 写的,通过 torch_tpu._internal.pallas.jax_op 包装成 torch custom op:
# src/vllm_torchtpu/layers/common/gdn_attention.py
import jax
import jax.numpy as jnp
from jax.sharding import PartitionSpec as P
from vllm_torchtpu.kernels.gdn.v3 import wrapper as gdn_v3_wrapper
所以真正的分界线不在"JAX 还是 PyTorch",而在模型定义与服务层用哪一套:
| 层 | JAX 路线 | torchtpu-vllm 路线 | 共用? |
|---|---|---|---|
| 公共算子库 | 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¶
理由有三:
- 工作量约为一半,且关键路径(G2)已经从"写 kernel"降级为"接线" —— 逐通道分块前向在 tokamax 和 tops 各有一份现成实现,
torch_tpu调用 tokamax 的机制也已验证过。 - MXFP4 能用,权重 1.553 TB 而非 2.83 TB。这直接决定 16 chips 是否可行,也给 block-FP8 留了退路。
- 战略方向一致 —— 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_attention 经 pallas.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-type(TPU-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.py 的 check_inputs_support 在任何预处理之前就硬性检查 segment_ids is None and seq_len % chunk_size != 0 并直接 NotImplementedError。common.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=65 比 T=64 贵 19.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 µs。429 vs 128 —— 变长路径给 decode 用会慢约 3.3×。 它不是"以后有空再做",是方向就不对。
唯一的填法是写一个专用的单步 kernel:initial_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::_fwd 用 jnp.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_id、jax.process_index()、mesh 里的设备顺序 —— 它们两两不相等:
- 启动器给的
process_id≠jax.process_index()。实测:传process_id=0的那台机器,process_index()是 2。这是个非恒等置换,两轮独立运行完全一致(不是随机的)—— TPU 上进程编号由运行时按物理拓扑定,process_id只用来跟 coordinator 握手。这跟 vLLM 那条线的"ray rank 与地址表排序不一致"是同一类陷阱的翻版。 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.py的fused_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
FusedMoE的reduce_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,已经排进队列。
方法上有四条限制,抄这些数的人必须一起抄走:
- N < 1000 量不出重叠。
unchained N=100是 9.4 µs,跟串行一样;到 N=1000 才掉到 2.35 µs。第一轮我只跑了 N=100,据此得出过"XLA 拿不到重叠收益"的结论 —— 那是错的,独立调用不够多时调度器摊不开。 - 这是代理指标。 量的是一维 32 设备 mesh 上的
jax.lax.all_to_all,K3 真实的 dispatch 是 ragged 的、走shard_map。它量准的是"这个 collective 的固定开销有多大",不是端到端的 MoE 通信时间。 ep_degree=64那个假设没验。 1792 B 是按 EP=64 反推的,而我们在 32 设备上量。因为成本对字节数不敏感,EP=64 时字节减半不会让它变快;参与设备翻倍反而大概率更慢。没测,别推。- 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_tpuvs 它自带的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 栈跑贪心生成,跟 HFgenerate(do_sample=False)在同一份随机权重、同一个输入下比 token 序列:两组配置(全程xla、以及 prefill 强制走pallas_tpuchunk 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.py4 passed / 1 failed。 这是整个 v7x 路线上单点风险最高的一环,现在它是通的。 -
唯一那条红的,是我们自己的一个洞被 v7x 照出来了 —— 不是 kernel 的错。
pallas_tpuvsxla对拍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 除不尽"(MINIMAL的H=8、num_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。这一节是真 checkpoint(gs://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 × 4,2460 − 93 = 2367(93 处 gate_proj+up_proj → gate_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:
| 阶段 | 时间 | 数字 |
|---|---|---|
| 权重加载 | 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_dottile 参数是照 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 的状态。