作者​:昇腾实战派
知识地图​:https://blog.csdn.net/Lumos_Lovegood/article/details/161601003

背景概述

在多Token预测(MTP)投机推理场景中,模型通过同时预测多个后续token来加速推理过程。然而,部分场景下在vLLM-Ascend框架中接入MTP后,Full Graph图模式下可能会出现精度异常问题,表现为draft token接受率低于Piecewise图模式。

本文详细记录了该问题的定位过程、根因分析及修复方案,为类似场景下的精度问题排查提供参考。

问题现象

在vLLM-Ascend框架中,step3.5-flash模型接入MTP(Multi-Token Prediction)投机推理后,Full Graph图模式下出现精度异常,draft token接受率显著低于Piecewise图模式。

开MTP时断图精度正常,Full Decode Only模式下精度异常。
mtp3 + full-decode-only 模式下,mtp接受率异常,相较于mtp3 + piecewise模式下, 明显偏低。

定位思路

  1. 排除 MTP 投机推理本身:MTP 不影响主模型精度,因此问题不在 MTP 投机推理的逻辑中。
  2. 锁定图模式更新逻辑:Piecewise 图模式精度正常,Full Decode Only 图模式精度异常。两者的区别在于 Full Decode Only 模式在 Decode 阶段需要更新计算图(update-graph),怀疑该路径存在参数错误。
  3. 对比入图/不入图参数差异:最终定位到 FIA 算子的入参在入图和不入图时设置不一致,具体是 pre_tokensnext_tokens 的值存在差异:
    • 不入图时(eager 及非入图分支):pre_tokens = sliding_windownext_tokens = 0
    • 入图时(图参数更新分支):pre_tokens = sliding_window - 1next_tokens = 1

原因分析

对于FIA算子中 pre_tokensnext_tokens 两个参数的含义,可以理解为:

参数含义
pre_tokens当前 token 向前(past)能 attend 的 token 数量
next_tokens当前 token 向后(future)能 attend 的 token 数量
  • MTP1(单头):seq_len=1,当前 token 之后无其他 token,next_tokens=1 等效于 next_tokens=0
  • MTP3(三个mtp头,step 模型):seq_len=4,若 next_tokens=1,当前 token 可看到后续 1 个 token,破坏因果性,影响精度。

为什么 MTP1 无影响而 MTP3 有影响?

next_tokens 控制当前 token 向后能 attend 的 token 数量。在 Decode 阶段:

  • MTP1 场景:Draft token 数为 1(seq_len=1),当前 token 之后没有有效 token 可以 attend,因此 next_tokens=1next_tokens=0 效果等价,参数错误被掩盖。
  • MTP3 场景:Draft token 数为 3(总 seq_len=4),若 next_tokens=1,每个 token 可以 attend 到其后 1 个 token,这违反了因果 attention 的约束(token 不应看到未来的信息),导致 attention 分布偏移,精度下降。

为什么 eager 模式精度正常?

Eager 模式下,开 MTP 走的是非入图的 forward_fia 分支,该分支参数正确设置为 pre_tokens = self.sliding_window, next_tokens = 0

而在 Full Graph 模式下,主模型前 45 层走的是入图的 full_graph_fia 或图参数更新分支,这些分支的参数存在 pre_tokens = sw - 1, next_tokens = 1 的错误配置。

总结:问题的本质是同一个 FIA 算子在入图和不入图两条代码路径上参数设置不一致,入图路径的参数未经 MTP3 场景的正确性验证。

问题代码定位

Bug 位置

attention代码: forward_fia_slidingwindow 方法

问题代码(入图分支 / 图参数更新分支):

torch_npu.npu_fused_infer_attention_score.out(
    ...,
    pre_tokens=self.sliding_window - 1,  # 错误!应为 self.sliding_window
    next_tokens=1,                        # 错误!应为 0
    ...
)

对比 —— 正确代码(非入图分支):

torch_npu.npu_fused_infer_attention_score(
    ...,
    pre_tokens=self.sliding_window if self.sliding_window is not None else SWA_INT_MAX,
    next_tokens=0 if self.sliding_window is not None else SWA_INT_MAX,
    ...
)

bug代码位置截图

外链图片转存失败,源站可能有防盗链机制,建议将图片保存下来直接上传

def update_grpah_params() :

外链图片转存失败,源站可能有防盗链机制,建议将图片保存下来直接上传

修复方案

将入图分支及图参数更新分支中的 FIA 调用改为与非入图分支一致:

