[bugfix] [NPU] fix transpose batch matmul K*B exceed 65536. (#34715)
This commit is contained in:
@@ -860,9 +860,6 @@ class Envs:
|
|||||||
SGLANG_USE_AG_AFTER_QLORA = EnvBool(False)
|
SGLANG_USE_AG_AFTER_QLORA = EnvBool(False)
|
||||||
# Enable int4x2 weights loading
|
# Enable int4x2 weights loading
|
||||||
SGLANG_NPU_W4A4_NEW_PACKING = EnvBool(False)
|
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.
|
# Use the graph-safe Triton-Ascend kernel for masked speculative KV commits.
|
||||||
SGLANG_NPU_USE_TRITON_PREFIX_KV_CACHE_STORE = EnvBoolWithAlias(
|
SGLANG_NPU_USE_TRITON_PREFIX_KV_CACHE_STORE = EnvBoolWithAlias(
|
||||||
False, deprecated_name="SGLANG_NPU_USE_TRITON_KV_CACHE_STORE"
|
False, deprecated_name="SGLANG_NPU_USE_TRITON_KV_CACHE_STORE"
|
||||||
|
|||||||
@@ -303,8 +303,20 @@ def forward_mla_core_npu(
|
|||||||
attn_output = attn_output.view(-1, m.num_local_heads, m.kv_lora_rank)
|
attn_output = attn_output.view(-1, m.num_local_heads, m.kv_lora_rank)
|
||||||
|
|
||||||
attn_output = attn_output.contiguous()
|
attn_output = attn_output.contiguous()
|
||||||
# torch.ops.npu.batch_matmul_transpose is not numerically equivalent for
|
if (
|
||||||
# Kimi-K3, so use the numerically validated torch_npu implementation.
|
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_bmm_output = torch_npu.npu_transpose_batchmatmul(
|
||||||
attn_output,
|
attn_output,
|
||||||
m.w_vc,
|
m.w_vc,
|
||||||
|
|||||||
@@ -135,8 +135,6 @@ logger = logging.getLogger(__name__)
|
|||||||
_is_hip = is_hip()
|
_is_hip = is_hip()
|
||||||
_is_npu = is_npu()
|
_is_npu = is_npu()
|
||||||
_aiter_k3_opt = get_bool_env_var("SGLANG_AITER_K3_OPT")
|
_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:
|
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
|
# but allow the NPU launcher to retain the proven attention-TP layout
|
||||||
# without a device-type branch in shared model code.
|
# without a device-type branch in shared model code.
|
||||||
self._dense_attn_tp = (
|
self._dense_attn_tp = (
|
||||||
_k3_dense_mlp_attn_tp
|
get_parallel().enable_dense_mlp_attn_tp
|
||||||
and is_dp_attention_enabled()
|
and is_dp_attention_enabled()
|
||||||
and tp_rank is None
|
and tp_rank is None
|
||||||
and tp_size 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
|
# a2a: the block runs on partial batches (shard / DP-local rows), and
|
||||||
# a TP-sharded partial sum could never be reduced across ranks that
|
# a TP-sharded partial sum could never be reduced across ranks that
|
||||||
# hold different tokens.
|
# 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
|
# NPU compatibility mode keeps DeepEP's DP-local token dispatch but
|
||||||
# uses the original TP-sharded shared MLP. Gather only that branch's
|
# uses the original TP-sharded shared MLP. Gather only that branch's
|
||||||
# inputs, then reduce-scatter its output back to the DP-local rows.
|
# inputs, then reduce-scatter its output back to the DP-local rows.
|
||||||
self._shared_experts_attn_tp_comm = (
|
self._shared_experts_attn_tp_comm = (
|
||||||
_k3_shared_experts_attn_tp
|
get_parallel().enable_shared_experts_attn_tp
|
||||||
and self._ep_a2a
|
and self._ep_a2a
|
||||||
and self._dp_attention
|
and self._dp_attention
|
||||||
and get_parallel().attn_tp_size > 1
|
and get_parallel().attn_tp_size > 1
|
||||||
|
|||||||
@@ -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.",
|
"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"),
|
NS("parallel"),
|
||||||
] = False
|
] = 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[
|
disable_attn_tp_gather: A[
|
||||||
bool,
|
bool,
|
||||||
"Disable scheduler-side attn_tp_gather (the upstream SP path "
|
"Disable scheduler-side attn_tp_gather (the upstream SP path "
|
||||||
|
|||||||
Reference in New Issue
Block a user