vLLM-Ascend框架MTP投机推理在Full Graph模式下的精度问题分析与修复
作者:昇腾实战派
知识地图: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模式下, 明显偏低。
定位思路
- 排除 MTP 投机推理本身:MTP 不影响主模型精度,因此问题不在 MTP 投机推理的逻辑中。
- 锁定图模式更新逻辑:Piecewise 图模式精度正常,Full Decode Only 图模式精度异常。两者的区别在于 Full Decode Only 模式在 Decode 阶段需要更新计算图(update-graph),怀疑该路径存在参数错误。
- 对比入图/不入图参数差异:最终定位到 FIA 算子的入参在入图和不入图时设置不一致,具体是
pre_tokens和next_tokens的值存在差异:- 不入图时(eager 及非入图分支):
pre_tokens = sliding_window,next_tokens = 0 - 入图时(图参数更新分支):
pre_tokens = sliding_window - 1,next_tokens = 1
- 不入图时(eager 及非入图分支):
原因分析
对于FIA算子中 pre_tokens 和 next_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=1和next_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 | 精度 |
|---|---|---|---|---|---|
| Eager | 开 | Prefill | forward_fia | sw, 0(正确) | 正常 |
| Eager | 开 | Decode | forward_fia | sw, 0(正确) | 正常 |
| Eager | 关 | Prefill | forward_fia | sw, 0(正确) | 正常 |
| Eager | 关 | Decode (FA) | forward_fia | sw, 0(正确) | 正常 |
| Eager | 关 | Decode (SWA) | forward_fia_slidingwindow | sw-1, 1(BUG) | 正常(无MTP时seq_len=1) |
| Full Graph | 开 | Prefill | forward_fia | sw, 0(正确) | 正常 |
| Full Graph | 开 | Decode (主模型45层) | full_graph_fia | sw-1, 1 (BUG) | 已修复,sw, 0** |
| Full Graph | 开 | Decode (MTP 3层) | forward_fia | sw, 0(正确) | 正常 |
| Full Graph | 关 | Prefill | forward_fia | sw, 0(正确) | 正常 |
| Full Graph | 关 | Decode | full_graph_fia | sw-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

错误触发条件
只有当以下条件同时满足时才会触发精度问题:
- Sliding Window Attention 生效(
self.sliding_window is not None) - 图模式(进入
full_graph_fia或图参数更新路径) - MTP 头数 >= 2(seq_len > 1,
next_tokens=1产生影响)
Eager 模式下 SWA Decode 同样走了 forward_fia_slidingwindow(sw-1, 1),但由于无 MTP 时 seq_len=1,该错误被掩盖。
修正后结果
- AIME25 精度恢复正常
- MTP draft token 接受率回升至正常水平
鲲鹏昇腾开发者社区是面向全社会开放的“联接全球计算开发者,聚合华为+生态”的社区,内容涵盖鲲鹏、昇腾资源,帮助开发者快速获取所需的知识、经验、软件、工具、算力,支撑开发者易学、好用、成功,成为核心开发者。
更多推荐

所有评论(0)