跳转至

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(实测)。剩下的是「跑多快、值不值」,这两个问题 目前一个数字都没有。

六条主要判断:

  1. 显存可行域已实测钉死。 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。这是初版静态分析完全没预见、 真机才暴露的一类问题。
  2. KDA kernel 风险已解除(初版最担心的一环)。 tokamax 的逐通道 KDA Pallas kernel 不改任何参数 即在 TPU7x 上编译通过、数值正确(4 host 一致)。prefill 侧该 kernel 比 XLA 参考快 31.6×; decode 侧 torchtpu-vllm 有现成的逐通道单步融合 kernel(kernels/gdn/v2,已核实存在且接口匹配)。
  3. 长上下文是 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 起步。
  4. 低并发 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,在途)。
  5. 初版的「27–47 天」工期估算已被日历时间证伪,但证伪方式有信息量。 07-28 至今的 8 天没有花在预估的 G2「接线」上,而是花在初版未立项的三类事:整模型编译可行性(scan 折图)、多 host 数据摆放 (process_id ≠ process_index ≠ mesh 顺序)、以及一批「v6e 绿、v7x 才现形」的平台差异(默认 matmul 精度硬件相关等)。剩余工作量应按「vLLM 服务线真权重端到端 + 性能测量 + 加载工程化」重估。
  6. 上游依赖的最大现实风险是 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.pyK3_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」)

  • decodetorchtpu-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 直接 NotImplementedErrorattention_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)。

不同意或需要修正的:

  1. 工期框架整体失真。 初版按「kernel 缺口」估的 27–47 天,实际八天下来 kernel 缺口确实不是瓶颈, 瓶颈全在它没立项的地方:编译可行性(scan)、多 host 数据摆放、平台差异、加载工程、以及 vLLM 线上 一打「不报错只是卡住/算错」的坑。我的重估:从 08-05 现状到「v7x-16/32 上真权重端到端 + 首批性能数字」,还需要约 2–4 周(真权重端到端收口 ~1 周、collective 清点与首轮调优 1–2 周、 加载工程化 1 周,部分可并行);到「可谈 SLO 的生产服务」还要再加 4–8 周。合计与初版下限接近, 但结构完全不同——风险不在 kernel,在系统集成与测量。
  2. 「16 chips 最小可行」要加定语。 作为「权重装得下 + 图编得出」是实测成立的;作为「服务配置」 它只剩 96–178K 上下文量级,兑现不了 1M context 的卖点。生产口径的最小可行是 v7x-32。
  3. 初版对「MLA ragged paged attention 现成」的表述过于乐观。 kernel 存在但没接线(死代码), 实际是 k3serve 自己把 MLA 接起来的(一串七个 issue 的修复)。这类「上游有 ≠ 上游接着」的情况 在 torchtpu-vllm 里不是个例,评估上游能力时要把「有文件」和「有调用」分开数。
  4. §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 均无必要)。

建议的下一步排序(按信息价值,不按方便程度):

  1. 收口 vLLM 线真权重端到端(在途)→ 出了 token,「可行」在服务线上闭环。
  2. 拿第一批端到端性能数字:单步 decode 延迟、prefill 吞吐、对 §3.3 roofline 的对点。哪怕是 batch=1 单配置,也把「分析值」清单里最大的一项划掉。
  3. decode collective 清点(PR #312 合入)→ 决定小 batch 场景的 EP/TP 配比。
  4. 把 vLLM 线的权重加载做到与 JAX 线同一量级(10 分钟级),否则每次实验迭代都被加载时间税一遍。
  5. 跟进 tokamax#1103 合入;在合入前坚持钉 commit。
  6. 容量规划按 v7x-32 做生产基线,v7x-16 定位为开发/验证配置;把 GKE 排队提前期(实测 ~20 小时) 算进任何时间表。

本报告基于 2026-08-05 的代码与真机记录。模型事实已对回 HF 官方 config.json;上游代码声明已逐条核实; 实测数字引自 k3-jax 仓库的 issue/日志/提交记录。所有 roofline 性能数字仍为分析值,等待真机对点。