Prhub

#50487 [Model][Spec Decode] Tap the pre-norm AttnRes mixture as the Kimi K3 DFlash aux state

原始 PR 作者 rchalamala 合并时间 2026-08-15 00:03 文件变更 4 提交数 4 评论 29 代码增减 +350 / -1

执行摘要

Kimi K3 DFlash 改捕获 pre-norm AttnRes 混合流(默认关)

PR body 明确指出:“K3 captures the post-mixture stream, which is not what the drafter was trained against: the AttnRes residual mixture is applied before the layer norm, and the current capture site reads the value after it.” 即 AttnRes 混合发生在 layer norm 之前,而捕获点读取的是它之后的值,导致 DFlash 拿到与训练分布不一致的辅助张量。改为 tap pre-norm 混合可恢复 drafter 期望的流,维护者 zixi-qi 也评价“The change LGTM, this is probably more correct than our current approach”。因改变 speculator 看到的数值,PR 用默认关闭的环境开关控制,等待在更真实的缓存场景下复测。

值得精读。重点关注三点:捕获点与训练目标对齐的必要性;把 pending MLP 输出折叠进 prefix 而非作为 delta 传递,以规避内核就地写回导致的二次加和;以及 PP 边界上“值是否被消费”的论证方式。合并后建议在 prefix caching 开启的真实服务上复测接受率,并跟踪 revert 提交中提到的框架级配置校验是否落地。

讨论亮点

yubofredwang 对第三分支提出实质性质疑:“The first two branches look right. I think the third one — the fallback to the running prefix for the last layer of a non-final stage is wrong.” 他以 pp=2、16 层、attn_res_block_size=4 的构造证明:返回 prefix 等价于把 softmax 强制 one-hot 到最后源并丢弃 bank[0] 与 bank[1] 的学得权重贡献,合成权重下与真实混合的余弦相似度仅 0.05-0.60。作者先用 commit 2c46ba2 混合边界权重,随后用 commit f2c897f revert,论证 aux hidden states 不跨 PP 边界、非最后 stage 的 tap 值无人消费,最终保留回退并把论证写进 docstring。zixi-qi 评论“The change LGTM, this is probably more correct than our current approach”;ZJY0516 最后 APPROVED。njhill 在合并前确认是否 ready,作者确认并处理 env var 冲突后完成合并。

实现拆解

  1. 环境开关与配置期日志:在 vllm/envs.py 注册 VLLM_KIMI_K3_AUX_ATTN_RES_STREAM(默认 False);在 vllm/models/kimi_k3/nvidia/model.py 覆写 _set_aux_hidden_state_layers,用 logger.info_once 在 setup-time 一次性输出 tap 层元组与 capture 模式(attn_res_stream / prefix_only)。这避免在 torch.compile 的 forward 内给 nn.Module 设置属性引发 graph break 或重编译,也是日志的自然位置。
  2. 核心捕获逻辑 _capture_aux_hidden_stream:把 pending MLP 输出折叠进 prefix(prefix_sum + hidden_states),而不是作为 delta 传给内核——因为内核会把已应用的 delta 就地写回 prefix,造成双重加和;开关关闭或未启用 AttnRes 时原样返回该 prefix,保证与旧行为一致;否则按三分支选择权重:tap 层后还有层时用下一层的 self_attention_res_norm / self_attention_res_proj 与 prev_valid_blocks;tap 在最后层且最后 rank 时用模型的 output_attn_res_norm / output_attn_res_proj 与 num_attn_res_blocks;最后层而非最后 stage 时回退 prefix。每次混合都通过 attn_res 内核以无 delta、block_write_idx=-1、无 output norm 的方式只读计算 bank[:num_blocks] + prefix 的软混合,不触碰在线残差流。
  3. forward 接线:在 forward 的 use_attn_res 分支中,把原来的 aux_hidden_state = prefix_sum + hidden_states 替换为对 _capture_aux_hidden_stream(layer_idx, prefix_sum, hidden_states, residual) 的调用,并补上 residual 非空断言。
  4. 测试配套:新增 tests/models/kimi_k3/test_aux_attn_res_stream.py,用 SimpleNamespace stub 模型与 monkeypatch 的假 attn_res 内核覆盖三分支选择、关闭路径的精确等价性、以及 pending MLP 折叠的不变式(调用方 tensor 不被改写);test_eagle3.py 修复共享 stub 缺失 use_attn_res 属性(新 override 会无条件读取),并新增 forward 参数顺序测试,防止层输出与残差在调用点被静默交换。
  5. 一次失败的修正尝试:commit 2c46ba2 曾尝试让非最后 stage 的最后层也混合边界权重,但随即被 revert——aux hidden states 不会跨 pipeline-parallel 边界传递,每个 stage 只返回普通 IntermediateTensors,runner 只会在最后 rank 解包辅助输出,因此非最后 stage 的 tap 值任何情况下都会被丢弃,修正边界权重等于喂给无人消费的值。最终保留 prefix 回退并在 docstring 中写明依据。