pre_tokens=self.sliding_window,   # 原值 self.sliding_window - 1
next_tokens=0,                     # 原值 1

修复涉及的代码入口

  • full_graph_fia() —— 入图时的 FIA 调用
  • update_graph_params() —— 图参数更新函数中的 FIA 调用

FIA算子分支

整体分支流程

将FIA算子入图和不入图、开mtp和不开mtp分为4个分支进行拆解,流程如下:

外链图片转存失败,源站可能有防盗链机制,建议将图片保存下来直接上传

FIA 算子分支矩阵

模式MTP阶段FIA 分支pre/next_tokens精度
EagerPrefillforward_fiasw, 0(正确)正常
EagerDecodeforward_fiasw, 0(正确)正常
EagerPrefillforward_fiasw, 0(正确)正常
EagerDecode (FA)forward_fiasw, 0(正确)正常
EagerDecode (SWA)forward_fia_slidingwindowsw-1, 1(BUG)正常(无MTP时seq_len=1)
Full GraphPrefillforward_fiasw, 0(正确)正常
Full GraphDecode (主模型45层)full_graph_fiasw-1, 1 (BUG)已修复,sw, 0**
Full GraphDecode (MTP 3层)forward_fiasw, 0(正确)正常
Full GraphPrefillforward_fiasw, 0(正确)正常
Full GraphDecodefull_graph_fiasw-1, 1 (BUG)对seq_len=1无影响,不影响精度,但应修复为**sw, 0**

注:sw = self.sliding_window

FIA 代码分支代码位置

eager 模式下, 开mtp:

Prefill阶段:

模型48层

进入的是 forward_fia 的 如下分支

外链图片转存失败,源站可能有防盗链机制,建议将图片保存下来直接上传

Decode阶段

attn_state = SpecDecoding

seq_lens=1,query.size(0) = 4

模型48层(1fa + 3swa)

走的也是 forward_fia 的 npu_fused_infer_attention_score 分支(上图所示)

eager模式下, 不开mtp:

Prefill阶段:

主模型45层(1 fa + 3 swa) 都进入:

外链图片转存失败,源站可能有防盗链机制,建议将图片保存下来直接上传

Decode阶段:

full attention层时:

sw = None

外链图片转存失败,源站可能有防盗链机制,建议将图片保存下来直接上传

sw attention层时:

sw =512

seq_lens == query.size(0) = 1

atten_state = DecodeOnly

外链图片转存失败,源站可能有防盗链机制,建议将图片保存下来直接上传

外链图片转存失败,源站可能有防盗链机制,建议将图片保存下来直接上传

FULL图模式下, 开mtp:

Prefill阶段:

48层(fa + swa)

外链图片转存失败,源站可能有防盗链机制,建议将图片保存下来直接上传

Decode阶段:

模型共48层(主模型45 + mtp3)

主模型前45层, 进入 full_graph_fia

外链图片转存失败,源站可能有防盗链机制,建议将图片保存下来直接上传

后3层 mtp层 (未入图):

seq_lens =1, query.size(0) = 4

atten_state = SpecDecoding

外链图片转存失败,源站可能有防盗链机制,建议将图片保存下来直接上传

FULL图模式下,不开mtp

Prefill阶段:

模型45层(fa + swa):

外链图片转存失败,源站可能有防盗链机制,建议将图片保存下来直接上传

Decode阶段:

模型45层, 进入 full_graph_fia

外链图片转存失败,源站可能有防盗链机制,建议将图片保存下来直接上传

错误触发条件

只有当以下条件同时满足时才会触发精度问题:

  1. Sliding Window Attention 生效(self.sliding_window is not None
  2. 图模式(进入 full_graph_fia 或图参数更新路径)
  3. MTP 头数 >= 2(seq_len > 1,next_tokens=1 产生影响)

Eager 模式下 SWA Decode 同样走了 forward_fia_slidingwindowsw-1, 1),但由于无 MTP 时 seq_len=1,该错误被掩盖。

修正后结果

  • AIME25 精度恢复正常
  • MTP draft token 接受率回升至正常水平
Logo

鲲鹏昇腾开发者社区是面向全社会开放的“联接全球计算开发者,聚合华为+生态”的社区,内容涵盖鲲鹏、昇腾资源,帮助开发者快速获取所需的知识、经验、软件、工具、算力,支撑开发者易学、好用、成功,成为核心开发者。

更多推荐