Lesson 02 · Inference trace

一个请求走完 Kimi K3 的 prefill 与 decode

这次不再逐个孤立地看模块。我们让提示词 token [A,B,C,D] 一次穿过 93 层,采样得到 E;再把 E 作为下一轮输入,更新 KDA 与 MLA cache,得到用于采样 F 的 logits。

版本边界:shape 对照 vLLM main f4b161d7(2026-08-18)。为了看清时间线,示例只有一个请求、4 个 prompt token、非 speculative decode;A...F 是教学符号,不代表真实词表 id。真实模型仍是 93 层、hidden 7168、词表 160K。

0. 最先澄清:模型输出的是 logits,不是“下一 token”

prefill 输入: [A,B,C,D] 模型最后一个位置输出: logits_D [160000] 采样器: E ~ Softmax(logits_D / temperature) 第 1 次 decode 输入: [E] 模型输出: logits_E [160000] 采样器: F ~ Softmax(logits_E / temperature)

因此 prefill 已经计算出了第一个生成 token E 的分布;第一次 decode 消费的是 E,并产生 F 的分布。把“decode 输入 E”和“decode 输出 E”混为一谈,会让 cache 的写入时点整体错一格。

1. 先记住三类东西的生命周期

请求持久状态当前 forward 临时量模型常驻权重
类别prefill 后是否留下例子
请求持久状态是,decode 下一步读取并更新69 层的 KDA short-conv history 与 recurrent state;24 层的 MLA latent pages;长度、slot mapping、page table
当前 forward 临时量否,使用完成即可复用/释放Q/K/V、AttnRes depth buffer、router logits、expert activation、hidden states、LM logits
模型常驻权重一直在模型实例中,不属于某个请求投影矩阵、896 个 expert 权重、AttnRes score vectors、RMSNorm 参数、decode 吸收后的 W_UK_T/W_UV

2. Prefill:一次输入四个 token

2.1 embedding 与本轮 packed token 维

token_ids [4] embedding lookup → X0 [4,7168] attn_res buffer 初始化 [4,8,7168] positions / query_start_loc / slots metadata

vLLM 常把同一调度轮的多个请求 pack 到 token 维,所以生产代码里的第一维通常是本轮总 token 数 T,而不是整齐的 [batch,seq]。本例只有一个长度 4 的请求,所以 T=4

2.2 第 0 层是 KDA:四个位置并行准备,状态按顺序组合

AttnRes(X0) [4,7168] packed projection: Qraw/Kraw/Vraw/G² [4,96,128] F_a [4,128] beta_raw [4,96] ShortConv4 + Swish: q/k/v [4,96,128] decay α [4,96,128] FlashKDA / chunk KDA: S_0 --A--> S_A --B--> S_B --C--> S_C --D--> S_D output [4,96,128] flatten + output projection [4,12288] → [4,7168]

逻辑上仍按 A→B→C→D 更新每个 head 的 S:[128,128];prefill kernel 把 chunk 内可改写成 GEMM 的 token 交互并行计算,只在 chunk 边界传 state。该层完成后,请求保留最终 S_D:[96,128,128],而不是四份 state。

short conv 也只保留继续计算所需的最后 3 个 raw Q/K/V 输入。全局逻辑 shape 为 [36864,3];TP=8 时每 rank 是 [4608,3],recurrent state 是 [12,128,128]

2.3 经过第三个 KDA 后,进入第一层 Gated MLA

X_mla [4,7168] fused_qkv_a_proj → [q_c | c] [4,2112] q_c / latent c after RMSNorm [4,1536] / [4,576] q_b_proj → Q [4,96,192] kv_b_proj(c) → K / V(prefill 临时展开) [4,96,192] / [4,96,128] causal attention scores [96,4,4] attention output + full-rank gate [4,96,128] o_proj [4,7168] 写入 paged latent cache 4 × [576]

D 个位置只能看 A...D;第 A 个位置只能看自己。prefill 为了高吞吐会临时展开完整 K/V,但请求持久 cache 只留下四个 576 维 latent:[c_A,c_B,c_C,c_D]。相同动作在 24 个 MLA 层分别发生。

2.4 每层的 AttnRes 与 Stable LatentMoE

attention 前 AttnRes [4,R,7168] → [4,7168] KDA 或 MLA [4,7168] MLP 前 AttnRes [4,R,7168] → [4,7168] router logits [4,896] FP32 Top-16 ids / weights [4,16] / [4,16] routed down [4,3584] dispatch 后逻辑输入 [4×16,3584] expert gate/up [4×16,3072] 各一份 combine + norm + up [4,3584] → [4,7168] shared experts [4,7168] 合并输出 [4,7168]

R 是该层可见的深度来源数,随深度 block 推进而增加,最多是 embedding、已完成的 block 和当前 block 部分和。它不是序列长度 4。四个 token 的 AttnRes buffer 在这次 93 层 forward 内流转,最后即失去用途;decode E 时会为 E 新建一份。

2.5 第 92 层、LM head 与第一次采样

最后一层 Gated MLA 输出 [4,7168] 最终 AttnRes + RMSNorm [4,7168] LM head(逻辑上) [4,160000] 只取本请求最后位置 D 的 logits [160000] 采样 → E scalar token id

服务引擎通常无需保留 prompt 每个位置的全词表 logits;生成只需要最后一个有效位置。至此 prefill 结束,cache 已经包含 A...D,但还不包含刚采样出的 E

3. Prefill 结束后的请求状态快照