文件 模块 状态 重要度
vllm/models/kimi_k3/nvidia/model.py 模型前向 modified 7.97
tests/models/kimi_k3/test_aux_attn_res_stream.py 单元测试 added 7.93
tests/models/kimi_k3/test_eagle3.py 投机解码 modified 5.48
vllm/envs.py 环境变量 modified 5.27

关键符号

_capture_aux_hidden_stream _set_aux_hidden_state_layers _aux_attn_res_stream KimiLinearModel.forward

关键源码片段

vllm/models/kimi_k3/nvidia/model.py data-contract

核心实现文件:新增捕获助手决定 DFlash 看到的具体张量语义,并覆写配置期钩子做一次性日志。

# 核心捕获逻辑:DFlash 看到的辅助 hidden state 必须是消费者实际读到的
# pre-norm AttnRes 混合值,而不是 post-mixture 之和。
# pending_mlp_out 不能作为 delta 传给内核:内核会把已应用的 delta 就地写回
# prefix,导致在线残差流被二次相加,因此必须先折叠进 prefix。
def _capture_aux_hidden_stream(
    self,
    layer_idx: int, # 刚执行完的层索引(0-based)
    prefix_sum: torch.Tensor, # 当前 block 的运行前缀
    pending_mlp_out: torch.Tensor | None, # 刚算完的 MLP 输出(hidden_states)
    block_residual: torch.Tensor, # 已提交 block 的残差 bank
) -> torch.Tensor:
    prefix = prefix_sum if pending_mlp_out is None else prefix_sum + pending_mlp_out
    # 开关默认关闭;use_attn_res 为 False 时本就没有 norm/proj 权重可读,
    # 原样返回就是旧行为(prefix_sum + hidden_states)。
    if not (self._aux_attn_res_stream and self.use_attn_res):
        return prefix
​
    # 分支 1:后面还有层,用下一层自己的 AttnRes 权重与已提交 block 数。
    if layer_idx + 1 < self.end_layer:
        consumer = self.layers[layer_idx + 1]
        score_norm = consumer.self_attention_res_norm
        score_proj = consumer.self_attention_res_proj
        num_blocks = consumer.prev_valid_blocks
    elif get_pp_group().is_last_rank:
        # 分支 2:最后层的最后 rank,无下游层,用模型输出侧聚合权重。
        score_norm = self.output_attn_res_norm
        score_proj = self.output_attn_res_proj
        num_blocks = self.num_attn_res_blocks
    else:
        # 分支 3:非最后 stage 的最后层,本 rank 无输出侧权重;该 tap 值
        # 不会跨 PP 边界传输,回退 prefix 只保证计算有定义。
        return prefix
​
    # 无 delta、不写 block、无 output norm 的 attn_res 调用正好计算
    # bank[:num_blocks] 与 prefix 的软混合,且不改变在线残差流。
    return attn_res(
        prefix,
        None,
        block_residual,
        score_norm.weight,
        score_proj.weight.squeeze(0),
        None,
        num_blocks=num_blocks,
        block_write_idx=-1,
        eps=score_norm.variance_epsilon,
        output_norm_eps=0.0,
    )
