执行摘要
NPU fused_rope_qk_mqa 阈值条件修复
NPU 上 fused_rope_qk_mqa 的调用条件过于严格,使用 query.shape[0] * query.shape[1] < 65535 可能导致某些形状无法触发融合算子,本质是修复条件使其与算子的实际支持范围一致。
建议关注此改动是否引入 NPU 上的精度或性能回归,可通读 fused_rope_qk_mqa 算子实现确认其上限约束。
无相关 review 讨论,PR 获得批准后合并。
NPU 上 fused_rope_qk_mqa 的调用条件过于严格,使用 query.shape[0] * query.shape[1] < 65535 可能导致某些形状无法触发融合算子,本质是修复条件使其与算子的实际支持范围一致。
建议关注此改动是否引入 NPU 上的精度或性能回归,可通读 fused_rope_qk_mqa 算子实现确认其上限约束。
无相关 review 讨论,PR 获得批准后合并。
python/sglang/srt/layers/rotary_embedding/base.py 的 forward_npu 方法中,将 fused_rope_qk_mqa 的调用条件从 query.shape[0] * query.shape[1] < 65535 简化为 query.shape[0] < 65535,移除对第二维的乘法检查。| 文件 | 模块 | 状态 | 重要度 |
|---|---|---|---|
python/sglang/srt/layers/rotary_embedding/base.py |
旋转位置编码 | modified | 4.73 |
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 链接,后续同步到相关引用后会出现在这里。
参与讨论