持久项每层/每请求逻辑内容共几组此刻代表的历史
KDA recurrent state[96,128,128];TP=8 每 rank [12,128,128]69 层A...D 经各层输入形成的最终 state
KDA conv stateQ/K/V 各通道最近 3 项;全局 [36864,3]69 层能继续对 E 做 causal Conv4
MLA latent pages每 token [576]24 层每层各有 [c_A,c_B,c_C,c_D]
调度/cache metadatarequest length=4、block table、slot mapping 等请求级下一步把 E 写到逻辑位置 4
不要算错显存:上表是逻辑 shape。实际字节数还依赖 TP/PP、state dtype、latent cache dtype、对齐、page 分配和是否启用 RecoverSSM。没有这些假设时,只能比较元素量和增长阶数。

4. Decode:把 E 作为一个新输入 token

4.1 新 forward 从头穿过 93 层

decode input token_id = E [1] embedding → X_E [1,7168] 新的 AttnRes buffer [1,8,7168] 注意:E 仍会依次经过第 0...92 层; “只算一个 token”不等于“只算一层”。

4.2 在每个 KDA 层原地前进一步

读取该层旧 conv history + S_D E 的 raw Q/K/V → causal Conv4,丢掉最老项并追加 E 生成 q_E,k_E,v_E,α_E,β_E,g²_E S_before = Diag(α_E) S_D S_E = S_before + β_E k_E (v_E - S_before^T k_E)^T o_E = S_E^T q_E [96,128] gate + o_proj [1,7168] 写回该层 conv history 与 S_E

纯 decode 正是 vLLM kda.py 的 fused recurrent 路径:按 metadata 找到该请求的 state slot,更新 conv state 与 recurrent state。计算量不会随 A...D 的长度线性增长,但必须为 69 个 KDA 层分别维护状态。

4.3 在每个 MLA 层追加 E,并查询 A...E

E → q_E [1,96,192], c_E [1,576] 把 c_E 写入本层 cache,历史变成 [c_A,...,c_E] BMM1: q_latent = q_E W_UK_T [1,96,576] latent MQA scores [1,96,5] latent_out = Softmax(scores) · [c_A...c_E] [1,96,576] BMM2: head_out = latent_out W_UV [1,96,128] full-rank gate + o_proj [1,7168]

decode 路径没有为五个历史位置物化 K:[5,96,192]V:[5,96,128]。vLLM 在模型加载后的 weight processing 阶段准备 W_UK_T/W_UV,运行时走 BMM1 → latent MQA → BMM2。这是 MLA 真正降低 decode cache 读取与中间展开的地方。

4.4 MoE 与 AttnRes 没有历史 cache

E,router 重新生成 [1,896] logits,并可能选择与 D 完全不同的 16 个专家;expert activation 用完就释放。AttnRes 也只混合 E 在深度方向的表示,不会读取 A...D 的 AttnRes buffer。跨 token 信息来自 KDA state 与 MLA latent history。

4.5 得到 F 的分布

E 经过最后层 → hidden_E [1,7168] LM head → logits_E [1,160000] F ~ Softmax(logits_E / temperature) 此刻 cache 已包含 A...E;F 尚未写入。 下一次 decode 消费 F,产生 G 的 logits。

5. vLLM 为什么需要一份混合 metadata

一次 scheduler batch 可以同时放入长 prompt、chunked prefill 的续块、普通 decode token,甚至 speculative token。token 张量会被 pack 在一起,但 KDA kernel 要区分“可 chunk 并行的 prefill”和“原地前进一步的 decode”,MLA backend 又要知道每个 query 对应哪些 latent pages。

metadata 作用KDA 使用方式MLA 使用方式
请求边界与 query 长度query_start_loc 切分各请求递推序列构造 causal mask 与每请求 context 长度
state/cache 索引state_indices 定位 conv/recurrent slotslot_mapping/block_table 定位 latent page
prefill/decode 分类选择 FlashKDA/chunk 或 fused recurrent选择 full prefill attention 或 absorbed latent MQA
重新排序spec 与 non-spec token 可先拆开再写回 packed 顺序decode query 与 cache insert 在 fused kernel 中对齐

所以 Python forward 看起来是一条统一数据流,底层却按 token 类型分派不同 kernel;这属于 vLLM 的批处理实现,不是模型公式额外增加了一种 attention。

6. Speculative decode 为什么对 KDA 更难

普通 MLA 可以给 draft token 分配临时 cache slots,拒绝后丢弃。KDA 的问题是 state 更新具有顺序性:若一次验证 [E,F,G],执行完后 S 已前进三步;若只接受 E,不能通过删除两条 KV 记录恢复 S_E

验证前 checkpoint: S_D 草稿验证: S_D --E--> S_E --F--> S_F --G--> S_G 只接受 E 时,正确提交状态应是 S_E,而不是 S_G

vLLM 当前 K3 支持把 speculative token 与 non-spec token 分开。ReplaySSM/RecoverSSM 路径保留 checkpoint 与更小的校正/投影记录,在验证结果已知后重放或恢复正确边界,而不是为每个 draft 位置复制完整 [96,128,128] state。这是后端优化;不开 speculative decode 时,本课前面的 plain decode 时间线不需要这些附加记录。

7. 用一张表复盘整轮生成

时刻本轮模型输入cache 提交后包含模型给出的分布随后采样
prefillA,B,C,DA...Dp(next | A...D)E
decode #1EA...Ep(next | A...E)F
decode #2FA...Fp(next | A...F)G

检索练习

问题:prefill 刚采样出 E、但还没执行第一次 decode 时,哪些请求状态已经包含 E?

先读:第一课:四个模块的公式与 shape。实现对照:vLLM Kimi K3 NVIDIA backend at f4b161d7。速查:推理状态台账