tests/models/kimi_k3/test_aux_attn_res_stream.py test-coverage

新增测试精确覆盖三分支选择与关闭路径,用 stub 模型和假内核校验权重来源与不变量。

# 关闭路径的等价性测试:开关关闭或不启用 AttnRes 时,tap 必须严格等于
# 它替换掉的 prefix_sum + hidden_states,且内核不得被调用。
@pytest.mark.parametrize(
    'enabled,use_attn_res', [(False, True), (True, False), (False, False)]
)
def test_disabled_reproduces_the_plain_residual_sum(
    recorder, monkeypatch, enabled, use_attn_res
):
    _set_last_rank(monkeypatch, True)
    prefix_sum = torch.tensor([1.0, 2.0])
    pending = torch.tensor([0.5, 0.25])
​
    got = _call(
        _stub_model(enabled=enabled, use_attn_res=use_attn_res),
        0,
        prefix_sum,
        pending,
        torch.zeros(2),
    )
​
    torch.testing.assert_close(got, prefix_sum + pending)
    assert not recorder, 'the kernel must not run when the tap is off'

评论区精华

PP 非最后 stage 的 prefix 回退在语义上错误 正确性

yubofredwang 指出返回 prefix 等价于把 softmax 强制 one-hot 到最后源,丢弃 bank[0] 与 bank[1] 的学得权重贡献;合成权重下与真实混合的余弦相似度仅 0.05-0.60。

结论:作者先以 commit 2c46ba2 尝试混合边界权重,随后因 aux hidden state 不会跨 PP 边界传递、非最后 stage 的 tap 值无人消费而 revert,最终保留 prefix 回退并在 docstring 中注明依据。 · closed-with-argument

维护者对方案正确性的认可 设计

zixi-qi 评论 “The change LGTM, this is probably more correct than our current approach”;ZJY0516 随后 APPROVED。

结论:维护者认可 pre-norm 混合捕获方向并批准合并。 · 已解决

合并前确认与环境变量冲突 question

njhill 询问 “is this ready to merge now from your pov?”;作者确认后修复一个 env var 冲突并再次请求合并。

结论:合并者完成合并。 · 已解决

风险与影响

PP 边界语义争议未完全消解:yubofredwang 指出的数值不等价在最终代码中依旧存在,作者以“aux 不跨 PP 边界、该值无消费者”消解其影响;一旦未来 runner 允许跨 stage 传递 aux hidden states,或捕获 ID 落在 stage 边界且该 capture 被消费,分支 3 会静默喂给 drafter 错误张量。当前测试为单 GPU,未覆盖真实多 GPU PP 组合。数值分布风险:开启后 speculator 输入从 post-mixture 之和变为 pre-norm 混合,e2e 在 prefix caching 关闭下测得 +1.9,缓存开启时仅 +0.083,落在臂内波动内,需在有缓存的生产路径复测。内核语义依赖:正确性依赖 attn_res 在 block_write_idx=-1、delta=None 时不写回 prefix 的约定,测试用 fake 内核,未在真实 CUDA 内核上验证该只读语义。

默认关闭且仅影响 Kimi K3 NVIDIA 模型 + AttnRes + DFlash 的专用捕获路径,不改变默认模型行为;开启后投机解码接受长度显著提升(e2e 均值 +1.9,p=1.5e-8),有利于 Kimi K3 的 DFlash 部署。新增测试集中在 tests/models/kimi_k3 目录,对单 GPU 环境无风险。对团队的直接影响是明确了 capture 点必须与 drafter 训练目标对齐的原则,并为后续缓存场景复测与框架级配置校验留下任务。

默认关闭的实验性开关 PP 边界语义争议未完全解决 改变 speculator 输入分布 缓存命中场景未复测

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论