From c439e77872c9c89dce9cf00558cc327fcf6fb79c Mon Sep 17 00:00:00 2001 From: McZyWu Date: Mon, 24 Aug 2026 15:22:59 +0800 Subject: [PATCH] [bugfix] [NPU] fix transpose batch matmul K*B exceed 65536. (#34715) --- python/sglang/srt/environ.py | 3 -- .../modules/deepseek_v2_attention_mla_npu.py | 30 +++++++++++++------ python/sglang/srt/models/kimi_k3.py | 10 +++---- python/sglang/srt/server_args.py | 10 +++++++ 4 files changed, 36 insertions(+), 17 deletions(-) diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index ac34b2c93..b9be96fc2 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -860,9 +860,6 @@ class Envs: SGLANG_USE_AG_AFTER_QLORA = EnvBool(False) # Enable int4x2 weights loading SGLANG_NPU_W4A4_NEW_PACKING = EnvBool(False) - # Keep K3 shared experts and dense MLPs sharded over attention TP. - SGLANG_K3_SHARED_EXPERTS_ATTN_TP = EnvBool(False) - SGLANG_K3_DENSE_MLP_ATTN_TP = EnvBool(False) # Use the graph-safe Triton-Ascend kernel for masked speculative KV commits. SGLANG_NPU_USE_TRITON_PREFIX_KV_CACHE_STORE = EnvBoolWithAlias( False, deprecated_name="SGLANG_NPU_USE_TRITON_KV_CACHE_STORE" diff --git a/python/sglang/srt/hardware_backend/npu/modules/deepseek_v2_attention_mla_npu.py b/python/sglang/srt/hardware_backend/npu/modules/deepseek_v2_attention_mla_npu.py index 55949768f..7badaf1d4 100644 --- a/python/sglang/srt/hardware_backend/npu/modules/deepseek_v2_attention_mla_npu.py +++ b/python/sglang/srt/hardware_backend/npu/modules/deepseek_v2_attention_mla_npu.py @@ -303,15 +303,27 @@ def forward_mla_core_npu( attn_output = attn_output.view(-1, m.num_local_heads, m.kv_lora_rank) attn_output = attn_output.contiguous() - # torch.ops.npu.batch_matmul_transpose is not numerically equivalent for - # Kimi-K3, so use the numerically validated torch_npu implementation. - attn_bmm_output = torch_npu.npu_transpose_batchmatmul( - attn_output, - m.w_vc, - perm_x1=(1, 0, 2), - perm_x2=(0, 1, 2), - perm_y=(1, 0, 2), - ) + if ( + attn_output.shape[0] >= 65536 + or attn_output.shape[-1] * attn_output.shape[-2] >= 65536 + or m.w_vc.shape[-1] >= 65536 + ): + # npu_transpose_batchmatmul does not support dimensions >= 65536. + attn_bmm_output = torch.empty( + (attn_output.shape[0], m.num_local_heads, m.v_head_dim), + dtype=attn_output.dtype, + device=attn_output.device, + ) + torch.ops.npu.batch_matmul_transpose(attn_output, m.w_vc, attn_bmm_output) + else: + # Use the numerically validated torch_npu implementation when supported. + attn_bmm_output = torch_npu.npu_transpose_batchmatmul( + attn_output, + m.w_vc, + perm_x1=(1, 0, 2), + perm_x2=(0, 1, 2), + perm_y=(1, 0, 2), + ) attn_bmm_output = attn_bmm_output.reshape(-1, m.num_local_heads * m.v_head_dim) output, _ = m.o_proj(attn_bmm_output) diff --git a/python/sglang/srt/models/kimi_k3.py b/python/sglang/srt/models/kimi_k3.py index 5c1034f57..b1d460c2c 100644 --- a/python/sglang/srt/models/kimi_k3.py +++ b/python/sglang/srt/models/kimi_k3.py @@ -135,8 +135,6 @@ logger = logging.getLogger(__name__) _is_hip = is_hip() _is_npu = is_npu() _aiter_k3_opt = get_bool_env_var("SGLANG_AITER_K3_OPT") -_k3_shared_experts_attn_tp = envs.SGLANG_K3_SHARED_EXPERTS_ATTN_TP.get() -_k3_dense_mlp_attn_tp = envs.SGLANG_K3_DENSE_MLP_ATTN_TP.get() def _cdiv(a: int, b: int) -> int: @@ -297,7 +295,7 @@ class KimiK3MLP(nn.Module): # but allow the NPU launcher to retain the proven attention-TP layout # without a device-type branch in shared model code. self._dense_attn_tp = ( - _k3_dense_mlp_attn_tp + get_parallel().enable_dense_mlp_attn_tp and is_dp_attention_enabled() and tp_rank is None and tp_size is None @@ -553,12 +551,14 @@ class KimiK3MoE(nn.Module): # a2a: the block runs on partial batches (shard / DP-local rows), and # a TP-sharded partial sum could never be reduced across ranks that # hold different tokens. - self._shared_experts_tp1 = self._ep_a2a and not _k3_shared_experts_attn_tp + self._shared_experts_tp1 = ( + self._ep_a2a and not get_parallel().enable_shared_experts_attn_tp + ) # NPU compatibility mode keeps DeepEP's DP-local token dispatch but # uses the original TP-sharded shared MLP. Gather only that branch's # inputs, then reduce-scatter its output back to the DP-local rows. self._shared_experts_attn_tp_comm = ( - _k3_shared_experts_attn_tp + get_parallel().enable_shared_experts_attn_tp and self._ep_a2a and self._dp_attention and get_parallel().attn_tp_size > 1 diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index aeafa014e..a63a403e1 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -1221,6 +1221,16 @@ class ServerArgs: "Allow input of attention to be scattered when only using tensor parallelism, to reduce the computational load of operations such as qkv latent.", NS("parallel"), ] = False + enable_shared_experts_attn_tp: A[ + bool, + "Shard shared expert weights across the attention TP group when using an expert-parallel all-to-all backend.", + NS("parallel"), + ] = False + enable_dense_mlp_attn_tp: A[ + bool, + "Shard dense MLP weights across the attention TP group under DP attention.", + NS("parallel"), + ] = False disable_attn_tp_gather: A[ bool, "Disable scheduler-side attn_tp_gather (the upstream SP path "