[NPU] support kimi k3 on A5 and improve performance (#39589)

This commit is contained in:
zhaozx-cn
2026-09-18 16:33:54 +08:00
committed by GitHub
parent 6de4666e43
commit 8ac39c66d8
27 changed files with 922 additions and 249 deletions
@@ -8,15 +8,22 @@ import triton
import triton.language as tl
from sglang.kernels.ops.speculative.dspark.dispatch import inputs_on_cuda
from sglang.kernels.ops.speculative.reject_sampling import (
chain_speculative_sampling_triton,
)
from sglang.srt.speculative.dflash_info_v2 import DFlashDraftInputV2
from sglang.srt.speculative.dflash_utils import (
_get_or_create_chain_verify_buffers,
build_dflash_verify_target_probs,
compute_dflash_correct_drafts_and_bonus,
)
from sglang.srt.utils import is_npu
_is_npu = is_npu()
if _is_npu:
from sgl_kernel_npu.sample import chain_speculative_sampling_triton
else:
from sglang.kernels.ops.speculative.reject_sampling import (
chain_speculative_sampling_triton,
)
class AcceptSampling:
@@ -117,7 +124,12 @@ def _accept_sampling_core(
draft_token_num=verify_num_draft_tokens,
device=device,
)
uniform_samples = torch.rand((bs, gamma), dtype=torch.float32, device=device)
# The NPU implementation uses the candidate width as its row stride. The
# last value is intentionally unused because candidate slot 0 is the root.
uniform_width = candidates.shape[1] if _is_npu else gamma
uniform_samples = torch.rand(
(bs, uniform_width), dtype=torch.float32, device=device
)
uniform_samples_final = torch.rand((bs,), dtype=torch.float32, device=device)
chain_speculative_sampling_triton(
predicts=predicts,
@@ -103,6 +103,7 @@ POSITIONAL_FIELD_ORDER = (
"enable_tp_lm_head_all_to_all",
"enable_attn_tp_input_scattered",
"enable_shared_experts_attn_tp",
"shared_experts_tp_size",
"enable_dense_mlp_attn_tp",
"enable_layernorm_sp",
"disable_attn_tp_gather",
@@ -208,6 +208,12 @@ class Parallel(msgspec.Struct):
bool,
"Shard shared expert weights across the attention TP group when using an expert-parallel all-to-all backend.",
] = False
shared_experts_tp_size: A[
Optional[int],
"Shared-expert TP size for Kimi-K3 with an expert-parallel all-to-all "
"backend. Must divide attention TP size. Overrides "
"--enable-shared-experts-attn-tp when set; 1 replicates the weights.",
] = None
enable_dense_mlp_attn_tp: A[
bool,
"Shard dense MLP weights across the attention TP group under DP attention.",
@@ -121,6 +121,46 @@ def handle_context_parallelism(server_args: Any):
)
def handle_shared_experts_tp(server_args: Any):
cfg = resolving_view(server_args)
size = cfg.shared_experts_tp_size
if size is None:
return
from sglang.srt.runtime_context import derive_attention_widths
view = resolved_view(server_args)
_, attn_tp_size = derive_attention_widths(
tp_size=cfg.tp_size,
attn_cp_size=view.attn_cp_size,
dp_size=cfg.dp_size,
enable_dp_attention=view.enable_dp_attention,
)
if size < 1 or attn_tp_size % size != 0:
raise ValueError(
f"--shared-experts-tp-size ({size}) must be a positive divisor "
f"of attention TP size ({attn_tp_size})."
)
if parse_connector_type(cfg.model_path) == ConnectorType.INSTANCE:
raise ValueError(
"--shared-experts-tp-size requires a Kimi-K3 model configuration."
)
model_arch = model_config_of(server_args).hf_config.architectures[0]
if model_arch != "KimiK3ForConditionalGeneration":
raise ValueError("--shared-experts-tp-size is only supported for Kimi-K3.")
if cfg.moe_a2a_backend not in (
"deepep",
"megamoe",
"mooncake",
"ascend_fuseep",
"mori",
):
raise ValueError(
"--shared-experts-tp-size requires an expert-parallel all-to-all "
"backend (deepep, megamoe, mooncake, ascend_fuseep or mori)."
)
def handle_decode_context_parallelism(server_args: Any):
run_post_process_pass(server_args, _dcp_comm_backend_default)
cfg = resolving_view(server_args)
+2
View File
@@ -176,6 +176,7 @@ def run_resolution_pipeline(server_args: Any) -> None:
handle_elastic_ep,
handle_eplb_and_dispatch,
handle_expert_distribution_metrics,
handle_shared_experts_tp,
)
run_hook(validate_prefill_only_disable_kv_cache_args, server_args)
@@ -311,6 +312,7 @@ def run_resolution_pipeline(server_args: Any) -> None:
run_hook(handle_moe_kernel_config, server_args)
run_hook(handle_a2a_moe, server_args)
run_hook(handle_shared_experts_tp, server_args)
run_hook(handle_eplb_and_dispatch, server_args)
run_hook(handle_expert_distribution_metrics, server_args)
run_hook(handle_elastic_ep, server_args)
@@ -100,6 +100,7 @@ _OVERRIDABLE_HOOKS: FrozenSet[str] = frozenset(
"handle_a2a_moe",
"handle_eplb_and_dispatch",
"handle_expert_distribution_metrics",
"handle_shared_experts_tp",
"handle_elastic_ep",
"validate_experimental_sgl_marlin",
"handle_speculative_decoding",
@@ -291,6 +291,7 @@ def _init_parallel_groups(
attention_context_model_parallel_size=attn_cp_size,
moe_data_model_parallel_size=moe_dp_size,
decode_context_parallel_size=dcp_size,
shared_experts_tensor_parallel_size=get_parallel().shared_experts_tp_size,
duplicate_tp_group=get_disagg().enable_pdmux,
enable_symm_mem=get_exec().comm.enable_symm_mem,
# Only WORLD is extended during scale-up. The joiner's model-parallel
@@ -2142,6 +2142,7 @@ def init_model_parallel_group(
_TP: Optional[GroupCoordinator] = None
_ATTN_TP: Optional[GroupCoordinator] = None
_SHARED_EXPERTS_TP: Optional[GroupCoordinator] = None
_ATTN_CP: Optional[GroupCoordinator] = None
_DCP: Optional[GroupCoordinator] = None
@@ -2173,6 +2174,13 @@ def get_attn_tp_group() -> GroupCoordinator:
return _ATTN_TP
def get_shared_experts_tp_group() -> GroupCoordinator:
assert _SHARED_EXPERTS_TP is not None, (
"shared-expert tensor model parallel group is not initialized"
)
return _SHARED_EXPERTS_TP
def get_attn_cp_group() -> GroupCoordinator:
assert _ATTN_CP is not None, (
"attention context model parallel group is not initialized"
@@ -2263,7 +2271,7 @@ def graph_capture(stream=None):
):
with contextlib.ExitStack() as stack:
seen = {id(_TP), id(_PP)}
for group in (_DCP, _ATTN_TP, _MOE_EP, _MOE_TP):
for group in (_DCP, _ATTN_TP, _SHARED_EXPERTS_TP, _MOE_EP, _MOE_TP):
if group is not None and id(group) not in seen:
seen.add(id(group))
stack.enter_context(group.graph_capture(context))
@@ -2518,6 +2526,7 @@ def initialize_model_parallel(
recovered_rank: bool = False,
rank_offset: int = 0,
max_world_size: Optional[int] = None,
shared_experts_tensor_parallel_size: Optional[int] = None,
) -> None:
"""
Initialize model parallel groups.
@@ -2540,6 +2549,8 @@ def initialize_model_parallel(
tensor-parallel group during decoding. Must be a divisor of
tensor_model_parallel_size and is currently only supported on the
AMD HIP platform.
shared_experts_tensor_parallel_size: optional shared-expert TP width.
Must divide attention TP; subgroups never cross attention replicas.
Let's say we have a total of 8 GPUs denoted by g0 ... g7 and we
use 2 GPUs to parallelize the model tensor, and 4 GPUs to parallelize
@@ -2762,6 +2773,34 @@ def initialize_model_parallel(
max_world_size=max_world_size,
)
global _SHARED_EXPERTS_TP
assert _SHARED_EXPERTS_TP is None, "shared-expert TP group already initialized"
if (
shared_experts_tensor_parallel_size is not None
and shared_experts_tensor_parallel_size > 1
):
if shared_experts_tensor_parallel_size == attn_tp_size:
_SHARED_EXPERTS_TP = _ATTN_TP
else:
# Attention TP groups are contiguous, and the requested width
# divides each one. These groups also stay inside their PP stage.
shared_size = shared_experts_tensor_parallel_size
shared_group_ranks = [
list(range(start, start + shared_size))
for start in range(0, world_size, shared_size)
]
_SHARED_EXPERTS_TP = init_model_parallel_group(
shared_group_ranks,
get_world_group().local_rank,
backend,
use_custom_allreduce=False,
use_torch_symm_mem_allreduce=False,
group_name="shared_experts_tp",
recovered_rank=recovered_rank,
rank_offset=rank_offset,
max_world_size=max_world_size,
)
moe_ep_size = expert_model_parallel_size
moe_dp_size = moe_data_model_parallel_size
moe_tp_size = derived_widths["moe_tp_size"]
@@ -3130,7 +3169,17 @@ def destroy_model_parallel():
dwdp_mgr.cleanup()
set_global_dwdp_manager(None)
global _SHARED_EXPERTS_TP
global _ATTN_TP
global _TP
if (
_SHARED_EXPERTS_TP is not None
and _SHARED_EXPERTS_TP is not _ATTN_TP
and _SHARED_EXPERTS_TP is not _TP
):
_SHARED_EXPERTS_TP.destroy()
_SHARED_EXPERTS_TP = None
if _TP:
_TP.destroy()
_TP = None
@@ -3166,7 +3215,6 @@ def destroy_model_parallel():
_ATTN_CP.destroy()
_ATTN_CP = None
global _ATTN_TP
if _ATTN_TP:
_ATTN_TP.destroy()
_ATTN_TP = None
+9
View File
@@ -948,8 +948,17 @@ class Envs:
# ===================================================================
SGLANG_NPU_DISABLE_ACL_FORMAT_WEIGHT = EnvBool(False)
SGLANG_NPU_USE_MULTI_STREAM = EnvBool(False)
# Kimi-K3 attention-TP shared experts: overlap AG / MLP / RS with the
# routed front / DeepEP dispatch / routed GEMMs, respectively.
SGLANG_NPU_FINE_GRAINED_MOE_DUAL_STREAM = EnvBool(False)
SGLANG_NPU_USE_MLAPO = EnvBool(False)
# Fuse grouped Kimi-K3 SiTU with valid-row MXFP8 quantization before GMM2.
# Set to 0 to restore the separate SiTU + npu_dynamic_mx_quant path.
SGLANG_NPU_MOE_SITU_MXFP8_FUSED = EnvBool(True)
SGLANG_NPU_ENABLE_SPARSE_KV_OFFLOAD = EnvBool(False)
# Use FIAS V2 for DSpark MLA target verify and MHA draft paths. Graph
# replay requires torch_npu's V2 handler to update actual_seq_kvlen.
SGLANG_NPU_USE_FIAS_V2_BSND = EnvBool(False)
# BF16 wo_a: use F.linear for single-local-group decode (Flash TP8),
# retaining the original weight layout. Opt-in for A/B.
SGLANG_OPT_NPU_BF16_WO_A_GEMM = EnvBool(False)
@@ -15,6 +15,7 @@ from sglang.srt.dllm.config import DllmConfig
from sglang.srt.hardware_backend.npu.attention.ascend_torch_native_backend import (
AscendTorchNativeAttnBackend,
)
from sglang.srt.hardware_backend.npu.attention.mla_cache import gather_mla_cache_pages
from sglang.srt.hardware_backend.npu.attention.mla_preprocess import (
is_fia_nz,
is_mla_preprocess_enabled,
@@ -370,6 +371,10 @@ class AscendAttnBackend(AttentionBackend):
self.sparse_kv_manager,
)
self.use_fia = get_bool_env_var("ASCEND_USE_FIA", "False")
self.use_fias_v2_bsnd = (
get_bool_env_var("SGLANG_NPU_USE_FIAS_V2_BSND", "False")
and model_runner.spec_algorithm.is_dspark()
)
self.enable_torch_compile = get_flags().capture.enable_torch_compile
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
if (
@@ -1769,11 +1774,15 @@ class AscendAttnBackend(AttentionBackend):
k_buffer = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
v_buffer = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
kv_cached = torch.index_select(
k_buffer, 0, self.forward_metadata.flatten_prefix_block_tables
kv_cached = gather_mla_cache_pages(
k_buffer,
self.forward_metadata.flatten_prefix_block_tables,
is_nz=is_fia_nz(),
)
k_rope_cached = torch.index_select(
v_buffer, 0, self.forward_metadata.flatten_prefix_block_tables
k_rope_cached = gather_mla_cache_pages(
v_buffer,
self.forward_metadata.flatten_prefix_block_tables,
is_nz=is_fia_nz(),
).flatten(0, 1)
assert layer.kv_b_proj is not None
@@ -2346,48 +2355,107 @@ class AscendAttnBackend(AttentionBackend):
q_nope = torch.cat([q_nope, nope_padding], dim=1).contiguous()
q_rope = torch.cat([q_rope, rope_padding], dim=1).contiguous()
workspace = torch_npu._npu_fused_infer_attention_score_get_max_workspace(
q_nope,
c_kv_cache,
c_kv_cache,
query_rope=q_rope,
key_rope=k_rope_cache,
num_heads=self.q_head_num_padding,
num_key_value_heads=layer.tp_k_head_num,
input_layout="TND",
scale=layer.scaling,
antiquant_mode=0,
antiquant_scale=None,
block_table=block_table,
block_size=self.page_size,
sparse_mode=3,
atten_mask=self.mtp_mask,
actual_seq_lengths=actual_seq_lengths,
actual_seq_lengths_kv=actual_seq_lengths_kv,
)
attn_output = torch.empty_like(q_nope, dtype=q.dtype, device=q.device)
softmax_lse = torch.empty(1, dtype=q.dtype, device=q.device)
torch_npu.npu_fused_infer_attention_score.out(
q_nope,
c_kv_cache,
c_kv_cache,
query_rope=q_rope,
key_rope=k_rope_cache,
num_heads=self.q_head_num_padding,
num_key_value_heads=layer.tp_k_head_num,
input_layout="TND",
scale=layer.scaling,
antiquant_mode=0,
antiquant_scale=None,
block_table=block_table,
block_size=self.page_size,
sparse_mode=3,
atten_mask=self.mtp_mask,
actual_seq_lengths=actual_seq_lengths,
actual_seq_lengths_kv=actual_seq_lengths_kv,
workspace=workspace,
out=[attn_output, softmax_lse],
)
num_query_heads = q_nope.shape[1]
if self.use_fias_v2_bsnd:
# The existing paged MLA cache is [block, KV_N, page, D].
# V2 consumes it with BNSD queries; keep the cache unchanged.
batch_size = len(actual_seq_lengths_kv)
query_seq_len = self.speculative_num_draft_tokens
assert q_nope.shape[0] == batch_size * query_seq_len, (
"FIAS V2 target verify requires one fixed draft block per request"
)
if batch_size == 0:
attn_output = torch.empty_like(q_nope)
else:
q_nope_bnsd = (
q_nope.view(
batch_size,
query_seq_len,
num_query_heads,
self.kv_lora_rank,
)
.transpose(1, 2)
.contiguous()
)
q_rope_bnsd = (
q_rope.view(
batch_size,
query_seq_len,
num_query_heads,
self.qk_rope_head_dim,
)
.transpose(1, 2)
.contiguous()
)
attn_output, _ = torch_npu.npu_fused_infer_attention_score_v2(
q_nope_bnsd,
c_kv_cache,
c_kv_cache,
query_rope=q_rope_bnsd,
key_rope=k_rope_cache,
num_query_heads=num_query_heads,
num_key_value_heads=layer.tp_k_head_num,
input_layout="BNSD",
softmax_scale=layer.scaling,
block_table=block_table,
block_size=self.page_size,
sparse_mode=3,
atten_mask=self.mtp_mask,
actual_seq_qlen=[query_seq_len] * batch_size,
actual_seq_kvlen=actual_seq_lengths_kv,
pre_tokens=FULL_ATTENTION_WINDOW,
next_tokens=0,
)
attn_output = (
attn_output.transpose(1, 2)
.contiguous()
.reshape(-1, num_query_heads, self.kv_lora_rank)
)
else:
workspace = (
torch_npu._npu_fused_infer_attention_score_get_max_workspace(
q_nope,
c_kv_cache,
c_kv_cache,
query_rope=q_rope,
key_rope=k_rope_cache,
num_heads=num_query_heads,
num_key_value_heads=layer.tp_k_head_num,
input_layout="TND",
scale=layer.scaling,
antiquant_mode=0,
antiquant_scale=None,
block_table=block_table,
block_size=self.page_size,
sparse_mode=3,
atten_mask=self.mtp_mask,
actual_seq_lengths=actual_seq_lengths,
actual_seq_lengths_kv=actual_seq_lengths_kv,
)
)
attn_output = torch.empty_like(q_nope, dtype=q.dtype, device=q.device)
softmax_lse = torch.empty(1, dtype=q.dtype, device=q.device)
torch_npu.npu_fused_infer_attention_score.out(
q_nope,
c_kv_cache,
c_kv_cache,
query_rope=q_rope,
key_rope=k_rope_cache,
num_heads=num_query_heads,
num_key_value_heads=layer.tp_k_head_num,
input_layout="TND",
scale=layer.scaling,
antiquant_mode=0,
antiquant_scale=None,
block_table=block_table,
block_size=self.page_size,
sparse_mode=3,
atten_mask=self.mtp_mask,
actual_seq_lengths=actual_seq_lengths,
actual_seq_lengths_kv=actual_seq_lengths_kv,
workspace=workspace,
out=[attn_output, softmax_lse],
)
attn_output = attn_output[:, : layer.tp_q_head_num, :]
attn_output = attn_output.view(-1, layer.tp_q_head_num * layer.v_head_dim)
if (
@@ -2,20 +2,9 @@ import math
from typing import Optional
import torch
from sgl_kernel_npu.fla.kda_chunk_delta_h import (
chunk_gated_delta_rule_fwd_h_npu,
)
from sgl_kernel_npu.fla.kda_gate import fused_kda_gate_npu
from sgl_kernel_npu.fla.kda_prefill import (
chunk_gla_fwd_o_gk_npu,
recompute_w_u_fwd_npu,
)
from sgl_kernel_npu.fla.kda_target_verify import kda_target_verify_npu
from sgl_kernel_npu.fla.solve_tril import solve_tril_npu
from sgl_kernel_npu.fla.utils import prepare_chunk_indices
from sglang.kernels.ops.attention.fla.cumsum import chunk_local_cumsum
from sglang.kernels.ops.attention.fla.kda import chunk_kda_scaled_dot_kkt_fwd
from sglang.kernels.ops.attention.fla.l2norm import l2norm_fwd
from sglang.srt.layers.attention.linear.kda_backend import (
KDAAttnBackend,
@@ -48,67 +37,54 @@ class _AscendKDAExtendKernel:
q = l2norm_fwd(q.contiguous())
k = l2norm_fwd(k.contiguous())
v = v.contiguous()
g = g.contiguous()
beta = beta.contiguous()
chunk_indices = prepare_chunk_indices(query_start_loc, chunk_size)
g = chunk_local_cumsum(
g.contiguous(),
chunk_size=chunk_size,
scale=_LOG2_E,
# chunk_kda_fwd accepts one initial state per logical sequence, while
# SGLang owns a slot-indexed persistent pool. Gather the active slots in
# canonical contiguous [N, H, V, K] layout and scatter final_state back.
num_sequences = query_start_loc.shape[0] - 1
source_indices = cache_indices[:num_sequences].to(torch.long)
valid_state_mask = source_indices >= 0
# Forward metadata may use -1 for a padded request. index_select would
# otherwise read the last cache slot and index_copy_ would overwrite it.
# Slot 0 is a gather placeholder for that padded row; its computed result
# is irrelevant because padded rows are filtered before state writeback.
gather_indices = source_indices.clamp_min(0)
initial_state = (
ssm_states.index_select(0, gather_indices)
.to(dtype=torch.float32)
.contiguous()
)
scale = k.shape[-1] ** -0.5
query_start_loc = query_start_loc.to(dtype=torch.int64).contiguous()
outputs = torch.ops.npu.chunk_kda_fwd(
q,
k,
v,
g,
beta,
scale=scale,
initial_state=initial_state,
output_final_state=True,
cu_seqlens=query_start_loc,
chunk_indices=chunk_indices,
chunk_size=chunk_size,
layout="BSND",
safe_gate=False,
use_gate_in_kernel=False,
state_v_first=True,
output_h=return_intermediate_states,
)
out, final_state, chunk_states = outputs[0], outputs[1], outputs[10]
valid_positions = valid_state_mask.nonzero(as_tuple=False).flatten()
ssm_states.index_copy_(
0,
source_indices.index_select(0, valid_positions),
final_state.index_select(0, valid_positions).to(dtype=ssm_states.dtype),
)
triangular, query_key = chunk_kda_scaled_dot_kkt_fwd(
q=q,
k=k,
gk=g,
beta=beta,
scale=k.shape[-1] ** -0.5,
cu_seqlens=query_start_loc,
output_dtype=torch.float32,
)
triangular = solve_tril_npu(
A=triangular,
cu_seqlens=query_start_loc,
output_dtype=k.dtype,
)
w, u, gated_k = recompute_w_u_fwd_npu(
k=k,
v=v,
beta=beta,
A=triangular,
gk=g,
cu_seqlens=query_start_loc,
chunk_indices=chunk_indices,
)
del triangular
chunk_states, new_values = chunk_gated_delta_rule_fwd_h_npu(
k=gated_k,
w=w,
u=u,
gk=g,
initial_state=ssm_states,
initial_state_indices=cache_indices,
cu_seqlens=query_start_loc,
chunk_indices=chunk_indices,
use_exp2=True,
)
del w, u, gated_k
out = chunk_gla_fwd_o_gk_npu(
q=q,
v=new_values,
g=g,
A=query_key,
h=chunk_states,
out=v,
scale=k.shape[-1] ** -0.5,
cu_seqlens=query_start_loc,
chunk_size=chunk_size,
chunk_indices=chunk_indices,
)
del query_key, new_values
if return_intermediate_states:
return out, chunk_states.transpose(-1, -2).contiguous()
return out, chunk_states
return out
@@ -304,6 +280,7 @@ class AscendKDAAttnBackend(KDAAttnBackend):
layer, a, b
)
track_ssm = self.forward_metadata.has_mamba_track_mask
core_attn_out = self.kernel_dispatcher.extend(
q=q,
k=k,
@@ -323,6 +300,7 @@ class AscendKDAAttnBackend(KDAAttnBackend):
self.forward_metadata.track_ssm_h_src if track_ssm else None
),
)
if track_ssm:
core_attn_out, h = core_attn_out
self._track_mamba_state_extend(
@@ -458,7 +436,6 @@ class AscendKDAAttnBackend(KDAAttnBackend):
intermediate_states_buffer=intermediate_state,
intermediate_state_indices=intermediate_indices,
cache_steps=draft_token_num,
lower_bound=None,
gates_are_preactivated=True,
)
if dense_token_indices is None:
@@ -599,6 +576,8 @@ class AscendKDAHybridLinearAttnBackend:
mamba_steps_to_track,
)
else:
# No-op self-copy for non-tracked entries so we never run
# bool-mask indexing (aten::nonzero) or a host numel check.
track_mask = mamba_steps_to_track >= 0
src_slots = torch.where(
track_mask, dst_indices_tensor, mamba_track_indices
@@ -0,0 +1,23 @@
"""Read logical pages from the explicit PA-NZ storage of the NPU MLA cache."""
import torch
def gather_mla_cache_pages(
cache: torch.Tensor, block_ids: torch.Tensor, *, is_nz: bool
) -> torch.Tensor:
"""Return selected pages in logical [blocks, page_size, 1, head_dim] order.
NZ buffers retain that public shape, but their physical contents are
[blocks, head_dim // 16, page_size, 16]. Restore token-major order before
projecting cached latent vectors or concatenating their RoPE features.
"""
pages = torch.index_select(cache, 0, block_ids)
if not is_nz:
return pages
page_size, head_dim = cache.shape[1], cache.shape[-1]
return (
pages.view(block_ids.numel(), head_dim // 16, page_size, 16)
.permute(0, 2, 1, 3)
.reshape(block_ids.numel(), page_size, 1, head_dim)
)
@@ -24,12 +24,14 @@ def is_mla_preprocess_enabled() -> bool:
@lru_cache(maxsize=1)
def is_fia_nz() -> bool:
is_fia_nz_ = get_bool_env_var("SGLANG_USE_FIA_NZ")
if is_fia_nz_:
assert is_mla_preprocess_enabled(), (
"SGLANG_USE_FIA_NZ must be enable with SGLANG_NPU_USE_MLAPO"
)
return is_fia_nz_
"""Whether MLA KV cache uses the FIA NZ physical layout.
This is a cache-layout choice, not an MLAPO-only optimization. MLAPO can
write NZ cache directly, while the ordinary MLA path writes the same
layout through ``NPUMLATokenToKVPool``. Keeping the switch independent
lets models such as Kimi-K3 use FIA NZ without selecting MLAPO.
"""
return get_bool_env_var("SGLANG_USE_FIA_NZ")
def round_up(val: int, align: int) -> int:
@@ -12,7 +12,7 @@ non-NPU hosts.
from __future__ import annotations
import threading
from concurrent.futures import ThreadPoolExecutor
from contextlib import AbstractContextManager, contextmanager
from functools import partial
from typing import TYPE_CHECKING, Any, Callable, Dict, Optional
@@ -63,6 +63,13 @@ class NPUCudaGraphBackend(BaseCudaGraphBackend):
self._enable_torch_compile = getattr(
cuda_graph_runner, "enable_torch_compile", False
)
# Reuse one device-bound worker for graph input updates.
self._update_executor = ThreadPoolExecutor(
max_workers=1,
thread_name_prefix="npu-graph-update",
initializer=self._device_module.set_device,
initargs=(self._device_id,),
)
@contextmanager
def capture_session(self, stream):
@@ -150,8 +157,10 @@ class NPUCudaGraphBackend(BaseCudaGraphBackend):
attr_type: Any = None,
cpu_update_input: list = None,
) -> Any:
"""Rebind seq_lens on the recorded NPU graph in a background
thread, then replay. Used when the model is not deepseek-nsa.
"""Rebind seq_lens on the recorded NPU graph, then replay.
NPUGraph.update must complete before replay can consume the updated
KV lengths. Used when the model is not deepseek-nsa.
Two calling conventions:
1. (legacy) seq_lens + attr_name + attr_type:
@@ -166,17 +175,15 @@ class NPUCudaGraphBackend(BaseCudaGraphBackend):
graph = self._graphs[shape_key]
def _update():
self._device_module.set_device(self._device_id)
graph.update(cpu_update_input=cpu_update_input)
thread = threading.Thread(target=_update)
thread.start()
update_future = self._update_executor.submit(
graph.update, cpu_update_input=cpu_update_input
)
update_future.result()
graph.replay()
thread.join()
return self._outputs[shape_key]
def cleanup(self) -> None:
self._update_executor.shutdown(wait=True, cancel_futures=True)
self._graphs.clear()
self._outputs.clear()
self._pool = None
@@ -131,6 +131,10 @@ class NPUGraphRunner(DecodeCudaGraphRunner):
in ("MiMoV2ForCausalLM", "MiMoV2FlashForCausalLM", "Step3p5ForCausalLM")
for arch in (model_runner.model_config.hf_config.architectures or [])
)
self.use_fias_v2_bsnd = (
envs.SGLANG_NPU_USE_FIAS_V2_BSND.get()
and model_runner.spec_algorithm.is_dspark()
)
def _init_arch_map(self):
if self.is_dllm:
@@ -156,6 +160,14 @@ class NPUGraphRunner(DecodeCudaGraphRunner):
def _create_device_graph(self):
return torch.npu.NPUGraph()
def _uses_v2_seq_len_update(self):
# IDLE DP ranks replay the same target-verify graph as active ranks.
# Select the handler from the captured graph, not the runtime mode;
# a V1 update key is ignored by V2 and leaves stale KV lengths behind.
return self.if_use_v2 or (
self.use_fias_v2_bsnd and self.capture_forward_mode.is_target_verify()
)
def _capture_graph(self, graph, pool, stream, run_once_fn):
if self.enable_torch_compile:
skip_guard_context = torch.compiler.set_stance(skip_guard_eval_unsafe=True)
@@ -175,12 +187,12 @@ class NPUGraphRunner(DecodeCudaGraphRunner):
return out
def _get_update_attr_name(self):
if self.if_use_v2:
if self._uses_v2_seq_len_update():
return self.attr_name["TARGET_VERIFY"]
return self.attr_name[AttentionArch.MLA]
def _get_update_attr_type(self):
if self.if_use_v2:
if self._uses_v2_seq_len_update():
return self.attr_type["TARGET_VERIFY"]
return self.attr_type[AttentionArch.MLA]
@@ -22,6 +22,34 @@ if is_npu():
import torch_npu
def _mla_fia_nz_scatter_indices(
loc: torch.Tensor, head_dim: int, page_size: int
) -> torch.Tensor:
"""Return physical rows for token-wise writes into an MLA NZ cache.
The storage allocation remains page-major ``[page, slot, 1, D]`` for
transfer and bookkeeping compatibility. FIA reads that storage as
``[page, 1, D / 16, page_size, 16]``. A token-major scatter would therefore
write the wrong physical rows, so every logical token expands to its
``D / 16`` NZ tiles.
"""
if head_dim % 16:
raise ValueError(
"FIA NZ MLA cache requires a head dimension divisible by 16, "
f"got {head_dim}."
)
if page_size <= 0:
raise ValueError(f"page_size must be positive, got {page_size}.")
num_tiles = head_dim // 16
page = torch.div(loc, page_size, rounding_mode="floor")
slot = torch.remainder(loc, page_size)
tiles = torch.arange(num_tiles, dtype=loc.dtype, device=loc.device)
# Flatten [token, tile] in the same order as source.view(T, tiles, 16).
rows = ((page[:, None] * num_tiles + tiles) * page_size) + slot[:, None]
return rows.reshape(-1, 1)
def _init_npu_conv_state(
conv_state_in,
conv_state_shape,
@@ -562,6 +590,10 @@ class NPUMLATokenToKVPool(MLATokenToKVPool):
indexer_layer_ids: Optional[Sequence[int]] = None,
kv_cache_dim: Optional[int] = None,
):
# MLAPO historically owned NZ writes. Keep the allocation unchanged and
# write into the NZ-addressed view below so ordinary MLA (including
# Kimi-K3 MTP) can use FIA NZ without MLAPO.
self.use_fia_nz = get_bool_env_var("SGLANG_USE_FIA_NZ")
super(MLATokenToKVPool, self).__init__(
size=size,
page_size=page_size,
@@ -851,6 +883,12 @@ class NPUMLATokenToKVPool(MLATokenToKVPool):
packed.view(-1, 1, self.kv_cache_dim),
)
return
if cache_v is None:
cache_k, cache_v = cache_k.split(
[self.kv_lora_rank, self.qk_rope_head_dim], dim=-1
)
if cache_k.dtype != self.dtype:
cache_k = cache_k.to(self.dtype)
cache_v = cache_v.to(self.dtype)
@@ -859,10 +897,9 @@ class NPUMLATokenToKVPool(MLATokenToKVPool):
cache_k = cache_k.view(self.store_dtype)
cache_v = cache_v.view(self.store_dtype)
if cache_v is None:
cache_k, cache_v = cache_k.split(
[self.kv_lora_rank, self.qk_rope_head_dim], dim=-1
)
if self.use_fia_nz:
self._set_fia_nz_kv_buffer(layer_id, loc, cache_k, cache_v)
return
torch_npu.npu_scatter_nd_update_(
self.k_buffer[layer_id - self.start_layer].view(-1, 1, self.kv_lora_rank),
@@ -877,6 +914,28 @@ class NPUMLATokenToKVPool(MLATokenToKVPool):
cache_v.view(-1, 1, self.qk_rope_head_dim),
)
def _set_fia_nz_kv_buffer(
self,
layer_id: int,
loc: torch.Tensor,
cache_k: torch.Tensor,
cache_v: torch.Tensor,
) -> None:
"""Store MLA latent and RoPE KV tensors in FIA's NZ tile order."""
def scatter(cache: torch.Tensor, values: torch.Tensor, head_dim: int):
num_tiles = head_dim // 16
indices = _mla_fia_nz_scatter_indices(loc, head_dim, self.page_size)
# Destination rows are ordered [page, tile, slot]. Source rows use
# the matching [token, tile] order after this reshape.
dst = cache.view(-1, 1, num_tiles, self.page_size, 16).view(-1, 16)
src = values.contiguous().view(-1, num_tiles, 16).view(-1, 16)
torch_npu.npu_scatter_nd_update_(dst, indices, src)
offset = layer_id - self.start_layer
scatter(self.k_buffer[offset], cache_k, self.kv_lora_rank)
scatter(self.v_buffer[offset], cache_v, self.qk_rope_head_dim)
def set_index_k_buffer(
self,
layer_id: int,
@@ -7,7 +7,6 @@ import torch.nn.functional as F
from sglang.srt.distributed.communication_op import (
tensor_model_parallel_all_gather,
)
from sglang.srt.layers.activation import GeluAndMul
from sglang.srt.runtime_context import get_parallel
@@ -148,8 +147,35 @@ class NPUSitu(BaseActivation):
)
class NPUSituMXFP8Quant(BaseActivation):
"""A5 AscendC grouped SiTU with valid-row MXFP8 quantization."""
def __init__(self, *, beta: float = 4.0, linear_beta: float = 25.0):
from sgl_kernel_npu.activation.situ_mxfp8_quant import situ_mxfp8_quant
self.situ_mxfp8_quant = situ_mxfp8_quant
self.beta = float(beta)
self.linear_beta = float(linear_beta)
def _apply_activation(
self,
hidden_states: torch.Tensor,
group_list: torch.Tensor,
group_list_type: int,
):
return self.situ_mxfp8_quant(
hidden_states,
group_list,
group_list_type,
beta=self.beta,
linear_beta=self.linear_beta,
)
class NPUGeluAndMul(BaseActivation):
def __init__(self):
from sglang.srt.layers.activation import GeluAndMul
self._gelu = GeluAndMul()
def _apply_activation(self, hidden_states: torch.Tensor):
@@ -298,10 +298,17 @@ class NPUW4A8MXFP4MoEMethod(_NPUMoEMethodBase):
weight.data, weight_scale.data
)
# The refactored NPU dispatchers currently support BF16 and INT8.
# Keep dispatch in BF16 and quantize to MXFP8 immediately before GMM.
if weight_prefix == "w13":
self._set_dispatcher_output_dtype(layer, "bf16")
# A5 DeepEP low-latency dispatch quantizes the valid received rows to
# MXFP8 and returns the matching E8M0 block scales. Keep normal mode
# in BF16 because A5 MXFP8 normal dispatch is intranode-only; this also
# preserves multi-node prefill when ``deepep-mode=auto``.
if weight_prefix == "w13" and hasattr(layer, "dispatcher"):
layer.dispatcher.set_quant_config(
{
"normal_dispatcher_output_dtype": "bf16",
"low_latency_dispatcher_output_dtype": "mxfp8",
}
)
def apply(
self,
@@ -7,10 +7,12 @@ from typing import TYPE_CHECKING, Any, Optional
import torch
from sglang.srt.environ import envs
from sglang.srt.hardware_backend.npu.moe.activation import (
AllGatherActivationWrapper,
NPUGeluAndMul,
NPUSitu,
NPUSituMXFP8Quant,
NPUSwiglu,
NPUSwigluDeepEPKernel,
NPUSwigluMxfp8Quant,
@@ -35,16 +37,16 @@ from sglang.srt.layers.moe.moe_runner.base import (
)
if TYPE_CHECKING:
from sglang.srt.layers.moe.token_dispatcher.ascend_tp import (
AscendTPCombineInput,
AscendTPDispatchOutput,
)
from sglang.srt.layers.moe.token_dispatcher.deepep import (
DeepEPLLCombineInput,
DeepEPLLDispatchOutput,
DeepEPNormalCombineInput,
DeepEPNormalDispatchOutput,
)
from sglang.srt.layers.moe.token_dispatcher.ascend_tp import (
AscendTPDispatchOutput,
AscendTPCombineInput,
)
from sglang.srt.layers.moe.utils import (
MoeRunnerBackend,
@@ -111,13 +113,25 @@ class AscendRunnerCore(MoeRunnerCore):
kernel, (NPUW4A8Int8MoEMethod, NPUW8A8Int8MoEMethod)
)
if config.activation == "situ":
self.activation = NPUSitu(
need_quant=is_quant_kernel,
beta=(
config.gemm1_alpha if config.gemm1_alpha is not None else 4.0
),
linear_beta=config.gemm1_clamp_limit,
)
beta = config.gemm1_alpha if config.gemm1_alpha is not None else 4.0
if (
isinstance(kernel, NPUW4A8MXFP4MoEMethod)
and envs.SGLANG_NPU_MOE_SITU_MXFP8_FUSED.get()
):
if config.gemm1_clamp_limit is None:
raise ValueError(
"fused SiTU MXFP8 quantization requires gemm1_clamp_limit"
)
self.activation = NPUSituMXFP8Quant(
beta=beta,
linear_beta=config.gemm1_clamp_limit,
)
else:
self.activation = NPUSitu(
need_quant=is_quant_kernel,
beta=beta,
linear_beta=config.gemm1_clamp_limit,
)
else:
self.activation = NPUSwigluDeepEPKernel(
need_quant=is_quant_kernel,
@@ -129,6 +143,15 @@ class AscendRunnerCore(MoeRunnerCore):
# 1. Choose the base activation according to the quant method
if isinstance(kernel, (NPUW4A8Int8MoEMethod, NPUW8A8Int8MoEMethod)):
inner = NPUSwigluQuant()
elif config.activation == "situ":
# Grouped SiTU (Kimi-K3). need_quant=False: the MXFP4 / BF16
# gmm2 requantizes the activations itself, so no quant is
# fused here. Matches the DeepEP branch below.
inner = NPUSitu(
need_quant=False,
beta=config.gemm1_alpha if config.gemm1_alpha is not None else 4.0,
linear_beta=config.gemm1_clamp_limit,
)
else:
if config.activation == "npu_swiglu_oai":
# NPUSwigluOAI requires the runner config to pass
@@ -194,7 +217,12 @@ class AscendRunnerCore(MoeRunnerCore):
# Grouped-row activations require dispatch metadata.
if isinstance(
self.activation,
(NPUSwigluDeepEPKernel, NPUSitu, NPUSwigluMxfp8Quant),
(
NPUSwigluDeepEPKernel,
NPUSitu,
NPUSituMXFP8Quant,
NPUSwigluMxfp8Quant,
),
):
hidden_states, pertoken_scale = self.activation._apply_activation(
hidden_states,
@@ -53,6 +53,7 @@ from sglang.srt.layers.quantization.compressed_tensors.schemes import (
CompressedTensorsWNA16MoE,
CompressedTensorsWNA16TritonMoE,
NPUCompressedTensorsW4A8Int8DynamicMoE,
NPUCompressedTensorsW4A8mxfp4MoE,
NPUCompressedTensorsW4A16Int4DynamicMoE,
NPUCompressedTensorsW8A8Int8,
NPUCompressedTensorsW8A8Int8DynamicMoE,
@@ -225,7 +226,10 @@ class CompressedTensorsConfig(QuantizationConfig):
# Detect MXFP4 before the scheme-based path: MXFP4 uses a
# dedicated FusedMoEMethodBase (Mxfp4MoEMethod) that already
# handles all MoE backends, bypassing the scheme abstraction.
if self._is_mxfp4_moe(layer_name=prefix):
# On NPU the dedicated Mxfp4MoEMethod does not apply, so fall
# through to the scheme-based path and let get_moe_scheme select
# NPUCompressedTensorsW4A8mxfp4MoE.
if self._is_mxfp4_moe(layer_name=prefix) and not _is_npu:
from sglang.srt.layers.quantization.mxfp4 import Mxfp4MoEMethod
logger.info_once(
@@ -819,15 +823,20 @@ class CompressedTensorsConfig(QuantizationConfig):
weight_quant = scheme_dict.get("weights")
input_quant = scheme_dict.get("input_activations")
# MXFP4 MoE on NPU is served by NPUCompressedTensorsW4A8mxfp4MoE. Detect
# it before the WNA16 branch: MXFP4 weights are FP4 (float) group-32 so
# `_is_wNa16_group_channel` / `_is_dynamic_token_w4a8` would otherwise
# misroute them to the INT4 WNA16 or W4A8-int8 schemes.
if _is_npu and self._is_mxfp4_moe(layer_name=layer_name):
logger.info_once("Using NPUCompressedTensorsW4A8mxfp4MoE")
return NPUCompressedTensorsW4A8mxfp4MoE()
if self._is_wNa16_group_channel(weight_quant, input_quant):
if not _is_npu:
if (
self._is_mxint4a16(weight_quant, input_quant)
and get_moe_runner_backend().is_flashinfer_trtllm()
):
logger.info_once(
"Using CompressedTensorsMxInt4MoE with flashinfer_trtllm backend"
)
return CompressedTensorsMxInt4MoE(self, weight_quant=weight_quant)
elif _is_hip:
logger.info_once("Using CompressedTensorsWNA16TritonMoE (ROCm)")
@@ -9,6 +9,7 @@ from .compressed_tensors_w4a4_nvfp4 import CompressedTensorsW4A4Fp4
from .compressed_tensors_w4a4_nvfp4_moe import CompressedTensorsW4A4Nvfp4MoE
from .compressed_tensors_w4a8_fp8_moe import CompressedTensorsW4AFP8MoE
from .compressed_tensors_w4a8_int8_moe import NPUCompressedTensorsW4A8Int8DynamicMoE
from .compressed_tensors_w4a8_mxfp4_moe import NPUCompressedTensorsW4A8mxfp4MoE
from .compressed_tensors_w8a8_fp8 import CompressedTensorsW8A8Fp8
from .compressed_tensors_w8a8_fp8_moe import CompressedTensorsW8A8Fp8MoE
from .compressed_tensors_w8a8_int8 import (
@@ -43,4 +44,5 @@ __all__ = [
"NPUCompressedTensorsW4A8Int8DynamicMoE",
"CompressedTensorsMxInt4MoE",
"CompressedTensorsW4AFP8MoE",
"NPUCompressedTensorsW4A8mxfp4MoE",
]
@@ -0,0 +1,154 @@
from __future__ import annotations
import logging
from typing import TYPE_CHECKING
import torch
from sglang.srt.hardware_backend.npu.quantization.moe_methods import (
NPUW4A8MXFP4MoEMethod,
)
from sglang.srt.layers.moe.moe_runner import MoeRunner, MoeRunnerConfig
from sglang.srt.layers.moe.moe_runner.ascend import AscendQuantInfo
from sglang.srt.layers.moe.utils import MoeRunnerBackend, get_moe_runner_backend
from sglang.srt.layers.quantization.compressed_tensors.schemes import (
CompressedTensorsMoEScheme,
)
from sglang.srt.utils import set_weight_attrs
if TYPE_CHECKING:
from sglang.srt.layers.moe.token_dispatcher import (
CombineInput,
StandardDispatchOutput,
)
__all__ = ["NPUCompressedTensorsW4A8mxfp4MoE"]
logger = logging.getLogger(__name__)
class NPUCompressedTensorsW4A8mxfp4MoE(CompressedTensorsMoEScheme):
"""Compressed-tensors MXFP4 MoE scheme for Ascend NPU.
Follows the same structure as the other NPU MoE schemes: the MXFP4
payload / scale layout transforms live in ``NPUW4A8MXFP4MoEMethod``
(shared with the ModelSlim path), and the runner drives w13 -> activation
-> w2 through the Ascend ``MoeRunner``.
"""
def __init__(self):
self.group_size = 32
self.w13_kernel = NPUW4A8MXFP4MoEMethod()
self.w2_kernel = NPUW4A8MXFP4MoEMethod()
def create_weights(
self,
layer: torch.nn.Module,
num_experts: int,
hidden_size: int,
intermediate_size_per_partition: int,
params_dtype: torch.dtype,
**extra_weight_attrs,
):
from sglang.srt.layers.moe.fused_moe_triton import FusedMoeWeightScaleSupported
layer.params_dtype = params_dtype
# Weights are stored as packed FP4 (two FP4 items per byte), so the
# K dimension is halved. The compressed-tensors loader writes the
# payload under the `_packed` suffix; process_weights_after_loading
# renames it to the kernel's `w{13,2}_weight`.
w13_weight = torch.nn.Parameter(
torch.empty(
num_experts,
2 * intermediate_size_per_partition,
hidden_size // 2,
requires_grad=False,
dtype=torch.uint8,
),
requires_grad=False,
)
layer.register_parameter("w13_weight_packed", w13_weight)
set_weight_attrs(w13_weight, extra_weight_attrs)
w2_weight = torch.nn.Parameter(
torch.empty(
num_experts,
hidden_size,
intermediate_size_per_partition // 2,
dtype=torch.uint8,
),
requires_grad=False,
)
layer.register_parameter("w2_weight_packed", w2_weight)
set_weight_attrs(w2_weight, extra_weight_attrs)
# Weight scales: one e8m0 block scale per 32-value group.
extra_weight_attrs.update(
{"quant_method": FusedMoeWeightScaleSupported.GROUP.value}
)
w13_weight_scale = torch.nn.Parameter(
torch.empty(
num_experts,
2 * intermediate_size_per_partition,
hidden_size // self.group_size,
dtype=torch.uint8,
),
requires_grad=False,
)
layer.register_parameter("w13_weight_scale", w13_weight_scale)
set_weight_attrs(w13_weight_scale, extra_weight_attrs)
w2_weight_scale = torch.nn.Parameter(
torch.empty(
num_experts,
hidden_size,
intermediate_size_per_partition // self.group_size,
dtype=torch.uint8,
),
requires_grad=False,
)
layer.register_parameter("w2_weight_scale", w2_weight_scale)
set_weight_attrs(w2_weight_scale, extra_weight_attrs)
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
# The compressed-tensors MXFP4 loader stores the packed FP4 payloads
# under the `_packed` suffix; rename them to the kernel's expected
# names before delegating the NZ layout / scale transform.
layer.w13_weight = torch.nn.Parameter(
layer.w13_weight_packed.data, requires_grad=False
)
delattr(layer, "w13_weight_packed")
layer.w2_weight = torch.nn.Parameter(
layer.w2_weight_packed.data, requires_grad=False
)
delattr(layer, "w2_weight_packed")
self.w13_kernel.process_weights_after_loading(layer, "w13")
self.w2_kernel.process_weights_after_loading(layer, "w2")
def create_moe_runner(
self, layer: torch.nn.Module, moe_runner_config: MoeRunnerConfig
):
layer.w13_kernel = self.w13_kernel
layer.w2_kernel = self.w2_kernel
moe_runner_config.layer = layer
self.moe_runner_config = moe_runner_config
backend = get_moe_runner_backend()
if backend.is_auto():
backend = MoeRunnerBackend.ASCEND
self.runner = MoeRunner(backend, moe_runner_config)
def apply_weights(
self,
layer: torch.nn.Module,
dispatch_output: StandardDispatchOutput,
) -> CombineInput:
quant_info = AscendQuantInfo(
w13_weight=layer.w13_weight,
w2_weight=layer.w2_weight,
w13_weight_scale=layer.w13_weight_scale,
w2_weight_scale=layer.w2_weight_scale,
)
return self.runner.run(dispatch_output, quant_info)
@@ -1166,7 +1166,7 @@ class KVCacheConfigurator:
# case above: None => the pool skips SpeculativeState).
speculative_num_draft_tokens=(
None
if get_disagg().disaggregation_mode == "prefill"
if get_disagg().disaggregation_mode == "prefill" and not _is_npu
else max_speculative_num_draft_tokens()
),
speculative_eagle_topk=get_spec().speculative_eagle_topk,
+1 -8
View File
@@ -754,15 +754,8 @@ class MambaPool:
)
if speculative_num_draft_tokens is not None:
if _is_npu:
temporal_state = temporal_state.transpose(-1, -2)
temporal_state_shape = (
*temporal_state_shape[:-2],
temporal_state_shape[-1],
temporal_state_shape[-2],
)
# Cache intermediate SSM states per draft token during target verify
# Shape: [num_layers, size + 1, speculative_num_draft_tokens, HV, K, V]
# Shape: [num_layers, size + 1, speculative_num_draft_tokens, HV, V, K]
#
# ReplaySSM spec-verify owns rollback via the ring + cursors (the
# verify kernel never writes per-draft snapshots; the commit never
+235 -74
View File
@@ -22,6 +22,7 @@ from sglang.srt.configs.kimi_linear import KimiLinearConfig
from sglang.srt.distributed import (
divide,
get_pp_group,
get_shared_experts_tp_group,
get_tp_group,
tensor_model_parallel_all_reduce,
)
@@ -40,8 +41,6 @@ from sglang.srt.layers.activation import SiluAndMul, SituAndMul
from sglang.srt.layers.attn_residual import AttnResidual, aggregate_stream, get_cw
from sglang.srt.layers.dcp.planner import prepare_decode_context_parallel_metadata
from sglang.srt.layers.dp_attention import (
attn_tp_all_gather_into_tensor,
attn_tp_reduce_scatter_tensor,
dp_gather_replicate,
dp_scatter,
get_global_dp_buffer,
@@ -76,6 +75,7 @@ from sglang.srt.layers.moe.utils import (
)
from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.layers.quantization.fp8_utils import block_quant_dequant
from sglang.srt.layers.quantization.modelslim.modelslim import ModelSlimConfig
from sglang.srt.layers.radix_linear_attention import RadixLinearAttention
from sglang.srt.layers.utils import PPMissingLayer, get_layer_id
from sglang.srt.layers.vocab_parallel_embedding import (
@@ -453,14 +453,20 @@ class KimiK3MoE(nn.Module):
# full precision (matches GateLinear in mke). codespell:ignore mke
self.gate = MoEGate(config, quant_config=None, prefix=f"{prefix}.gate")
# For MXFP4 compressed-tensors, replace quant_config with Mxfp4Config
# so FusedMoE's weight_loader uses the MXFP4 fast path
# For MXFP4 compressed-tensors on non-NPU, replace quant_config with
# Mxfp4Config so FusedMoE's weight_loader uses the MXFP4 fast path. On
# NPU the compressed-tensors config is kept so the scheme-based path
# selects NPUCompressedTensorsW4A8mxfp4MoE (see get_moe_scheme).
moe_quant_config = quant_config
if quant_config is not None and getattr(quant_config, "quant_format", None):
if "mxfp4" in quant_config.quant_format:
from sglang.srt.layers.quantization.mxfp4 import Mxfp4Config
if (
quant_config is not None
and getattr(quant_config, "quant_format", None)
and "mxfp4" in quant_config.quant_format
and not _is_npu
):
from sglang.srt.layers.quantization.mxfp4 import Mxfp4Config
moe_quant_config = Mxfp4Config(is_checkpoint_mxfp4_serialized=True)
moe_quant_config = Mxfp4Config(is_checkpoint_mxfp4_serialized=True)
# Routed experts (operate in moe_hidden_size space)
# gate_up_interleaved=False: K3 loads per-expert w1/w3 into non-interleaved layout
@@ -561,32 +567,42 @@ class KimiK3MoE(nn.Module):
and config.hidden_act == "situ"
)
# Shared experts (operate in original hidden_size space).
# Replicate the shared-expert weights (tp1, DSv2 convention) under EP
# 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 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 = (
get_parallel().enable_shared_experts_attn_tp
and self._ep_a2a
and get_parallel().attn_tp_size > 1
# Shared experts operate on original hidden states. EP a2a gives each
# rank a token shard: either replicate the weights, or gather within
# the shared-expert TP subgroup and reduce-scatter back to those rows.
parallel = get_parallel()
requested_shared_tp = parallel.shared_experts_tp_size
shared_tp = requested_shared_tp
if shared_tp is None and parallel.enable_shared_experts_attn_tp:
shared_tp = parallel.attn_tp_size
if requested_shared_tp is not None and not self._ep_a2a:
raise ValueError("Independent shared-expert TP requires an EP a2a backend.")
self._shared_experts_tp1 = self._ep_a2a and shared_tp in (None, 1)
self._shared_experts_tp_comm = (
self._ep_a2a and shared_tp is not None and shared_tp > 1
)
self._shared_experts_tp_group = None
shared_experts_tp_kwargs = {}
if self._shared_experts_tp1:
shared_experts_tp_kwargs = dict(tp_rank=0, tp_size=1)
elif self._shared_experts_attn_tp_comm:
elif self._shared_experts_tp_comm:
group = (
get_shared_experts_tp_group()
if requested_shared_tp is not None
else parallel.attn_tp_group
)
assert group.world_size == shared_tp
self._shared_experts_tp_group = group
shared_experts_tp_kwargs = dict(
tp_rank=get_parallel().attn_tp_rank,
tp_size=get_parallel().attn_tp_size,
tp_rank=group.rank_in_group, tp_size=group.world_size
)
if self.num_shared_experts is not None and self.num_shared_experts > 0:
shared_intermediate_size = moe_intermediate_size * self.num_shared_experts
if shared_tp is not None and shared_intermediate_size % shared_tp != 0:
raise ValueError(
f"Shared-expert intermediate size ({shared_intermediate_size}) "
f"must be divisible by shared-expert TP size ({shared_tp})."
)
self.shared_experts = KimiK3MLP(
hidden_size=config.hidden_size,
intermediate_size=shared_intermediate_size,
@@ -611,12 +627,11 @@ class KimiK3MoE(nn.Module):
# (TP8/EP8 MegaMoE + SP-MoE): +4~5% output tok/s and 5% ITL over
# bs 132, GSM8K unchanged — so it is on whenever the shape allows,
# no flag.
# EP a2a only: with plain-TP experts the fused front already lands both
# partial sums in one collective (_forward_fused), a strictly better
# overlap than two streams.
# NPU shared-expert TP can also overlap the shared
# collectives using SGLANG_NPU_FINE_GRAINED_MOE_DUAL_STREAM. Otherwise
# the collectives stay on the current stream.
self._sbo_shared_overlap = (
self._ep_a2a
and not self._shared_experts_attn_tp_comm
and self.shared_experts is not None
and self.alt_stream is not None
)
@@ -1000,20 +1015,62 @@ class KimiK3MoE(nn.Module):
return self._latent_norm(latent)
return self._latent_norm(tensor_model_parallel_all_reduce(latent))
def _gather_shared_expert_inputs(self, hidden_states: torch.Tensor) -> torch.Tensor:
group = self._shared_experts_tp_group
# The attention DP buffer spans the entire attention-TP replica.
# Size this buffer from the subgroup's actual token shards instead.
with use_symmetric_memory(group, disabled=not is_allocation_symmetric()):
gathered = hidden_states.new_empty(
(hidden_states.shape[0] * group.world_size, *hidden_states.shape[1:])
)
group.all_gather_into_tensor(gathered, hidden_states)
return gathered
def _reduce_scatter_shared_experts(
self, shared_output: torch.Tensor, hidden_states: torch.Tensor
) -> torch.Tensor:
output = torch.empty_like(hidden_states)
self._shared_experts_tp_group.reduce_scatter_tensor(output, shared_output)
return output
def _forward_shared_experts(self, hidden_states: torch.Tensor) -> torch.Tensor:
"""Run TP-sharded shared experts while DeepEP tokens stay scattered."""
if not self._shared_experts_attn_tp_comm:
if not self._shared_experts_tp_comm:
return self.shared_experts(hidden_states)
group = get_parallel().attn_tp_group
# SP-MoE presents one contiguous token shard per attention-TP rank;
# the DP local buffer is the full reassembled per-replica batch.
gathered_hidden_states = get_local_dp_buffer(group)
attn_tp_all_gather_into_tensor(gathered_hidden_states, hidden_states)
gathered_hidden_states = self._gather_shared_expert_inputs(hidden_states)
gathered_shared_output = self.shared_experts(gathered_hidden_states)
shared_output = torch.empty_like(hidden_states)
attn_tp_reduce_scatter_tensor(shared_output, gathered_shared_output)
return shared_output
return self._reduce_scatter_shared_experts(
gathered_shared_output, hidden_states
)
def _can_overlap_shared_experts_npu(self, hidden_states: torch.Tensor) -> bool:
if not (
_is_npu
and envs.SGLANG_NPU_FINE_GRAINED_MOE_DUAL_STREAM.get()
and self._sbo_shared_overlap
and self._shared_experts_tp_comm
and self.use_latent_moe
and hidden_states.shape[0] > 0
and get_moe_a2a_backend().is_deepep()
):
return False
from sglang.srt.batch_overlap.two_batch_overlap import (
MaybeTboDeepEPDispatcher,
)
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
is_in_tc_piecewise_cuda_graph,
)
# The hooks must surround the complete dispatch, including its receive
# wait. Fused EP bypasses these hooks. An eager/piecewise graph break
# must not split the side-stream event record from its wait.
return (
isinstance(self.experts.dispatcher, MaybeTboDeepEPDispatcher)
and not is_in_breakable_cuda_graph()
and not is_in_tc_piecewise_cuda_graph()
)
def _forward_unfused(
self,
@@ -1026,33 +1083,117 @@ class KimiK3MoE(nn.Module):
# Shared experts on original hidden_states. Under SBO they go to the
# side stream and are joined at the tail (see _sbo_shared_overlap).
#
# Issued *after* the front, deliberately: alt_stream.wait_stream() makes
# the side stream wait for whatever the main stream has enqueued so far,
# so issuing here means the shared experts overlap the routed a2a rather
# than the front GEMMs. The shared branch is the shorter of the two and
# does not need a head start; running it against the front only takes
# bandwidth away from the critical path.
# CUDA issues this after the front so the shared experts overlap the
# routed a2a rather than the front GEMMs. NPU starts the shared branch
# before the front. Fine-grained NPU overlap splits it at the complete
# dispatch boundaries:
# current: front ---------- dispatch ---------- routed GEMMs -- tail
# alt: all-gather ----- shared MLP -------- reduce-scatter
# Shared and routed GEMMs wait for each other at phase boundaries;
# each can run beside the other branch's communication.
fine_grained_overlap = self._can_overlap_shared_experts_npu(hidden_states)
shared_input = None
shared_output = None
shared_event = None
shared_compute_event = None
def issue_shared():
nonlocal shared_output, shared_event
nonlocal shared_input, shared_output, shared_event
if self.shared_experts is None or hidden_states.shape[0] == 0:
return
if self._sbo_shared_overlap:
if fine_grained_overlap:
# Fork before the routed front so HCCL's completion wait is
# queued on the side stream, leaving the front free to run.
self.alt_stream.wait_stream(torch.cuda.current_stream())
hidden_states.record_stream(self.alt_stream)
with torch.cuda.stream(self.alt_stream):
shared_output = self._forward_shared_experts(hidden_states)
shared_input = self._gather_shared_expert_inputs(hidden_states)
shared_input.record_stream(self.alt_stream)
return
if self._sbo_shared_overlap:
current_stream = torch.cuda.current_stream()
# Keep HCCL collectives on the current stream. The alternate
# stream only executes the shared-expert MLP.
shared_input = hidden_states
if self._shared_experts_tp_comm:
shared_input = self._gather_shared_expert_inputs(hidden_states)
shared_input.record_stream(self.alt_stream)
self.alt_stream.wait_stream(current_stream)
with torch.cuda.stream(self.alt_stream):
shared_output = self.shared_experts(shared_input)
shared_event = self.alt_stream.record_event()
else:
shared_output = self._forward_shared_experts(hidden_states)
def run_experts(expert_input, topk_output):
if not fine_grained_overlap:
return (
self._forward_mega_experts(expert_input, topk_output)
if self._use_mega_moe
else self.experts(expert_input, topk_output)
)
def pre_dispatch(dispatcher, dispatch_input, dispatch_topk):
nonlocal shared_output, shared_compute_event
# AllGather is already queued. Delay shared GEMMs until the
# gate, TopK and latent down projection finish on current.
self.alt_stream.wait_stream(torch.cuda.current_stream())
with torch.cuda.stream(self.alt_stream):
shared_output = self.shared_experts(shared_input)
shared_compute_event = self.alt_stream.record_event()
def post_dispatch(dispatcher, dispatch_output):
nonlocal shared_output, shared_event
current_stream = torch.cuda.current_stream()
# Dispatch has queued its receive wait. RS waits for that
# communication and the shared MLP, while routed GEMMs wait
# only for the MLP (not for RS).
self.alt_stream.wait_stream(current_stream)
with torch.cuda.stream(self.alt_stream):
shared_output = self._reduce_scatter_shared_experts(
shared_output, hidden_states
)
shared_event = self.alt_stream.record_event()
current_stream.wait_event(shared_compute_event)
dispatcher = self.experts.dispatcher
pre_handle = dispatcher.register_pre_dispatch_hook(pre_dispatch)
try:
post_handle = dispatcher.register_post_dispatch_hook(post_dispatch)
try:
return self.experts(expert_input, topk_output)
finally:
post_handle.remove()
finally:
# Remove outside hook iteration, including on dispatch/GEMM
# failures, so closures cannot leak into the next forward.
pre_handle.remove()
def wait_and_finalize_shared_experts():
nonlocal shared_output
if shared_event is None:
return
# Join just before consuming the shared result. The legacy path
# still needs to reduce-scatter its TP-partial MLP output here.
current_stream = torch.cuda.current_stream()
current_stream.wait_event(shared_event)
shared_output.record_stream(current_stream)
if self._shared_experts_tp_comm and not fine_grained_overlap:
shared_output = self._reduce_scatter_shared_experts(
shared_output, hidden_states
)
# Give the NPU shared-expert branch a head start. At this point
# hidden_states is the decoder layer's post-attention RMSNorm output.
if _is_npu and self._sbo_shared_overlap:
issue_shared()
# Front: gate + TopK (+ latent down-proj when the merged front covers it).
# The gate and the latent down-proj read the same hidden_states, so the
# merged-weight strategies compute both in one GEMM; see
# kernels/ops/moe/moe_front.py for the strategy table.
routed_input = self._ep_front(hidden_states)
if routed_input is None:
if routed_input is None and not fine_grained_overlap:
routed_input = self._ep_front_overlap(hidden_states)
topk_output = None
if routed_input is not None:
@@ -1063,15 +1204,17 @@ class KimiK3MoE(nn.Module):
# fp32 logits reach the radix router from moe_fused_gate.
router_logits = self.gate(hidden_states)
topk_output = self.topk(hidden_states, router_logits)
issue_shared()
if not (_is_npu and self._sbo_shared_overlap):
issue_shared()
if not self.use_latent_moe:
expert_output = self.experts(hidden_states, topk_output)
if shared_event is not None:
torch.cuda.current_stream().wait_event(shared_event)
wait_and_finalize_shared_experts()
if shared_output is not None:
expert_output = expert_output + shared_output
if self.tp_size > 1:
# EP combine and the shared-expert subgroup have already completed
# each source token. A global TP reduction would mix token shards.
if self.tp_size > 1 and not self._ep_a2a:
expert_output = tensor_model_parallel_all_reduce(expert_output)
if prefix_sum is not None:
expert_output = expert_output + prefix_sum
@@ -1093,11 +1236,7 @@ class KimiK3MoE(nn.Module):
routed_input = hidden_states.new_empty((0, self.moe_hidden_size))
else:
routed_input, _ = self.routed_expert_down_proj(hidden_states)
expert_output = (
self._forward_mega_experts(routed_input, topk_output)
if self._use_mega_moe
else self.experts(routed_input, topk_output)
)
expert_output = run_experts(routed_input, topk_output)
if expert_output.shape[0] == 0:
# The EP combine returns one row per source token. Keep the
# source-side empty result while avoiding empty RMSNorm/up-proj
@@ -1107,17 +1246,14 @@ class KimiK3MoE(nn.Module):
latent = self._reduce_latent(expert_output)
# up_proj is replicated, so the routed output is now fully reduced.
out, _ = self.routed_expert_up_proj(latent)
if shared_event is not None:
# SBO join: as late as possible, so the side-stream shared experts
# get the whole routed a2a + latent tail to hide under.
torch.cuda.current_stream().wait_event(shared_event)
wait_and_finalize_shared_experts()
if shared_output is not None:
# tp1 shared experts (SP-MoE) are complete per-rank; TP-sharded
# ones need the partial-sum reduction.
if (
self.tp_size > 1
and not self._shared_experts_tp1
and not self._shared_experts_attn_tp_comm
and not self._shared_experts_tp_comm
):
shared_output = tensor_model_parallel_all_reduce(shared_output)
out = _add3(out, shared_output, prefix_sum)
@@ -1441,13 +1577,15 @@ class KimiK3DeltaAttention(nn.Module):
quant_config, f"{prefix}.b_proj"
)
# The fused path hardcodes tp_size sharding, so require attn_tp == tp.
# Full-rank K3 also fuses mixed block-FP8 attention projections.
self.do_fuse_qkvbfg = self.attn_tp_size == self.tp_size and (
quant_config is None or self.use_full_rank_gate
)
# The full-rank [q, k, v, g] merged projection is explicitly sharded
# with attn_tp_rank/attn_tp_size, so it also supports DP attention.
# The low-rank fused path still uses full-TP-only projection helpers.
# For the full-rank gate (K3) the checkpoint quantizes only the MoE
# experts; attention linears resolve to UnquantizedLinearMethod, so a
# non-None quant_config is fine for the merged projection.
self.do_fuse_qkvbfg = quant_config is None and self.attn_tp_size == self.tp_size
if self.do_fuse_qkvbfg and self.use_full_rank_gate:
if self.use_full_rank_gate:
# Fuse only the alignment-friendly wide projections [q, k, v, g]
# (6144/rank at TP8). Folding b (12/rank) and f_a (128, replicated)
# in as well skews the output dim to 6284 and measurably degrades
@@ -1469,8 +1607,8 @@ class KimiK3DeltaAttention(nn.Module):
prefix=f"{prefix}.fused_qkvg_proj",
)
self.split_sizes = [
3 * projection_size // self.tp_size,
projection_size // self.tp_size,
3 * projection_size // self.attn_tp_size,
projection_size // self.attn_tp_size,
]
self.b_proj = ColumnParallelLinear(
self.hidden_size,
@@ -1998,7 +2136,7 @@ class KimiK3DeltaAttention(nn.Module):
defer_f_b = (
self._kda_hip_fused_decode_ready and forward_batch.forward_mode.is_decode()
)
if self.do_fuse_qkvbfg:
if self.do_fuse_qkvbfg or self.use_full_rank_gate:
mixed_qkv, beta, forget_gate, g_proj_states = self.forward_qkvbfg_fused(
hidden_states, defer_f_b=defer_f_b
)
@@ -3022,6 +3160,13 @@ class KimiK3LinearModel(nn.Module):
class KimiK3LinearForCausalLM(nn.Module):
"""Text-only K3 causal LM."""
# ModelSlim describes quantization with the original checkpoint module
# names. Register the runtime fused QKVG module so it can resolve the
# q_proj scheme while the weight loader packs q/k/v/g into its shards.
packed_modules_mapping = {
"fused_qkvg_proj": ["q_proj", "k_proj", "v_proj", "g_proj"],
}
def __init__(
self,
config: KimiLinearConfig,
@@ -3031,6 +3176,15 @@ class KimiK3LinearForCausalLM(nn.Module):
super().__init__()
self.config = config
self.quant_config = quant_config
if quant_config is not None:
if isinstance(quant_config, ModelSlimConfig):
model_mapping = {
**quant_config.packed_modules_mapping.get("model", {}),
**self.packed_modules_mapping,
}
quant_config.update_packed_modules_mapping({"model": model_mapping})
else:
quant_config.update_packed_modules_mapping(self.packed_modules_mapping)
self.model = KimiK3LinearModel(
config, quant_config, prefix=maybe_prefix(prefix, "model")
)
@@ -3196,7 +3350,8 @@ class KimiK3LinearForCausalLM(nn.Module):
continue
# compressed-tensors MXFP4 stores as weight_packed; Mxfp4MoEMethod uses weight
if "weight_packed" in name:
# (NPU keeps weight_packed for NPUCompressedTensorsW4A8mxfp4MoE).
if "weight_packed" in name and not _is_npu:
name = name.replace("weight_packed", "weight")
# MLA: fuse q_a_proj + kv_a_proj_with_mqa → fused_qkv_a_proj_with_mqa
@@ -3242,7 +3397,13 @@ class KimiK3LinearForCausalLM(nn.Module):
if not self.config.is_kda_layer(layer_id):
continue
layer = self.model.layers[layer_id].self_attn
if not getattr(layer, "do_fuse_qkvbfg", False):
# Full-rank K3 always instantiates fused_qkvg_proj, including
# ModelSlim-quantized models. The low-rank fused modules are
# still conditional on do_fuse_qkvbfg.
if param_name == ".fused_qkvg_proj":
if not getattr(layer, "use_full_rank_gate", False):
continue
elif not getattr(layer, "do_fuse_qkvbfg", False):
continue
if weight_name in {".q_proj", ".k_proj", ".v_proj"}:
layer_id = int(name.split(".")[2])
@@ -913,9 +913,32 @@ class DSparkWorkerV2(BaseSpecWorker):
last_correct_step_indices = commit_lens.to(torch.int64) - 1
mamba_steps_to_track = None
mamba_track_indices = batch.mamba_track_indices
if batch.mamba_track_indices is not None:
if mamba_track_indices is not None:
mamba_track_interval = mamba_track_grid(batch.tree_cache.page_size)
seq_lens_cpu = batch.seq_lens_cpu
if (
_is_npu
and seq_lens_cpu is not None
and seq_lens_cpu.device.type == "cpu"
and seq_lens_cpu.ndim == 1
and seq_lens_cpu.numel() == seq_lens_pre_verify.numel()
and seq_lens_cpu.dtype in (torch.int32, torch.int64)
):
# Verify restores the CPU prefix lengths before the forward.
# Acceptance can commit at most this many tokens, so this
# check needs no device readback. Passing None also avoids
# the NPU backend's conv-state self-copy for untracked rows.
if all(
seq_len >= 0
and seq_len // mamba_track_interval
== (seq_len + self.verify_num_draft_tokens) // mamba_track_interval
for seq_len in seq_lens_cpu.tolist()
):
mamba_track_indices = None
if mamba_track_indices is not None:
to_track_mask = (
seq_lens_pre_verify // mamba_track_interval
!= seq_lens_post_verify // mamba_track_interval
@@ -935,7 +958,7 @@ class DSparkWorkerV2(BaseSpecWorker):
attn_backend.update_mamba_state_after_mtp_verify(
last_correct_step_indices=last_correct_step_indices,
mamba_track_indices=batch.mamba_track_indices,
mamba_track_indices=mamba_track_indices,
mamba_steps_to_track=mamba_steps_to_track,
model=self.target_worker.model_runner.model,
req_pool_indices=batch.req_pool_indices,