Kimi K3 迁移到 Google TPU 的可行性评估(v2 · 2026-08-05)¶
本报告是对 2026-07-28 初版(
kimi-k3-tpu-feasibility.md,含截至 2026-08-05 的全部实测补记)的独立复核与更新。 方法:初版逐章重读;其引用的模型事实、上游代码声明、真机实测结论由独立核查逐条对回当前代码、HF 官方配置与 k3-jax 仓库的实际提交记录;11 个相关仓库已全部 fetch 到 2026-08-05 最新状态。 凡标「实测」的数字来自初版 §12 与 k3-jax 的真机记录,凡标「分析」的数字来自 roofline 推算,两者严格分开。目标硬件:TPU v7x(Ironwood)主选,v6e(Trillium)仅作开发平台。范围:推理移植与服务化。
0. 与初版报告的关系¶
初版不是纯静态分析:报告发出后配套实现项目(k3-jax)已在 v6e-4 / v6e-16 / v7x-16 真机上跑了八天
(2026-07-29 至 08-05,274 个提交),初版 §12 以「实测补记」持续吸收实测,多次推翻过自己的结论
(显存膝点、a2a 延迟、合成权重误导等)。所以本报告面对的不是「一份待验证的分析」,而是
「一份已经过一轮真机筛选、但收尾未完成的分析」。
复核结论先说:初版的架构事实与工程账基本可靠(本报告第 2、4 章逐条核验通过,代码引用抽查全部属实), 它的缺口不在「算错」而在「尚未完成」——端到端性能数字至今为零,vLLM 服务线的真权重端到端 截至 2026-08-05 上午仍在第三次尝试中。本报告的工作是把「已验证 / 未完成 / 被高估 / 被低估」重新分清, 并给出独立于初版的判断。
1. 执行摘要¶
结论:可行。「能不能跑」在 2026-08-04 后已基本被回答——真权重(1558.5 GiB 进 HBM)已在 v7x-16 上 完成加载、93 层全部参与计算并产出健康的 logits(实测)。剩下的是「跑多快、值不值」,这两个问题 目前一个数字都没有。
六条主要判断:
- 显存可行域已实测钉死。 v7x-16(16 芯 / 32 device / 实测 3031.6 GiB HBM)是最小可行拓扑,8 芯装不下。
整模型 93 层在 TP=32 下可编译运行(实测 77.59 / 94.74 GiB/device),前提是产品代码引入
nnx.scan折图—— 纯展开图在 ~36 层后越过 XLA 调度膝点,temp 从 7.7 GiB 爆到 347.85 GiB。这是初版静态分析完全没预见、 真机才暴露的一类问题。 - KDA kernel 风险已解除(初版最担心的一环)。 tokamax 的逐通道 KDA Pallas kernel 不改任何参数
即在 TPU7x 上编译通过、数值正确(4 host 一致)。prefill 侧该 kernel 比 XLA 参考快 31.6×;
decode 侧 torchtpu-vllm 有现成的逐通道单步融合 kernel(
kernels/gdn/v2,已核实存在且接口匹配)。 - 长上下文是 TPU 方案的结构性卖点,但在 v7x-16 上吃不满。 absorbed MLA 的 KV cache 实测 55296 B/token(fp32),1M context 约 14.5 GB/序列,远优于 DeepSeek-V3 量级。但 v7x-16 单 slice 上 prefill chunk 与上下文长度共享同一份 HBM 余量:32K chunk ↔ ~178K 上下文,52K chunk ↔ ~96K 上下文。 要兑现 1M context 的价值主张,生产配置应是 v7x-32 起步。
- 低并发 decode 的通信风险从 a2a 转移到了 all-reduce,仍未量化。 实测小消息 a2a per-call 9.88 µs (串行不可重叠),naive EP 的 184 次/步折合 HBM 下界的 331%——batch=1 下不可行。主线 torchtpu-vllm 的 MoE 路径不发 a2a(本地专家 + 标量偏移,combine 走 all-reduce,已核实),结构性避开;但替代路径 一步有多少次 collective、能否与计算重叠,至今未测(对应 issue #310 / PR #312,在途)。
- 初版的「27–47 天」工期估算已被日历时间证伪,但证伪方式有信息量。 07-28 至今的 8 天没有花在预估的 G2「接线」上,而是花在初版未立项的三类事:整模型编译可行性(scan 折图)、多 host 数据摆放 (process_id ≠ process_index ≠ mesh 顺序)、以及一批「v6e 绿、v7x 才现形」的平台差异(默认 matmul 精度硬件相关等)。剩余工作量应按「vLLM 服务线真权重端到端 + 性能测量 + 加载工程化」重估。
- 上游依赖的最大现实风险是 tokamax#1103 未合入且 API 在漂移。 截至 2026-08-05,KDA 代码不在
openxla/tokamax 的 main 上(origin/main 已前进到 dc54348,不含 KDA);PR head 在 08-05 当天发生
API 重命名(
chunk_size被删),直接导致一次两台机器装同名不同 API 的包的事故。依赖必须钉 commit。
2. 模型解剖(已逐条核验 ✅)¶
以 HF 官方 config.json(moonshotai/Kimi-K3,2026-08-05 实时拉取)为准,与本地 k3-jax/k3/config.py
的 K3_FULL 逐字段比对零出入;官方 model card 总参 2.8T / 激活 104B,按 config 逐层重算得
2.7795T / 104.19B,吻合。
| 维度 | 值 | 对 TPU 移植的意义 |
|---|---|---|
| 层结构 | 93 层 = 69 KDA + 24 MLA,full_attn_layers=[4,8,…,88,92,93](1-indexed;92/93 连续,不能用 i%4 反推) |
每 4 层一块(KDA×3+MLA)的结构恰好允许 scan 折图 |
| hidden / vocab | 7168 / 163840,tie_word_embeddings=false |
embedding + lm_head 双份,BF16 共 ~4.7 GB |
| KDA | 96 头 × head_dim 128,conv k=4,低秩衰减门 f_a(7168→128)→f_b(128→12288),全秩输出门,gate_lower_bound=-5.0 |
逐通道门是 kernel 选型的分水岭(见 4.2) |
| MLA | NoPE(rotary_emb=None;注意 qk_rope_head_dim=64 的通道存在、进 QK 点积、K 侧 MQA 广播,只是不旋转);q_lora 1536 / kv_lora 512 / 96 头;带输出门 |
absorbed 后 KV 仅 576 元素/token/层 |
| MoE | 896 路由专家 top-16 + 2 shared;latent 3584(down/up 投影夹在专家两侧);expert inter 3072;noaux_tc + sigmoid + renormalize | 专家在 3584 维工作,EP 通信载荷减半;98% 参数在 MoE |
| 量化 | compressed-tensors / mxfp4-pack-quantized,W4 组 32 + E8M0 scale(4.25 bit);仅 routed experts 与 latent 投影;ignore 排除 self_attn / shared_experts / dense MLP / lm_head / vision / mm_projector;无激活量化条目(W4A16) | 权重 1.45 TiB 而非 5.06 TiB 的前提 |
| 其余 | SiTU 激活(β=4.0)、AttnRes 12 层跨块残差、MoonViT-V2 视觉塔(~400M,本报告不展开) | AttnRes 影响 scan/PP 的切分边界 |
对初版的两处口径修正(不影响结论):HF 字段名是 num_experts / num_experts_per_token,不是
DeepSeek 风格的 n_routed_experts / num_experts_per_tok;「NoPE」不等于没有 rope 通道。
3. 显存与拓扑(实测为主 ✅)¶
3.1 权重账(两条独立代码路径逐位对上的实测)¶
| 部分 | 参数量 | 存储 | 字节 |
|---|---|---|---|
| routed experts + latent 投影 | 2.727 T | MXFP4 packed(U8) | 1347.12 GiB |
| attention / shared / dense / embed / head | 52.0 B | ckpt BF16;自建栈按 fp32 存,进 HBM 翻倍 | 105.67 GiB(ckpt)/ 211.34(HBM) |
| 其余零星 | — | F32 | 0.04 GiB |
| 合计(进 HBM,实测) | 2.7795 T | 1558.5 GiB |
3.2 拓扑与编译账(v7x-16 真机实测)¶
- v7x 一芯 = 2 device,各 94.74 GiB HBM(实测),「192 GB/芯」是两块的标称和、不是共享池—— 并行度按 device 数:v7x-16 = TP=32。
- 整模型 prefill 前向(
scan折图后):TP=32 下每 device 权重 50.27 + temp 11–40 GiB(随 chunk 长度), 合计 60.7–89.8 / 94.74 GiB。单次 prefill chunk 上限实测 53K token 量级(53248 装得下,57344 差 82 MiB)。 - decode 图自身 temp 29.8 GiB(对上下文长度不敏感)——「decode 很省」的直觉在 K3 上不成立, MoE 解量化那块地板两张图都要付。
- KV cache 实测 55296 B/token;KDA state 13.5 MiB/device 恒定。
- 「chunk 大小 × 上下文长度」是同一份余量的两种花法:chunk 32K → ~178K ctx;chunk 52K → ~96K ctx (由两组实测相减得到,未重新编译验证,取下限时需留余量)。
- v6e(32 GB/芯)装权重需 ~56 芯且带宽只有 v7x 的 22%,仅适合算子开发——维持初版判断。
3.3 Roofline(分析值 ⚠️ 至今无真机对点)¶
- decode batch=1:每步搬运 130 GB(注意:BF16 attention 权重 72.4 GB > 命中专家 25.8 GB), v7x-16 聚合带宽 118 TB/s → 理论 1.10 ms/步,按 50% 带宽效率 ~454 tok/s。
- prefill:每 token 208 GFLOPs,v7x-16 @ MFU 30–45% → 53K–80K tok/s(BF16)。
- ~45K token 以上上下文,MLA 注意力算力反超 MoE;1M 时达 23.7×——长上下文 decode 转向算力瓶颈, 是 v7x 算力密度有利的区间。这组数字至今没有任何真机对点,是下一阶段测量首当其冲的目标。
4. 生态盘点(2026-08-05 fetch + 代码核实)¶
4.1 各仓库现状¶
| 仓库 | 状态(2026-08-05 核实) | 对 K3 的意义 |
|---|---|---|
openxla/tokamax |
KDA 不在 main(origin/main dc54348 08-04 无 KDA);代码在 PR #1103 分支,本地钉在 29fa18c;PR head 已漂到 4d2135bd 且 API 重命名 | prefill 侧逐通道 KDA 的唯一公开来源;未合入 + API 漂移是活风险 |
google-pytorch/torchtpu-vllm |
本地与上游一致(88f359b);GDN v2 逐通道 decode kernel、MXFP4 MoE、循环状态管理均已核实存在 | PyTorch 侧服务栈,主线 |
vllm-project/vllm |
有 kimi_linear.py(48B 版,KimiMLAAttention 写死 use_nope/q_lora_rank is None 两条断言,213-214 行逐字属实);主干无任何 kimi_k3 支持;register_oot 是通用接管机制 |
模型定义来源;K3 与 48B 的差要补 |
google-pytorch/torch_tpu |
torch_tpu._internal.pallas.jax_op 已核实(pallas.py:781-806),JAX kernel → torch op 的桥 |
两条路线共用的缝合层 |
primatrix/maxtext + primatrix/pallas-kernel(tops) |
fork 主体是一套内部 KDA 混合架构模型的训练栈;KDA 走「maxtext 薄封装 → tops Pallas kernel」,prefill/varlen/CP 完备,KDA decode 至今是 NotImplementedError(attention_kda.py:370);其同型混合架构(KDA+MLA MoE)在 v7x 上有 64–2048 芯片级训练 CI 实测 |
JAX 路线备选 + 数值对拍基线;v7x 上 KDA kernel 大规模实战的旁证 |
google/tpu-raiden |
活跃开发中,JAX/torch 双绑定 | KV 传输 / P/D 分离(P4 优化项) |
google-pytorch/sglang-torchtpu |
已废弃(最新提交即 deprecate-repo) | 排除,不是候选路径 |
google/tpu-inference(k25-jax 经验) |
只支持 block-wise FP8;GKE v7x 上曾有 backbone 编译 crash(疑似镜像 AOT CPU feature 不匹配) | K3 的 MXFP4 与之正面冲突;crash 问题在 k3-jax 的 GKE JobSet 路径上未复现,判断为镜像特定问题而非平台问题 |
4.2 已核实的关键代码事实(支撑「没有要从零写的 kernel」)¶
- decode:
torchtpu-vllm/kernels/gdn/v2/gdn_decode_kernel.py是完整的融合逐通道单步 kernel—— state_indices 索引、L2norm、门控(A_log/dt_bias/lower_bound)全部内联,接口与 K3 的 KDA 逐条对应。 - prefill:tokamax
kimi_delta_attention(pr1103 分支)支持逐通道门、varlen、CP,自带 XLA 参考实现 (白捡的对拍基准);K3 配置逐条落在其支持范围内。差异仅四处:布局[H,B,T,K]↔[B,T,H,K]、 beta sigmoid 外提、CP 模式禁用 initial_state、T%64≠0时静默退回 XLA(decode 单步掉出 kernel, 慢 ~19×——但 vLLM 线的 gdn/v2 不受此限)。 - MoE 无 a2a:torchtpu-vllm 全 src 仅 4 处
all_to_all,均与 MoE 无关;EP 是「本地专家 + experts_start 编译期偏移」,combine 走 vLLM 的 TP all-reduce。 - MXFP4:TPU 上没有 mxfp4 linear kernel,上游既定行为是 LinearBase 一律退回 unquantized
(mxfp4.py:115-128 属实);MoE 走
VllmMxfp4MoEMethod现解现算。 - MLA kernel 是死代码:
kernels/mla/v1/kernel.py(1327 行 ragged paged MLA)存在但没有任何代码 import 它——k3serve 的 MLA 是自己接的线(k3serve/models/mla.py,一串 #35-#173 修复)。初版 「MLA ragged paged attention 现成」的说法要打折:kernel 在,接线不在。
4.3 maxtext / tops(JAX 路线备选,已核实)¶
primatrix/maxtext 的主体是一套内部 KDA 混合架构模型的训练栈(非 Google 官方上游),KDA 能力是「maxtext 薄封装 →
tops Pallas kernel」的结构(tops pin 在 [email protected])。核查结论:
- 训练侧成熟度高:KDA chunked 前向 + varlen + CP(a2a 与 all-gather 两种)完备,测试齐全; 其内部与 K3 同型的混合架构(KDA + MLA 混合 MoE、数百专家)在 v7x 上有 64 到 2048 芯片级的训练 CI 实测记录(含 kernel block size sweep)。 这反过来是 v7x + KDA kernel 组合被大规模实战过的旁证。
- 推理侧对 K3 是空白:KDA autoregressive 直接
NotImplementedError(attention_kda.py:370, 无 recurrent state cache、未接入 maxengine);没有 K3 的 config / 模型定义 / HF→maxtext 权重转换; 没有 MXFP4 的 TPU 推理路径(只有一个 GPU/torch 的反量化脚本);MLA 是 latent 存储但非 absorbed 计算(decode 每步上投影全部历史 latent),且 MLA cache 不支持量化。 - 结论:JAX 推理路线要补的四块(decode kernel + cache、模型定义、MXFP4、absorbed MLA)一块都没少, 初版对 JAX 路线 56–87 天的估算方向正确、甚至偏乐观。maxtext 的定位应保持在「KDA 训练侧参照 + 数值对拍基线(tops 的 CPU naive 实现)」,不建议作为 K3 推理路线重启。
5. 八天真机进展:两条线的真实状态(2026-08-05 核实)¶
初版发出后,项目实际走出了双线并行的格局,这一点初版没有预料(它设想的是「选定主线、单线推进」):
k3/ 自建 JAX 栈(flax.nnx) |
k3serve/ vLLM 服务线(torch_tpu) |
|
|---|---|---|
| 定位 | 数值对拍基准、语义真相 | 生产服务路径 |
| 已完成 | 逐层对拍 HF(atol 2e-5);整模型贪心生成与 HF 逐 token 全等(MINIMAL);多 host SPMD(v6e-16);真权重在 v7x-16 加载(10.6 分钟)并生成 4 步、logits 统计健康、4 pod 逐字节一致 | 48B→K3 六个结构差异全部补齐(SiTU、MLA 输出门、满秩门、LatentMoE、跨层残差、q_lora_rank);MLA 接线 + eager grouped_topk + mxfp4 Linear 规则;tiny 真权重 smoke 35/35;TP=32 + dummy 权重在 v7x-16 端到端出 token(08-05,35/35) |
| 未完成 | 不是服务栈(无连续批处理/paged KV);decode 循环形态决定了它只是参照物 | 真权重端到端:三次尝试中——第一次被 7200s 超时掐断(42/96 分片),第二次栽在 /dev/shm 孤儿段导致 gloo barrier 崩,第三次进行中;TP 数值判据 #94 仍 OPEN(随机权重 90.5% bf16 平局,判不了切分对错) |
| 显眼的坑 | 93 层纯展开图 XLA 调度塌掉 → scan 折图;多 host 三套编号互不相同的静默错位 | Ray 按 vfio 数到 8 芯片 vs webhook 注入 4 → 只起 16 worker 无声卡死;driver 必须落在内网 IP 字典序最小 pod |
一个值得单独记录的反差:JAX 线的权重加载被优化到了 10.6 分钟(1561 GB),而 vLLM 线加载 96 个分片 两小时还没过半。JAX 线在 08-03~04 做过一轮 GCS 加载性能攻坚(#267–#296,吞吐从 51–114 MiB/s 修到 接近基线),vLLM 线还没有做对应的工作——这是服务化清单上一个具体、可估的条目。
6. 风险与未知量(按 2026-08-05 现状重排)¶
| # | 项 | 状态 | 评注 |
|---|---|---|---|
| 1 | 端到端性能数字为零 | 未解 | 吞吐/延迟/带宽效率全部没有真机对点;§3.3 的 roofline 是纯分析。这是当前信息价值最高的测量 |
| 2 | vLLM 线真权重端到端 | 在途(第三次尝试) | 出了 token 之后,「可行」才算在服务线上闭环 |
| 3 | decode 路径 collective 清点 | 在途(#310 / PR #312) | 主线用 all-reduce 替代 a2a 后,一步的总 collective 开销无人知道;a2a 实测 9.88 µs 说明小消息固定开销这个量级不能假设 |
| 4 | tokamax#1103 未合入 + API 漂移 | 活风险 | 08-05 当天已造成一次真实事故;缓解 = 钉 commit(已做)+ 跟进合入 |
| 5 | 16 设备 decode 撞 XLA RET_CHECK | OPEN(#313) | 新出现,性质未定 |
| 6 | TP 数值判据缺失 | OPEN(#94) | 随机权重上判不了;真权重上去之后应第一时间补 |
| 7 | vLLM 线权重加载 ~4 小时级 | 未立项 | 对照 JAX 线 10.6 分钟;生产可接受度的问题,不是可行性问题 |
| 8 | v7x 供给 | 外部约束 | Cloud TPU API 查不到,只以 GKE 机型存在;上次交付排队 19h51m、租期 7 天。容量规划必须把这个提前期算进去 |
| 9 | KDA 状态 fp32 434 MB/序列(13.5 MiB/device 恒定) | 已知 | 短上下文并发上限由它决定;降 bf16 的数值稳定性未评估 |
| 10 | tpu-inference 的 FP8-only 与 GKE 编译 crash | 降级 | 那是另一条栈(k25 线)的问题;k3-jax 的 GKE JobSet 路径未复现 crash,判断为镜像特定 |
7. 独立判断:同意初版什么、不同意什么¶
同意的,且证据比初版写作时更强:
- 路线选型(torchtpu-vllm 为主线)。 后续八天完全站在这一侧:GDN decode kernel、mxfp4 MoE、状态管理 全部核实存在;JAX 路线(maxtext/tops)的 KDA decode 缺口(NotImplementedError)至今没人填。 实际项目自发演化成「JAX 栈做数值基准 + vLLM 栈做服务」的双线,等于用行动确认了初版的判断。
- 「没有要从零写的 kernel」。 唯一的例外后来被初版自己找到了:tokamax 没有单步 decode kernel—— 但主线上 gdn/v2 有,所以这个例外不落在关键路径上。
- 长上下文锚定。 KDA 恒定状态 + 24 层 absorbed MLA 的账被实测反复确认(55296 B/token)。
不同意或需要修正的:
- 工期框架整体失真。 初版按「kernel 缺口」估的 27–47 天,实际八天下来 kernel 缺口确实不是瓶颈, 瓶颈全在它没立项的地方:编译可行性(scan)、多 host 数据摆放、平台差异、加载工程、以及 vLLM 线上 一打「不报错只是卡住/算错」的坑。我的重估:从 08-05 现状到「v7x-16/32 上真权重端到端 + 首批性能数字」,还需要约 2–4 周(真权重端到端收口 ~1 周、collective 清点与首轮调优 1–2 周、 加载工程化 1 周,部分可并行);到「可谈 SLO 的生产服务」还要再加 4–8 周。合计与初版下限接近, 但结构完全不同——风险不在 kernel,在系统集成与测量。
- 「16 chips 最小可行」要加定语。 作为「权重装得下 + 图编得出」是实测成立的;作为「服务配置」 它只剩 96–178K 上下文量级,兑现不了 1M context 的卖点。生产口径的最小可行是 v7x-32。
- 初版对「MLA ragged paged attention 现成」的表述过于乐观。 kernel 存在但没接线(死代码), 实际是 k3serve 自己把 MLA 接起来的(一串七个 issue 的修复)。这类「上游有 ≠ 上游接着」的情况 在 torchtpu-vllm 里不是个例,评估上游能力时要把「有文件」和「有调用」分开数。
- §6.4 的 a2a 恐慌有正确的形状、错误的落点。 实测 9.88 µs 证明它对延迟的担心是对的,但主线 根本不发 a2a。真正要盯的量(all-reduce 总开销)至今没数。在这一点上初版 §12.1⑩ 的自评是准确的, 我只是在优先级上把它提到第一梯队。
对初版可信度的总评: 本次抽查了它十几处代码引用(文件、行号、断言文本、kernel 接口),全部属实, 没有发现编造;它的实测章节有完整的自我纠错记录(同一节里保留错误归因与修正过程),这种写法让 「实测」标签的含金量明显高于一般 AI 生成物。它的实测数字我倾向于采信;它的分析数字(roofline、 MFU 假设、带宽效率 50%)和一切工期估算,应按未验证处理。
8. 结论与建议¶
可行,且比初版写作时更接近「已证明」。 截至 2026-08-05:真权重已在 v7x-16 完成加载与 93 层前向 (JAX 线),vLLM 服务线 TP=32 管道已在同一片机器上出 token(dummy 权重),真权重端到端在途。 不建议此时再开第三条技术路线(maxtext/tpu-inference/sglang 均无必要)。
建议的下一步排序(按信息价值,不按方便程度):
- 收口 vLLM 线真权重端到端(在途)→ 出了 token,「可行」在服务线上闭环。
- 拿第一批端到端性能数字:单步 decode 延迟、prefill 吞吐、对 §3.3 roofline 的对点。哪怕是 batch=1 单配置,也把「分析值」清单里最大的一项划掉。
- decode collective 清点(PR #312 合入)→ 决定小 batch 场景的 EP/TP 配比。
- 把 vLLM 线的权重加载做到与 JAX 线同一量级(10 分钟级),否则每次实验迭代都被加载时间税一遍。
- 跟进 tokamax#1103 合入;在合入前坚持钉 commit。
- 容量规划按 v7x-32 做生产基线,v7x-16 定位为开发/验证配置;把 GKE 排队提前期(实测 ~20 小时) 算进任何时间表。
本报告基于 2026-08-05 的代码与真机记录。模型事实已对回 HF 官方 config.json;上游代码声明已逐条核实; 实测数字引自 k3-jax 仓库的 issue/日志/提交记录。所有 roofline 性能数字仍为分析值,等待真机对点。