Prhub

#28872 【NPU】adapt_fused_rope_qk_mqa_optimize

原始 PR 作者 cen121212 合并时间 2026-06-25 08:55 文件变更 1 提交数 3 评论 3 代码增减 +1 / -4

执行摘要

NPU fused_rope_qk_mqa 阈值条件修复

NPU 上 fused_rope_qk_mqa 的调用条件过于严格,使用 query.shape[0] * query.shape[1] < 65535 可能导致某些形状无法触发融合算子,本质是修复条件使其与算子的实际支持范围一致。

建议关注此改动是否引入 NPU 上的精度或性能回归,可通读 fused_rope_qk_mqa 算子实现确认其上限约束。

讨论亮点

无相关 review 讨论,PR 获得批准后合并。

实现拆解

  1. python/sglang/srt/layers/rotary_embedding/base.pyforward_npu 方法中,将 fused_rope_qk_mqa 的调用条件从 query.shape[0] * query.shape[1] < 65535 简化为 query.shape[0] < 65535,移除对第二维的乘法检查。
  2. 同时删除原条件周围多余的空格和换行,使代码更紧凑。
  3. 该变更只涉及一行核心逻辑改动,其余为格式化清理。
文件 模块 状态 重要度
python/sglang/srt/layers/rotary_embedding/base.py 旋转位置编码 modified 4.73

关键符号

forward_npu

关键源码片段

python/sglang/srt/layers/rotary_embedding/base.py core-logic

核心修改文件,调整 fused_rope_qk_mqa 的调用条件,直接影响 NPU 上的 RoPE 计算路径。

# 变更前:条件为 query.shape[0] * query.shape[1] < 65535
# 变更后:仅检查 query.shape[0] < 65535
if fused_rope_qk_mqa is not None and query.shape[0] < 65535:
    return fused_rope_qk_mqa(
        query,
        key,
        cos_sin,
        self.rotary_dim,
        self.is_neox_style,
    )
else:
    return self.forward_native(positions, query, key, offsets)

评论区精华

没有提炼出高价值讨论线程

当前评论区没有形成足够清晰的争议点或结论,后续有更多讨论时会体现在这里。

风险与影响

风险极低:改动仅涉及一个条件表达式的简化,不会改变无融合算子时的回退行为。但需注意若 fused_rope_qk_mqa 内部对第二维有隐式限制,放宽条件可能导致运行时错误,不过从上下文看该算子应为通用实现。

影响范围仅限于 NPU 后端上使用 fused_rope_qk_mqa 的场景,缩小了原有限制,使更多形状可以受益于融合算子加速。

条件表达式修改

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论