Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion docs/docs/advanced_features/server_arguments.mdx
Original file line number Diff line number Diff line change
Expand Up @@ -1626,7 +1626,7 @@ Combining `--enable-response-store` with `--disaggregation-mode=prefill` or `dec
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>`--speculative-moe-a2a-backend`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>MOE A2A backend for EAGLE speculative decoding, see `--moe-a2a-backend` for options. Same as moe a2a backend if unset.</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>MoE A2A backend for speculative decoding; see `--moe-a2a-backend` for options. If unset or `none`, inherit the target model's backend. Explicit `none` does not select a separate local dispatcher when the target uses another backend.</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>`None`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>none</code>, <code>deepep</code>, <code>mooncake</code>, <code>nixl</code>, <code>mori</code>, <code>ascend_fuseep</code>, <code>flashinfer</code>, <code>megamoe</code>, <code>pplx</code></td>
</tr>
Expand Down
2 changes: 1 addition & 1 deletion docs/docs/advanced_features/speculative_decoding.mdx
Original file line number Diff line number Diff line change
Expand Up @@ -946,7 +946,7 @@ Below is a comprehensive list of all speculative decoding parameters available i
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>--speculative-moe-a2a-backend</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>str</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>None</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>MoE all-to-all backend for the draft model</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>MoE all-to-all backend for the draft model. If unset or <code>none</code>, inherit <code>--moe-a2a-backend</code>; other values select an explicit draft backend.</td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>--speculative-draft-model-quantization</code></td>
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1476,7 +1476,7 @@ non-default speculative acceptance thresholds or deterministic inference.
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>`--speculative-moe-a2a-backend`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>`None`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>`ascend_fuseep` (the only supported value on Ascend NPU)</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>`none` (inherit the target backend), `deepep`, `ascend_fuseep`; model and quantization support varies by backend.</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>A2/A3 Series</td>
</tr>
<tr>
Expand Down
5 changes: 4 additions & 1 deletion python/sglang/srt/arg_groups/fields/spec.py
Original file line number Diff line number Diff line change
Expand Up @@ -170,7 +170,10 @@ class Spec(msgspec.Struct):
speculative_moe_a2a_backend: A[
Optional[str],
Arg(
help="Choose the backend for MoE A2A in speculative decoding",
help=(
"Choose the backend for MoE A2A in speculative decoding. "
"If unset or 'none', inherit --moe-a2a-backend."
),
choices=[
"none",
"deepep",
Expand Down
2 changes: 1 addition & 1 deletion python/sglang/srt/arg_groups/moe_hook.py
Original file line number Diff line number Diff line change
Expand Up @@ -527,7 +527,7 @@ def validate_deepep_v2_speculative_draft(server_args: Any) -> None:
"""Reject an explicit or inherited DeepEP v2 draft backend."""
view = resolved_view(server_args)
draft_backend = view.speculative_moe_a2a_backend
if draft_backend is None and view.speculative_algorithm:
if draft_backend in (None, "none") and view.speculative_algorithm:
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm

algorithm = SpeculativeAlgorithm.from_string(view.speculative_algorithm)
Expand Down
2 changes: 1 addition & 1 deletion python/sglang/srt/arg_groups/speculative_hook.py
Original file line number Diff line number Diff line change
Expand Up @@ -608,7 +608,7 @@ def _handle_dspark(server_args: ServerArgs) -> None:
)
if (
not _is_npu
and cfg.speculative_moe_a2a_backend is not None
and cfg.speculative_moe_a2a_backend not in (None, "none")
and cfg.speculative_moe_a2a_backend != cfg.moe_a2a_backend
):
raise ValueError(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,10 @@ def _finalize_routing(
expanded_row_idx: torch.Tensor,
topk_ids: torch.Tensor,
) -> torch.Tensor:
if self.drop_pad_mode == 3 and hidden_states.ndim == 2:
# Drop mode requires [E, C, H], but uses flattened row indices.
# View the packed GMM rows as one capacity buffer without copying.
hidden_states = hidden_states.unsqueeze(0)
return torch.ops.npu.npu_moe_finalize_routing(
hidden_states,
skip1=None,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -96,6 +96,7 @@ def _init_routing(
topk_ids: torch.Tensor,
num_experts: int,
top_k: int,
active_expert_range: Optional[Tuple[int, int]] = None,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, Optional[torch.Tensor]]:
num_tokens = hidden_states.shape[0]
hidden_states, expanded_row_idx, expert_tokens, pertoken_scale = (
Expand All @@ -106,7 +107,11 @@ def _init_routing(
expert_num=num_experts,
expert_tokens_num_type=1,
expert_tokens_num_flag=True,
active_expert_range=[0, num_experts],
active_expert_range=(
list(active_expert_range)
if active_expert_range is not None
else [0, num_experts]
),
quant_mode=self.quant_mode,
)
)
Expand Down
49 changes: 46 additions & 3 deletions python/sglang/srt/layers/moe/token_dispatcher/ascend_tp.py
Original file line number Diff line number Diff line change
Expand Up @@ -52,6 +52,14 @@ class AscendTPDispatcher(BaseDispatcher):
def __init__(self, moe_runner_config: MoeRunnerConfig):
super().__init__()
self.num_experts = moe_runner_config.num_experts
self.num_local_experts = moe_runner_config.num_local_experts
self.num_local_shared_experts = moe_runner_config.num_fused_shared_experts
self.num_local_routed_experts = (
self.num_local_experts - self.num_local_shared_experts
)
self.moe_ep_size = get_parallel().moe_ep_size
self.moe_ep_rank = get_parallel().moe_ep_rank
self.local_expert_mapping = None
self.top_k = moe_runner_config.top_k
self._dispatch_output: Optional[AscendTPDispatchOutput] = None

Expand All @@ -76,18 +84,21 @@ def set_quant_config(self, quant_config: dict) -> None:
def set_ascend_dispatcher_output_dtype(self) -> None:
"""Choose init & finalize routing kernels based on quant config."""
self.ascend_dispatcher_output_dtype = get_ascend_dispatcher_output_dtype(self)
# EP filtering produces -1 row indices for nonlocal experts. Mode 3
# supports dropped routes in row-major order; mode 2 does not.
drop_pad_mode = 3 if self.moe_ep_size > 1 else 2

if self.ascend_dispatcher_output_dtype == DispatcherOutputDtype.BF16:
self.init = NPUMoEInitRouting_v2(quant_mode=-1)
self.finalize = NPUFinalizeRouting(drop_pad_mode=2)
self.finalize = NPUFinalizeRouting(drop_pad_mode=drop_pad_mode)
self.group_list_type = 1
elif self.ascend_dispatcher_output_dtype == DispatcherOutputDtype.INT8:
self.init = NPUMoEInitRouting_v2(quant_mode=1)
self.finalize = NPUFinalizeRouting(drop_pad_mode=2)
self.finalize = NPUFinalizeRouting(drop_pad_mode=drop_pad_mode)
self.group_list_type = 1
elif self.ascend_dispatcher_output_dtype == DispatcherOutputDtype.MXFP8:
self.init = NPUMoEInitRouting_v2(quant_mode=MXFP8_QUANT_MODE)
self.finalize = NPUFinalizeRouting(drop_pad_mode=2)
self.finalize = NPUFinalizeRouting(drop_pad_mode=drop_pad_mode)
self.group_list_type = 1
else:
raise ValueError(
Expand All @@ -102,6 +113,37 @@ def dispatch(
topk_ids = topk_ids.to(torch.int32)
top_k = topk_weights.shape[-1]

if self.moe_ep_size > 1:
if self.local_expert_mapping is None:
# Unlike the GPU dispatcher, use a nonnegative drop sentinel:
# init_routing_v2 filters IDs outside active_expert_range.
# The extra entry also maps padding IDs (-1) to that sentinel.
self.local_expert_mapping = torch.full(
(self.num_experts + 1,),
self.num_local_experts,
dtype=torch.int32,
device=topk_ids.device,
)
start = self.moe_ep_rank * self.num_local_routed_experts
self.local_expert_mapping[
start : start + self.num_local_routed_experts
] = torch.arange(
self.num_local_routed_experts,
dtype=torch.int32,
device=topk_ids.device,
)
if self.num_local_shared_experts > 0:
self.local_expert_mapping[
self.num_experts
- self.num_local_shared_experts : self.num_experts
] = torch.arange(
self.num_local_routed_experts,
self.num_local_experts,
dtype=torch.int32,
device=topk_ids.device,
)
topk_ids = self.local_expert_mapping[topk_ids]

(
permuted_hidden_states,
expanded_row_idx,
Expand All @@ -112,6 +154,7 @@ def dispatch(
topk_ids,
self.num_experts,
top_k,
active_expert_range=(0, self.num_local_experts),
)

self._dispatch_output = AscendTPDispatchOutput(
Expand Down
2 changes: 1 addition & 1 deletion python/sglang/srt/layers/moe/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -460,7 +460,7 @@ def initialize_moe_config():
)
moe.speculative_a2a_backend = (
MoeA2ABackend(spec.speculative_moe_a2a_backend)
if spec.speculative_moe_a2a_backend is not None
if spec.speculative_moe_a2a_backend not in (None, "none")
else moe.a2a_backend
)
moe.deepep_mode = DeepEPMode(exec_moe.deepep_mode)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -288,9 +288,9 @@ def refresh_deep_gemm_layout_memory_budget(
get_spec().speculative_moe_runner_backend
or get_exec().moe.moe_runner_backend
)
moe_a2a_backend = (
get_spec().speculative_moe_a2a_backend or get_exec().moe.moe_a2a_backend
)
moe_a2a_backend = get_spec().speculative_moe_a2a_backend
if moe_a2a_backend in (None, "none"):
moe_a2a_backend = get_exec().moe.moe_a2a_backend
else:
moe_runner_backend = get_exec().moe.moe_runner_backend
moe_a2a_backend = get_exec().moe.moe_a2a_backend
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -72,9 +72,9 @@ def should_run_flashinfer_autotune(
get_spec().speculative_moe_runner_backend
or get_exec().moe.moe_runner_backend
)
a2a_backend_str = (
get_spec().speculative_moe_a2a_backend or get_exec().moe.moe_a2a_backend
)
a2a_backend_str = get_spec().speculative_moe_a2a_backend
if a2a_backend_str in (None, "none"):
a2a_backend_str = get_exec().moe.moe_a2a_backend
else:
backend_str = get_exec().moe.moe_runner_backend
a2a_backend_str = get_exec().moe.moe_a2a_backend
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,148 @@
"""Check local EP routing against a dense reference using real NPU kernels.

One NPU runs each logical EP rank sequentially; summing its partial outputs
models the all-reduce performed by the model after the ``none`` dispatcher.
"""

import unittest
from types import SimpleNamespace
from unittest.mock import patch

import torch

from sglang.srt.layers.moe.moe_runner.base import MoeRunnerConfig
from sglang.srt.layers.moe.token_dispatcher.ascend_tp import (
AscendTPCombineInput,
AscendTPDispatcher,
)
from sglang.srt.layers.moe.topk import StandardTopKOutput
from sglang.srt.utils import is_npu
from sglang.test.ci.ci_register import register_npu_ci
from sglang.test.test_utils import CustomTestCase

register_npu_ci(est_time=30, suite="full-1-npu-a3", nightly=True)


@unittest.skipUnless(is_npu(), "Requires Ascend NPU kernels")
class TestNpuNoneDispatcher(CustomTestCase):
def _check_routing(self, ids, num_routed=4, ep_size=2, num_shared=0, int8=False):
ids = torch.tensor(ids, dtype=torch.int32).reshape(-1, 2)
num_tokens, top_k = ids.shape
num_experts = num_routed + num_shared
local_routed = num_routed // ep_size
num_local = local_routed + num_shared
# Binary fractions make BF16 matmul/combine exact; distinct expert
# transforms detect a wrong rank offset even when count shapes match.
x = ((torch.arange(num_tokens * 32) % 13 - 6) / 8).reshape(num_tokens, 32)
weights = torch.stack(
[torch.eye(32).roll(expert % 32, dims=1) for expert in range(num_experts)]
)
scores = torch.tensor([0.25, 0.75]).repeat(num_tokens, 1)
expected = torch.zeros_like(x)
for row in range(num_tokens):
for slot in range(top_k):
expert = int(ids[row, slot])
if expert >= 0:
expected[row] += scores[row, slot] * (x[row] @ weights[expert])

total = torch.zeros_like(x)
for rank in range(ep_size):
with self.subTest(rank=rank):
config = MoeRunnerConfig(
num_experts=num_experts,
num_local_experts=num_local,
num_fused_shared_experts=num_shared,
top_k=top_k,
)
with patch(
"sglang.srt.layers.moe.token_dispatcher.ascend_tp.get_parallel",
return_value=SimpleNamespace(moe_ep_size=ep_size, moe_ep_rank=rank),
):
dispatcher = AscendTPDispatcher(config)
if int8:
dispatcher.set_quant_config({"dispatcher_output_dtype": "int8"})
local_ids = list(range(rank * local_routed, (rank + 1) * local_routed))
local_ids += list(range(num_routed, num_experts))
local_weights = weights[local_ids].clone()
# Shared experts are replicated across EP ranks. Their caller
# supplies the 1/EP scaling; the dispatcher must keep them local.
if num_shared:
local_weights[-num_shared:] /= ep_size
local_weights = local_weights.to(device="npu", dtype=torch.bfloat16)
topk = StandardTopKOutput(
topk_weights=scores.to(device="npu", dtype=torch.bfloat16),
topk_ids=ids.npu(),
router_logits=None,
)
hidden = x.to(device="npu", dtype=torch.bfloat16)
# Repeat to cover cached expert mappings and cleared combine state.
for _ in range(2):
dispatched = dispatcher.dispatch(hidden, topk)
counts = dispatched.expert_tokens.cpu()
self.assertEqual(counts.numel(), local_weights.shape[0])
torch.testing.assert_close(
counts,
torch.tensor([(ids == expert).sum() for expert in local_ids]),
)
permuted = dispatched.hidden_states
if int8:
self.assertEqual(permuted.dtype, torch.int8)
permuted = (
permuted.float()
* dispatched.hidden_states_scale.float().reshape(-1, 1)
).to(torch.bfloat16)
expert_output = torch.ops.npu.npu_grouped_matmul(
[permuted],
[local_weights],
group_list=dispatched.expert_tokens,
group_type=0,
split_item=3,
group_list_type=dispatched.group_list_type,
)[0]
# Unused GMM rows are undefined. Poison them so a finalizer
# that accidentally reads dropped routes cannot pass by luck.
expert_output[int(counts.sum()) :].fill_(float("nan"))
actual = dispatcher.combine(AscendTPCombineInput(expert_output))
partial = torch.zeros_like(x)
for row in range(num_tokens):
for slot in range(top_k):
expert = int(ids[row, slot])
if expert in local_ids:
scale = 1 / ep_size if expert >= num_routed else 1
partial[row] += (
scores[row, slot]
* scale
* (x[row] @ weights[expert])
)
torch.testing.assert_close(
actual.float().cpu(), partial, rtol=0, atol=0.008 if int8 else 0
)
total += actual.float().cpu()
torch.testing.assert_close(total, expected, rtol=0, atol=0.008 if int8 else 0)

def test_ep1(self):
self._check_routing([[0, 2], [1, 3]], ep_size=1)

def test_ep2(self):
self._check_routing([[0, 2], [1, 3], [2, 3], [0, 1]])

def test_ep2_qwen36_expert_count(self):
self._check_routing(
[[0, 127], [128, 255], [255, 1], [129, 126]], num_routed=256
)

def test_ep2_no_local_tokens(self):
self._check_routing([[2, 3]])

def test_ep2_shared_experts_and_padding(self):
self._check_routing([[0, 4], [2, 4], [-1, 4]], num_shared=1)

def test_ep2_int8_routing(self):
self._check_routing([[0, 2], [1, 3], [2, 3], [0, 1]], int8=True)

def test_ep2_empty_batch(self):
self._check_routing([])


if __name__ == "__main__":
unittest.main()
Loading