[Feature] Add DeepEPv2 (ElasticBuffer) MoE A2A backend (#35634)

Co-authored-by: menyu <menyu@nvidia.com>
Co-authored-by: Jinyan Chen <93358689+liz-badada@users.noreply.github.com>
Co-authored-by: Han Yu <helloyu0903@gmail.com>
Co-authored-by: Cheng Wan <cheng.wan@radixark.ai>
This commit is contained in:
MengYu
2026-08-26 19:54:33 -07:00
committed by GitHub
co-authored by menyu Jinyan Chen Han Yu Cheng Wan
parent cbfe54fba8
commit a3ae667d67
20 changed files with 2414 additions and 35 deletions
@@ -1221,6 +1221,142 @@ def ep_scatter(
return return
@triton.jit
def _fwd_kernel_ep_scatter_psum_init(
psum_num_recv_tokens_per_expert,
expert_start_loc,
m_indices,
BLOCK_E: tl.constexpr,
):
cur_expert = tl.program_id(0)
cur_end = tl.load(psum_num_recv_tokens_per_expert + cur_expert)
cur_start = tl.load(
psum_num_recv_tokens_per_expert + cur_expert - 1,
mask=cur_expert > 0,
other=0,
)
cur_token_num = cur_end - cur_start
tl.store(expert_start_loc + cur_expert, cur_start)
off_expert = tl.arange(0, BLOCK_E)
for start_m in tl.range(0, cur_token_num, BLOCK_E, num_stages=4):
# Mask the tail because this expert is packed against the next one.
idx = cur_start + start_m + off_expert
tl.store(m_indices + idx, cur_expert, mask=idx < cur_end)
@torch.no_grad()
def ep_scatter_from_psum(
recv_x: torch.Tensor,
recv_x_scale: torch.Tensor,
recv_topk: torch.Tensor,
psum_num_recv_tokens_per_expert: torch.Tensor,
expert_start_loc: torch.Tensor,
output_tensor: torch.Tensor,
output_tensor_scale: torch.Tensor,
m_indices: torch.Tensor,
output_index: torch.Tensor,
scale_ue8m0: bool = False,
):
BLOCK_E = 128
BLOCK_D = 128
num_warps = 8
num_experts = psum_num_recv_tokens_per_expert.shape[0]
hidden_size = recv_x.shape[1]
scale_hidden_size = hidden_size // BLOCK_D
if scale_ue8m0:
scale_hidden_size = ceil_div(scale_hidden_size, 4)
assert m_indices.shape[0] % BLOCK_E == 0
is_fp8 = recv_x_scale is not None and recv_x.dtype != torch.bfloat16
if is_fp8:
assert recv_x_scale.dtype == output_tensor_scale.dtype
assert (
recv_x_scale.shape[1] == output_tensor_scale.shape[1] == scale_hidden_size
)
_fwd_kernel_ep_scatter_psum_init[(num_experts,)](
psum_num_recv_tokens_per_expert,
expert_start_loc,
m_indices,
num_warps=num_warps,
BLOCK_E=BLOCK_E,
)
grid = min(recv_topk.shape[0], 1024 * 8)
_fwd_kernel_ep_scatter_2[(grid,)](
recv_topk.shape[0],
expert_start_loc,
recv_x,
recv_x.stride(0),
recv_x.stride(1),
recv_x_scale,
recv_x_scale.stride(0) if is_fp8 else 0,
recv_x_scale.stride(1) if is_fp8 else 0,
recv_topk,
recv_topk.stride(0),
recv_topk.stride(1),
output_tensor,
output_tensor.stride(0),
output_tensor.stride(1),
output_tensor_scale,
output_tensor_scale.stride(0) if is_fp8 else 0,
output_tensor_scale.stride(1) if is_fp8 else 0,
output_index,
output_index.stride(0),
output_index.stride(1),
topk_num=recv_topk.shape[1],
num_warps=num_warps,
HIDDEN_SIZE=hidden_size,
HIDDEN_SIZE_PAD=triton.next_power_of_2(hidden_size),
SCALE_HIDDEN_SIZE=scale_hidden_size,
SCALE_HIDDEN_SIZE_PAD=triton.next_power_of_2(scale_hidden_size),
ATOMIC_ADD_SEM=None if not _is_musa else "relaxed",
IS_FP8=is_fp8,
)
return
@triton.jit
def _fwd_kernel_ep_expand_m_indices_init(
psum_num_recv_tokens_per_expert,
m_indices,
BLOCK_E: tl.constexpr,
):
cur_expert = tl.program_id(0)
cur_end = tl.load(psum_num_recv_tokens_per_expert + cur_expert)
prev_end = tl.load(
psum_num_recv_tokens_per_expert + cur_expert - 1,
mask=cur_expert > 0,
other=0,
)
cur_start = ((prev_end + BLOCK_E - 1) // BLOCK_E) * BLOCK_E
aligned_end = ((cur_end + BLOCK_E - 1) // BLOCK_E) * BLOCK_E
off_expert = tl.arange(0, BLOCK_E)
for start_m in tl.range(0, aligned_end - cur_start, BLOCK_E, num_stages=4):
idx = cur_start + start_m + off_expert
tl.store(m_indices + idx, cur_expert, mask=idx < aligned_end)
@torch.no_grad()
def ep_expand_init_m_indices_from_psum(
psum_num_recv_tokens_per_expert: torch.Tensor,
m_indices: torch.Tensor,
):
BLOCK_E = 128
num_warps = 8
num_experts = psum_num_recv_tokens_per_expert.shape[0]
assert m_indices.shape[0] % BLOCK_E == 0
_fwd_kernel_ep_expand_m_indices_init[(num_experts,)](
psum_num_recv_tokens_per_expert,
m_indices,
num_warps=num_warps,
BLOCK_E=BLOCK_E,
)
return
@triton.jit @triton.jit
def _fwd_kernel_ep_gather( def _fwd_kernel_ep_gather(
total_token_num, total_token_num,
@@ -2021,6 +2157,233 @@ def fp8_per_token_to_per_tensor_quant_triton(
) )
# Expanded psum starts each expert at align(psum[e-1]); contiguous psum already
# includes the alignment padding.
_DEEPEP_V2_REPACK_WORKERS_PER_EXPERT = 64
# Scale strides support row-major FP32 and packed column-major UE8M0.
@triton.jit
def _fwd_kernel_expand_to_masked_slab(
psum_ptr,
recv_x_ptr,
recv_x_stride0,
recv_x_scale_ptr,
recv_x_scale_stride0,
recv_x_scale_stride1,
output_tensor_ptr,
output_tensor_stride0,
output_tensor_scale_ptr,
masked_m_ptr,
overflow_ptr,
MAX_M: tl.constexpr,
ALIGN: tl.constexpr,
HIDDEN: tl.constexpr,
HIDDEN_PAD: tl.constexpr,
SCALE_HIDDEN: tl.constexpr,
SCALE_HIDDEN_PAD: tl.constexpr,
IS_FP8: tl.constexpr,
CHECK_OVERFLOW: tl.constexpr,
NUM_WORKERS: tl.constexpr,
):
# A fixed worker grid makes conservative max_m values graph-safe.
e = tl.program_id(0)
worker = tl.program_id(1)
prev_end = tl.load(psum_ptr + e - 1, mask=e > 0, other=0)
start = ((prev_end + ALIGN - 1) // ALIGN) * ALIGN
end = tl.load(psum_ptr + e)
raw_count = end - start
count = tl.minimum(raw_count, MAX_M)
if worker == 0:
tl.store(masked_m_ptr + e, count)
if CHECK_OVERFLOW:
# Graph capture omits this host-visible overflow flag.
ovf = tl.arange(0, 1)
tl.store(overflow_ptr + ovf, 1, mask=raw_count > MAX_M)
off = tl.arange(0, HIDDEN_PAD)
mask = off < HIDDEN
off_s = tl.arange(0, SCALE_HIDDEN_PAD)
mask_s = off_s < SCALE_HIDDEN
for j in tl.range(worker, count, NUM_WORKERS):
src = (start + j).to(tl.int64)
dst = (e * MAX_M + j).to(tl.int64)
v = tl.load(recv_x_ptr + src * recv_x_stride0 + off, mask=mask)
tl.store(output_tensor_ptr + dst * output_tensor_stride0 + off, v, mask=mask)
if IS_FP8:
vs = tl.load(
recv_x_scale_ptr
+ src * recv_x_scale_stride0
+ off_s * recv_x_scale_stride1,
mask=mask_s,
)
# Write physical [E, SCALE_HIDDEN, MAX_M] for an mn-major view.
tl.store(
output_tensor_scale_ptr + e * SCALE_HIDDEN * MAX_M + off_s * MAX_M + j,
vs,
mask=mask_s,
)
@torch.no_grad()
def expand_to_masked_slab(
recv_x: torch.Tensor,
recv_x_scale,
psum_num_recv_tokens_per_expert: torch.Tensor,
num_local_experts: int,
max_m: int,
expert_alignment: int,
):
"""expanded [total, hidden] -> ([E_local, max_m, hidden], [E_local, max_m, sh] or None, masked_m[E_local])."""
hidden = recv_x.shape[1]
is_fp8 = recv_x_scale is not None and recv_x.dtype != torch.bfloat16
output_tensor = torch.empty(
(num_local_experts * max_m, hidden), device=recv_x.device, dtype=recv_x.dtype
)
masked_m = torch.empty(
(num_local_experts,), device=recv_x.device, dtype=torch.int32
)
check_overflow = not torch.cuda.is_current_stream_capturing()
# Dummy pointer under capture; CHECK_OVERFLOW compiles out every store to it.
overflow = (
torch.zeros((1,), device=recv_x.device, dtype=torch.int32)
if check_overflow
else masked_m
)
if is_fp8:
sh = recv_x_scale.shape[1]
# Store [E, sh, max_m] so the returned transpose is mn-major.
output_tensor_scale = torch.empty(
(num_local_experts * sh, max_m),
device=recv_x.device,
dtype=recv_x_scale.dtype,
)
scale_arg = recv_x_scale
scale_s0 = recv_x_scale.stride(0)
scale_s1 = recv_x_scale.stride(1)
else:
sh = 1
output_tensor_scale = None
scale_arg = recv_x
scale_s0 = 0
scale_s1 = 0
num_workers = min(max_m, _DEEPEP_V2_REPACK_WORKERS_PER_EXPERT)
_fwd_kernel_expand_to_masked_slab[(num_local_experts, num_workers)](
psum_num_recv_tokens_per_expert,
recv_x,
recv_x.stride(0),
scale_arg,
scale_s0,
scale_s1,
output_tensor,
output_tensor.stride(0),
output_tensor_scale if is_fp8 else scale_arg,
masked_m,
overflow,
MAX_M=max_m,
ALIGN=expert_alignment,
HIDDEN=hidden,
HIDDEN_PAD=triton.next_power_of_2(hidden),
SCALE_HIDDEN=sh,
SCALE_HIDDEN_PAD=triton.next_power_of_2(sh),
IS_FP8=is_fp8,
CHECK_OVERFLOW=check_overflow,
NUM_WORKERS=num_workers,
num_warps=4,
)
# Capture relies on max_m = cap * ep_group_size; eager also checks counts.
if check_overflow and int(overflow.item()) != 0:
raise RuntimeError(
f"DeepEP v2 masked slab overflow: an expert received more than max_m="
f"{max_m} tokens; increase "
f"SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK."
)
output_tensor = output_tensor.view(num_local_experts, max_m, hidden)
if is_fp8:
output_tensor_scale = output_tensor_scale.view(
num_local_experts, sh, max_m
).transpose(1, 2)
return output_tensor, output_tensor_scale, masked_m
@triton.jit
def _fwd_kernel_masked_slab_to_expand(
psum_ptr,
input_tensor_ptr,
input_tensor_stride0,
output_tensor_ptr,
output_tensor_stride0,
weight_ptr,
MAX_M: tl.constexpr,
ALIGN: tl.constexpr,
HIDDEN: tl.constexpr,
HIDDEN_PAD: tl.constexpr,
HAS_W: tl.constexpr,
NUM_WORKERS: tl.constexpr,
):
e = tl.program_id(0)
worker = tl.program_id(1)
prev_end = tl.load(psum_ptr + e - 1, mask=e > 0, other=0)
start = ((prev_end + ALIGN - 1) // ALIGN) * ALIGN
end = tl.load(psum_ptr + e)
count = end - start
count = tl.minimum(count, MAX_M)
off = tl.arange(0, HIDDEN_PAD)
mask = off < HIDDEN
for j in tl.range(worker, count, NUM_WORKERS):
src = (e * MAX_M + j).to(tl.int64)
dst = (start + j).to(tl.int64)
v = tl.load(input_tensor_ptr + src * input_tensor_stride0 + off, mask=mask)
if HAS_W:
w = tl.load(weight_ptr + dst)
v = (v.to(tl.float32) * w).to(v.dtype)
tl.store(output_tensor_ptr + dst * output_tensor_stride0 + off, v, mask=mask)
@torch.no_grad()
def masked_slab_to_expand(
input_tensor: torch.Tensor,
psum_num_recv_tokens_per_expert: torch.Tensor,
total_expanded_tokens: int,
expert_alignment: int,
topk_weights=None,
):
"""Convert masked-GEMM output to expanded order.
Only real rows are written; combine ignores uninitialized padding through the
handle. Optional top-k weights are fused into the copy over real rows only.
"""
num_local_experts, max_m, hidden = input_tensor.shape
output_tensor = torch.empty(
(total_expanded_tokens, hidden),
device=input_tensor.device,
dtype=input_tensor.dtype,
)
input_tensor2d = input_tensor.view(num_local_experts * max_m, hidden)
has_w = topk_weights is not None
if has_w:
weight_arg = topk_weights.reshape(-1).to(torch.float32).contiguous()
else:
weight_arg = input_tensor2d
num_workers = min(max_m, _DEEPEP_V2_REPACK_WORKERS_PER_EXPERT)
_fwd_kernel_masked_slab_to_expand[(num_local_experts, num_workers)](
psum_num_recv_tokens_per_expert,
input_tensor2d,
input_tensor2d.stride(0),
output_tensor,
output_tensor.stride(0),
weight_arg,
MAX_M=max_m,
ALIGN=expert_alignment,
HIDDEN=hidden,
HIDDEN_PAD=triton.next_power_of_2(hidden),
HAS_W=has_w,
NUM_WORKERS=num_workers,
num_warps=4,
)
return output_tensor
def _moe_permute_rows( def _moe_permute_rows(
inputs: torch.Tensor, inputs: torch.Tensor,
topk_ids: torch.Tensor, topk_ids: torch.Tensor,
+5 -1
View File
@@ -2806,7 +2806,7 @@ def _moe_runner_fusion_disable(view: Any) -> dict:
def _a2a_fusion_adjustments(view: Any) -> dict: def _a2a_fusion_adjustments(view: Any) -> dict:
"""A2A-backend-driven shared-experts fusion adjustments, declared at the """A2A-backend-driven shared-experts fusion adjustments, declared at the
legacy write slots in _handle_a2a_moe: Waterfill requires the legacy write slots in _handle_a2a_moe: Waterfill requires the
fusion enabled; FlashInfer A2A requires it disabled.""" fusion enabled; FlashInfer and DeepEP v2 A2A require it disabled."""
if view.moe_a2a_backend in ("deepep", "megamoe") and view.enable_waterfill: if view.moe_a2a_backend in ("deepep", "megamoe") and view.enable_waterfill:
if view.disable_shared_experts_fusion: if view.disable_shared_experts_fusion:
logger.warning( logger.warning(
@@ -2819,6 +2819,9 @@ def _a2a_fusion_adjustments(view: Any) -> dict:
"Flashinfer MoE A2A is enabled. --disable-shared-experts-fusion is automatically set." "Flashinfer MoE A2A is enabled. --disable-shared-experts-fusion is automatically set."
) )
return {"disable_shared_experts_fusion": True} return {"disable_shared_experts_fusion": True}
if view.moe_a2a_backend == "deepep_v2":
# Fused shared experts are not validated with DeepEP v2.
return {"disable_shared_experts_fusion": True}
return {} return {}
@@ -2827,6 +2830,7 @@ _A2A_EP_SPANNING_BACKENDS = frozenset(
{ {
"megamoe", "megamoe",
"deepep", "deepep",
"deepep_v2",
"mooncake", "mooncake",
"nixl", "nixl",
"ascend_fuseep", "ascend_fuseep",
+4
View File
@@ -1043,6 +1043,10 @@ class Envs:
# read by several call sites; do not use in new code. # read by several call sites; do not use in new code.
SGLANG_DEEPEP_BF16_DISPATCH = EnvBool(False) SGLANG_DEEPEP_BF16_DISPATCH = EnvBool(False)
SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK = EnvInt(128) SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK = EnvInt(128)
# Per-rank buffer capacity, not a model token limit.
SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK = EnvInt(128)
# 0 lets ElasticBuffer select its theoretical communication SM/QP counts.
SGLANG_DEEPEP_V2_NUM_SMS = EnvInt(0)
SGLANG_DEEPEP_LL_COMBINE_SEND_NUM_SMS = EnvInt(32) SGLANG_DEEPEP_LL_COMBINE_SEND_NUM_SMS = EnvInt(32)
SGLANG_BLACKWELL_OVERLAP_SHARED_EXPERTS_OUTSIDE_SBO = EnvBool(False) SGLANG_BLACKWELL_OVERLAP_SHARED_EXPERTS_OUTSIDE_SBO = EnvBool(False)
SGLANG_ENABLE_QWEN_DEEPEP_SHARED_OVERLAP = EnvBool(True) SGLANG_ENABLE_QWEN_DEEPEP_SHARED_OVERLAP = EnvBool(True)
+4 -1
View File
@@ -103,7 +103,9 @@ class DeepEPMoE(FusedMoE):
and quant_config is not None and quant_config is not None
and quant_config.get_name() == "humming" and quant_config.get_name() == "humming"
) )
if is_humming: if get_moe_a2a_backend().is_deepep_v2():
self.deprecate_flag = True
elif is_humming:
self.deprecate_flag = True self.deprecate_flag = True
elif _use_aiter: elif _use_aiter:
self.deprecate_flag = True self.deprecate_flag = True
@@ -354,6 +356,7 @@ def get_moe_impl_class(quant_config: Optional[QuantizationConfig]):
if ( if (
get_moe_a2a_backend().is_mori() get_moe_a2a_backend().is_mori()
or get_moe_a2a_backend().is_deepep() or get_moe_a2a_backend().is_deepep()
or get_moe_a2a_backend().is_deepep_v2()
or get_moe_a2a_backend().is_mooncake() or get_moe_a2a_backend().is_mooncake()
or get_moe_a2a_backend().is_nixl() or get_moe_a2a_backend().is_nixl()
or get_moe_a2a_backend().is_pplx() or get_moe_a2a_backend().is_pplx()
@@ -38,6 +38,7 @@ from sglang.srt.layers.moe.token_dispatcher.ascend_tp import (
AscendTPDispatcher, AscendTPDispatcher,
) )
from sglang.srt.layers.moe.token_dispatcher.base import BaseDispatcher from sglang.srt.layers.moe.token_dispatcher.base import BaseDispatcher
from sglang.srt.layers.moe.token_dispatcher.deepep_v2 import DeepEPv2Dispatcher
from sglang.srt.layers.moe.token_dispatcher.flashinfer import FlashinferDispatcher from sglang.srt.layers.moe.token_dispatcher.flashinfer import FlashinferDispatcher
from sglang.srt.layers.moe.token_dispatcher.standard import ( from sglang.srt.layers.moe.token_dispatcher.standard import (
StandardDispatcher, StandardDispatcher,
@@ -189,6 +190,15 @@ def create_moe_dispatcher(moe_runner_config: MoeRunnerConfig) -> BaseDispatcher:
async_finish=True, async_finish=True,
return_recv_hook=True, return_recv_hook=True,
) )
elif a2a_backend.is_deepep_v2():
return DeepEPv2Dispatcher(
group=get_tp_group().device_group,
router_topk=moe_runner_config.top_k,
num_experts=moe_runner_config.num_experts,
num_local_experts=moe_runner_config.num_local_experts,
hidden_size=moe_runner_config.hidden_size,
params_dtype=moe_runner_config.params_dtype,
)
elif a2a_backend.is_flashinfer(): elif a2a_backend.is_flashinfer():
return FlashinferDispatcher( return FlashinferDispatcher(
group=get_tp_group().device_group, group=get_tp_group().device_group,
@@ -226,6 +236,34 @@ def _validate_hpc_ops_quant_method(quant_method) -> None:
) )
def _validate_deepep_v2_quant_method(quant_method) -> None:
"""Validate the FP8 contract consumed by the DeepEP v2 adapter."""
if not get_moe_a2a_backend().is_deepep_v2():
return
config = (
quant_method.quant_config if isinstance(quant_method, Fp8MoEMethod) else None
)
reason = None
if not isinstance(quant_method, Fp8MoEMethod):
reason = f"selected {type(quant_method).__name__}"
elif quant_method.use_mxfp8:
reason = "selected MXFP8 weights"
elif quant_method.is_fp4_expert:
reason = "selected FP4 experts"
elif list(quant_method.weight_block_size or []) != [128, 128]:
reason = f"has weight_block_size={quant_method.weight_block_size}"
elif config.activation_scheme != "dynamic":
reason = f"has activation_scheme={config.activation_scheme!r}"
if reason is not None:
raise ValueError(
"--moe-a2a-backend deepep_v2 requires 128x128 blockwise FP8 "
f"experts with dynamic activation scaling, but this layer {reason}. "
"Use a compatible checkpoint or --moe-a2a-backend deepep."
)
class FusedMoE(torch.nn.Module): class FusedMoE(torch.nn.Module):
"""FusedMoE layer for MoE models. """FusedMoE layer for MoE models.
@@ -407,6 +445,7 @@ class FusedMoE(torch.nn.Module):
self.use_deep_gemm, self.use_deep_gemm,
) )
_validate_hpc_ops_quant_method(self.quant_method) _validate_hpc_ops_quant_method(self.quant_method)
_validate_deepep_v2_quant_method(self.quant_method)
self.supports_deferred_finalize = ( self.supports_deferred_finalize = (
envs.SGLANG_ENABLE_MOE_DEFERRED_FINALIZE.get() envs.SGLANG_ENABLE_MOE_DEFERRED_FINALIZE.get()
and get_moe_runner_backend().is_flashinfer_trtllm() and get_moe_runner_backend().is_flashinfer_trtllm()
@@ -49,6 +49,10 @@ if TYPE_CHECKING:
DeepEPNormalCombineInput, DeepEPNormalCombineInput,
DeepEPNormalDispatchOutput, DeepEPNormalDispatchOutput,
) )
from sglang.srt.layers.moe.token_dispatcher.deepep_v2 import (
DeepEPv2CombineInput,
DeepEPv2DispatchOutput,
)
from sglang.srt.layers.moe.token_dispatcher.standard import ( from sglang.srt.layers.moe.token_dispatcher.standard import (
StandardCombineInput, StandardCombineInput,
StandardDispatchOutput, StandardDispatchOutput,
@@ -206,6 +210,7 @@ class DeepGemmRunnerInput(RunnerInput):
masked_m: Optional[torch.Tensor] = None masked_m: Optional[torch.Tensor] = None
expected_m: Optional[int] = None expected_m: Optional[int] = None
m_indices: Optional[torch.Tensor] = None m_indices: Optional[torch.Tensor] = None
hidden_states_scale_tma_aligned: bool = False
@property @property
def runner_backend(self) -> MoeRunnerBackend: def runner_backend(self) -> MoeRunnerBackend:
@@ -321,7 +326,10 @@ class DeepGemmRunnerCore(MoeRunnerCore):
device=hidden_states_device, device=hidden_states_device,
dtype=torch.bfloat16, dtype=torch.bfloat16,
) )
if deep_gemm_wrapper.DEEPGEMM_NEED_TMA_ALIGNED_SCALES: if (
deep_gemm_wrapper.DEEPGEMM_NEED_TMA_ALIGNED_SCALES
and not runner_input.hidden_states_scale_tma_aligned
):
hidden_states_scale = tma_align_input_scale(hidden_states_scale) hidden_states_scale = tma_align_input_scale(hidden_states_scale)
deep_gemm_wrapper.grouped_gemm_nt_f8f8bf16_contig( deep_gemm_wrapper.grouped_gemm_nt_f8f8bf16_contig(
@@ -1419,3 +1427,186 @@ def _apply_swiglu_limit(
out = torch.cat([gate, up], dim=-1) out = torch.cat([gate, up], dim=-1)
assert out.shape == (num_tokens, hidden_size_x2) assert out.shape == (num_tokens, hidden_size_x2)
return out return out
@register_pre_permute("deepep_v2", "deep_gemm")
def pre_permute_deepep_v2_to_deep_gemm(
dispatch_output: DeepEPv2DispatchOutput,
quant_info: DeepGemmMoeQuantInfo,
runner_config: MoeRunnerConfig,
running_state: dict,
) -> DeepGemmRunnerInput:
from sglang.kernels.ops.moe.ep_moe_kernels import (
ep_expand_init_m_indices_from_psum,
ep_scatter_from_psum,
)
hidden_states = dispatch_output.hidden_states
hidden_states_scale = dispatch_output.hidden_states_scale
topk_ids = dispatch_output.topk_ids
topk_weights = dispatch_output.topk_weights
psum_num_recv_tokens_per_expert = dispatch_output.psum_num_recv_tokens_per_expert
is_expanded = dispatch_output.is_expanded
hidden_states_scale_tma_aligned = dispatch_output.hidden_states_scale_tma_aligned
deepep_v2_use_masked = dispatch_output.use_masked_gemm
deepep_v2_expected_m = dispatch_output.expected_m
deepep_v2_masked_max_m = dispatch_output.masked_max_m
deepep_v2_total_expanded = dispatch_output.total_expanded
deepep_v2_expert_alignment = dispatch_output.expert_alignment
if hidden_states_scale is None:
raise RuntimeError(
"DeepEP v2 -> DeepGEMM requires FP8 dispatch output with activation "
"scales, but the dispatch output carried none."
)
assert runner_config.activation == "silu"
if is_expanded:
if psum_num_recv_tokens_per_expert is None:
raise RuntimeError(
"DeepEP v2 requires the native expert prefix sums from the "
"ElasticBuffer dispatch handle."
)
all_tokens = hidden_states.shape[0]
running_state["all_tokens"] = all_tokens
running_state["hidden_states_shape"] = hidden_states.shape
running_state["hidden_states_device"] = hidden_states.device
running_state["hidden_states_dtype"] = hidden_states.dtype
running_state["topk_ids"] = None
running_state["topk_weights"] = topk_weights
running_state["deepep_v2_expanded"] = True
if deepep_v2_use_masked:
# masked_m bounds each expert independently of buffer capacity.
from sglang.kernels.ops.moe.ep_moe_kernels import expand_to_masked_slab
num_local_experts = psum_num_recv_tokens_per_expert.shape[0]
input_tensor, input_tensor_scale, masked_m = expand_to_masked_slab(
hidden_states,
hidden_states_scale,
psum_num_recv_tokens_per_expert,
num_local_experts,
deepep_v2_masked_max_m,
deepep_v2_expert_alignment,
)
running_state["deepep_v2_masked"] = True
running_state["deepep_v2_psum"] = psum_num_recv_tokens_per_expert
running_state["deepep_v2_total_expanded"] = deepep_v2_total_expanded
running_state["deepep_v2_expert_alignment"] = deepep_v2_expert_alignment
return DeepGemmRunnerInput(
hidden_states=input_tensor,
hidden_states_scale=input_tensor_scale,
use_masked_gemm=True,
masked_m=masked_m,
expected_m=deepep_v2_expected_m,
)
# Mark aligned expert rows and leave the unused receive tail at -1.
m_indices = torch.full(
(all_tokens,), -1, device=hidden_states.device, dtype=torch.int32
)
ep_expand_init_m_indices_from_psum(psum_num_recv_tokens_per_expert, m_indices)
return DeepGemmRunnerInput(
hidden_states=hidden_states,
hidden_states_scale=hidden_states_scale,
use_masked_gemm=False,
m_indices=m_indices,
hidden_states_scale_tma_aligned=hidden_states_scale_tma_aligned,
)
all_tokens = int(psum_num_recv_tokens_per_expert[-1].item())
K = hidden_states.shape[1]
running_state["all_tokens"] = all_tokens
running_state["hidden_states_shape"] = hidden_states.shape
running_state["hidden_states_device"] = hidden_states.device
running_state["hidden_states_dtype"] = hidden_states.dtype
running_state["topk_ids"] = topk_ids
running_state["topk_weights"] = topk_weights
input_tensor = torch.empty(
(all_tokens, K), device=hidden_states.device, dtype=hidden_states.dtype
)
if deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0:
# Packed UE8M0 scales require zero padding lanes.
input_tensor_scale = torch.zeros(
(ceil_div(K // 128, 4), all_tokens),
device=hidden_states.device,
dtype=torch.int,
).transpose(0, 1)
else:
input_tensor_scale = torch.empty(
(all_tokens, K // 128), device=hidden_states.device, dtype=torch.float32
)
m_indices = torch.empty(all_tokens, device=hidden_states.device, dtype=torch.int32)
output_index = torch.empty_like(topk_ids)
# Contiguous psum already includes the 128-row expert alignment.
expert_start_loc = torch.empty_like(psum_num_recv_tokens_per_expert)
ep_scatter_from_psum(
hidden_states,
hidden_states_scale,
topk_ids,
psum_num_recv_tokens_per_expert,
expert_start_loc,
input_tensor,
input_tensor_scale,
m_indices,
output_index,
scale_ue8m0=deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0,
)
dispose_tensor(hidden_states)
dispose_tensor(hidden_states_scale)
running_state["output_index"] = output_index
return DeepGemmRunnerInput(
hidden_states=input_tensor,
hidden_states_scale=input_tensor_scale,
use_masked_gemm=False,
m_indices=m_indices,
)
@register_post_permute("deep_gemm", "deepep_v2")
def post_permute_deep_gemm_to_deepep_v2(
runner_output: DeepGemmRunnerOutput,
quant_info: DeepGemmMoeQuantInfo,
runner_config: MoeRunnerConfig,
running_state: dict,
) -> DeepEPv2CombineInput:
from sglang.kernels.ops.moe.ep_moe_kernels import ep_gather
from sglang.srt.layers.moe.token_dispatcher.deepep_v2 import DeepEPv2CombineInput
if running_state.get("deepep_v2_expanded", False):
hidden_states = runner_output.hidden_states
topk_weights = running_state["topk_weights"]
if running_state.get("deepep_v2_masked", False):
# Expanded combine does not consume top-k weights.
from sglang.kernels.ops.moe.ep_moe_kernels import masked_slab_to_expand
hidden_states = masked_slab_to_expand(
hidden_states,
running_state["deepep_v2_psum"],
running_state["deepep_v2_total_expanded"],
running_state["deepep_v2_expert_alignment"],
topk_weights=topk_weights,
)
return DeepEPv2CombineInput(hidden_states, None)
if topk_weights is not None:
# Expanded combine does not consume top-k weights.
hidden_states = hidden_states * topk_weights.to(
hidden_states.dtype
).unsqueeze(-1)
return DeepEPv2CombineInput(hidden_states, None)
hidden_states = runner_output.hidden_states
topk_ids = running_state["topk_ids"]
topk_weights = running_state["topk_weights"]
output_index = running_state["output_index"]
gather_out = torch.empty(
running_state["hidden_states_shape"],
device=running_state["hidden_states_device"],
dtype=torch.bfloat16,
)
ep_gather(hidden_states, topk_ids, topk_weights, output_index, gather_out)
return DeepEPv2CombineInput(
hidden_states=gather_out,
topk_weights=topk_weights,
)
@@ -50,6 +50,15 @@ class MoeRunner:
"--moe-runner-backend hpc_ops for this model." "--moe-runner-backend hpc_ops for this model."
) )
if get_moe_a2a_backend().is_deepep_v2() and not runner_backend.is_deep_gemm():
raise ValueError(
"--moe-a2a-backend deepep_v2 requires the deep_gemm MoE runner, "
f"but this MoE layer's quantization method selected the "
f"'{runner_backend.value}' runner. deepep_v2 dispatches FP8 "
"activations plus scales, which only deep_gemm consumes; use an "
"FP8 blockwise-quantized checkpoint, or --moe-a2a-backend deepep."
)
self.fused_func = None self.fused_func = None
if runner_backend.is_triton(): if runner_backend.is_triton():
@@ -21,6 +21,11 @@ from sglang.srt.layers.moe.token_dispatcher.deepep import (
DeepEPNormalCombineInput, DeepEPNormalCombineInput,
DeepEPNormalDispatchOutput, DeepEPNormalDispatchOutput,
) )
from sglang.srt.layers.moe.token_dispatcher.deepep_v2 import (
DeepEPv2CombineInput,
DeepEPv2Dispatcher,
DeepEPv2DispatchOutput,
)
from sglang.srt.layers.moe.token_dispatcher.flashinfer import ( from sglang.srt.layers.moe.token_dispatcher.flashinfer import (
FlashinferDispatcher, FlashinferDispatcher,
FlashinferDispatchOutput, FlashinferDispatchOutput,
@@ -72,6 +77,9 @@ __all__ = [
"MoriEPLLDispatchOutput", "MoriEPLLDispatchOutput",
"MoriEPLLCombineInput", "MoriEPLLCombineInput",
"MoriEPDispatcher", "MoriEPDispatcher",
"DeepEPv2Dispatcher",
"DeepEPv2DispatchOutput",
"DeepEPv2CombineInput",
"NixlEPCombineInput", "NixlEPCombineInput",
"NixlEPDispatchOutput", "NixlEPDispatchOutput",
"NixlEPDispatcher", "NixlEPDispatcher",
@@ -27,6 +27,8 @@ if TYPE_CHECKING:
DeepEPLLDispatchOutput, DeepEPLLDispatchOutput,
DeepEPNormalCombineInput, DeepEPNormalCombineInput,
DeepEPNormalDispatchOutput, DeepEPNormalDispatchOutput,
DeepEPv2CombineInput,
DeepEPv2DispatchOutput,
FlashinferCombineInput, FlashinferCombineInput,
FlashinferDispatchOutput, FlashinferDispatchOutput,
StandardCombineInput, StandardCombineInput,
@@ -165,6 +167,12 @@ class DispatchOutputChecker:
) -> TypeGuard[FlashinferDispatchOutput]: ) -> TypeGuard[FlashinferDispatchOutput]:
return dispatch_output.format.is_flashinfer() return dispatch_output.format.is_flashinfer()
@staticmethod
def format_is_deepep_v2(
dispatch_output: DispatchOutput,
) -> TypeGuard[DeepEPv2DispatchOutput]:
return dispatch_output.format.is_deepep_v2()
class DispatchOutputFormat(Enum): class DispatchOutputFormat(Enum):
@@ -172,6 +180,7 @@ class DispatchOutputFormat(Enum):
DEEPEP_NORMAL = "deepep_normal" DEEPEP_NORMAL = "deepep_normal"
DEEPEP_LL = "deepep_ll" DEEPEP_LL = "deepep_ll"
FLASHINFER = "flashinfer" FLASHINFER = "flashinfer"
DEEPEP_V2 = "deepep_v2"
ASCEND_TP = "ascend_tp" ASCEND_TP = "ascend_tp"
def is_standard(self) -> bool: def is_standard(self) -> bool:
@@ -195,6 +204,9 @@ class DispatchOutputFormat(Enum):
def is_flashinfer(self) -> bool: def is_flashinfer(self) -> bool:
return self == DispatchOutputFormat.FLASHINFER return self == DispatchOutputFormat.FLASHINFER
def is_deepep_v2(self) -> bool:
return self == DispatchOutputFormat.DEEPEP_V2
@runtime_checkable @runtime_checkable
class DispatchOutput(Protocol): class DispatchOutput(Protocol):
@@ -249,12 +261,19 @@ class CombineInputChecker:
) -> TypeGuard[FlashinferCombineInput]: ) -> TypeGuard[FlashinferCombineInput]:
return combine_input.format == CombineInputFormat.FLASHINFER return combine_input.format == CombineInputFormat.FLASHINFER
@staticmethod
def format_is_deepep_v2(
combine_input: CombineInput,
) -> TypeGuard[DeepEPv2CombineInput]:
return combine_input.format == CombineInputFormat.DEEPEP_V2
class CombineInputFormat(Enum): class CombineInputFormat(Enum):
STANDARD = "standard" STANDARD = "standard"
DEEPEP_NORMAL = "deepep_normal" DEEPEP_NORMAL = "deepep_normal"
DEEPEP_LL = "deepep_ll" DEEPEP_LL = "deepep_ll"
FLASHINFER = "flashinfer" FLASHINFER = "flashinfer"
DEEPEP_V2 = "deepep_v2"
ASCEND_TP = "ascend_tp" ASCEND_TP = "ascend_tp"
@@ -0,0 +1,460 @@
from __future__ import annotations
import logging
import os
from typing import NamedTuple, Optional
import torch
import torch.distributed as dist
from sglang.srt.environ import envs
from sglang.srt.layers.dp_attention import get_is_extend_in_batch
from sglang.srt.layers.moe.token_dispatcher.base import (
BaseDispatcher,
CombineInput,
CombineInputFormat,
DispatchOutput,
DispatchOutputFormat,
)
from sglang.srt.layers.moe.topk import TopKOutput
from sglang.srt.layers.moe.utils import (
DeepEPv2Fp8ScaleFormat,
get_deepep_v2_fp8_scale_format,
)
logger = logging.getLogger(__name__)
_SCALE_BLOCK_SIZE = 128
# Must match DeepGEMM's contiguous expert alignment.
_EXPERT_ALIGNMENT = 128
_deepep_v2_import_error: Optional[BaseException] = None
_fp8_quant_import_error: Optional[BaseException] = None
sglang_per_token_group_quant_fp8 = None
try:
from deep_ep import ElasticBuffer
use_deepep_v2 = True
except (ImportError, OSError) as exc:
use_deepep_v2 = False
_deepep_v2_import_error = exc
if use_deepep_v2:
try:
from sglang.kernels.ops.quantization.fp8_kernel import (
sglang_per_token_group_quant_fp8,
)
except (ImportError, OSError) as exc:
_fp8_quant_import_error = exc
class DeepEPv2DispatchOutput(NamedTuple):
hidden_states: torch.Tensor
hidden_states_scale: Optional[torch.Tensor]
topk_ids: Optional[torch.Tensor]
topk_weights: torch.Tensor
psum_num_recv_tokens_per_expert: Optional[torch.Tensor] = None
is_expanded: bool = False
hidden_states_scale_tma_aligned: bool = False
use_masked_gemm: bool = False
expected_m: int = 0
masked_max_m: int = 0
total_expanded: int = 0
expert_alignment: int = 128
@property
def format(self) -> DispatchOutputFormat:
return DispatchOutputFormat.DEEPEP_V2
class DeepEPv2CombineInput(NamedTuple):
hidden_states: torch.Tensor
topk_weights: Optional[torch.Tensor]
@property
def format(self) -> CombineInputFormat:
return CombineInputFormat.DEEPEP_V2
assert isinstance(DeepEPv2DispatchOutput, DispatchOutput)
assert isinstance(DeepEPv2CombineInput, CombineInput)
def _raise_deepep_v2_import_error() -> None:
detail = (
f" Original import error: {_deepep_v2_import_error}"
if _deepep_v2_import_error is not None
else ""
)
raise ImportError(
"DeepEP v2 (ElasticBuffer) is not available. Install DeepEP v2 from "
"https://github.com/deepseek-ai/DeepEP." + detail
)
def _ensure_deepep_v2_available() -> None:
if not use_deepep_v2:
_raise_deepep_v2_import_error()
def _ensure_fp8_quant_available() -> None:
_ensure_deepep_v2_available()
if sglang_per_token_group_quant_fp8 is None:
detail = (
f" Original import error: {_fp8_quant_import_error}"
if _fp8_quant_import_error is not None
else ""
)
raise ImportError(
"DeepEP v2 FP8 dispatch requires the SGLang FP8 quantization kernel."
+ detail
)
def _get_allow_hybrid_mode() -> bool:
from sglang.srt.runtime_context import get_exec
return get_exec().moe.deepep_v2_mode == "hybrid"
def _quantize_for_deepep_v2_dispatch(
hidden_states: torch.Tensor, scale_format: DeepEPv2Fp8ScaleFormat
):
_ensure_fp8_quant_available()
return sglang_per_token_group_quant_fp8(
hidden_states,
_SCALE_BLOCK_SIZE,
column_major_scales=scale_format.tma_aligned,
scale_tma_aligned=scale_format.tma_aligned,
scale_ue8m0=scale_format.ue8m0,
)
class DeepEPv2Buffer:
"""Facade for the process-wide ElasticBuffer stored in runtime resources."""
_STATE_KEY = "deepep_v2_ep_state"
@classmethod
def _state(cls):
from types import SimpleNamespace
from sglang.srt.runtime_context import get_resources
buffers = get_resources().buffers
state = buffers.get(cls._STATE_KEY)
if state is None:
state = SimpleNamespace(buffer=None, key=None)
buffers[cls._STATE_KEY] = state
return state
@classmethod
def get_buffer(
cls,
group: dist.ProcessGroup,
hidden_size: int,
router_topk: int,
num_max_dispatch_tokens_per_rank: int,
use_fp8_dispatch: bool,
allow_hybrid_mode: Optional[bool] = None,
) -> ElasticBuffer:
_ensure_deepep_v2_available()
if allow_hybrid_mode is None:
allow_hybrid_mode = _get_allow_hybrid_mode()
state = cls._state()
# A key change rebuilds ElasticBuffer collectively on every rank.
key = (
group,
hidden_size,
router_topk,
num_max_dispatch_tokens_per_rank,
use_fp8_dispatch,
allow_hybrid_mode,
dist.get_world_size(group),
)
if state.buffer is not None and state.key == key:
return state.buffer
# Native explicit teardown is unavailable unless explicitly_destroy=True.
cls.destroy()
# Communicator reuse requires a device-bound process group.
os.environ.setdefault("EP_REUSE_NCCL_COMM", "0")
buffer = ElasticBuffer(
group,
num_max_tokens_per_rank=num_max_dispatch_tokens_per_rank,
hidden=hidden_size,
num_topk=router_topk,
use_fp8_dispatch=use_fp8_dispatch,
allow_hybrid_mode=allow_hybrid_mode,
sl_idx=0,
prefer_overlap_with_compute=False,
)
# Publish only after collective construction succeeds.
state.buffer = buffer
state.key = key
logger.info(
"Initialized DeepEP v2 ElasticBuffer: world_size=%s hidden_size=%s "
"num_topk=%s max_dispatch_tokens_per_rank=%s use_fp8_dispatch=%s "
"allow_hybrid_mode=%s num_bytes=%s",
dist.get_world_size(group),
hidden_size,
router_topk,
num_max_dispatch_tokens_per_rank,
use_fp8_dispatch,
allow_hybrid_mode,
buffer.num_bytes,
)
return buffer
@classmethod
def destroy(cls) -> None:
state = cls._state()
state.buffer = None
state.key = None
class _DeepEPv2Impl:
def __init__(
self,
group: dist.ProcessGroup,
router_topk: int,
num_experts: int,
num_local_experts: int,
hidden_size: int,
scale_format: DeepEPv2Fp8ScaleFormat,
num_max_dispatch_tokens_per_rank: int,
):
self.group = group
self.router_topk = router_topk
self.num_experts = num_experts
self.num_local_experts = num_local_experts
self.hidden_size = hidden_size
self.scale_format = scale_format
self.num_max_dispatch_tokens_per_rank = num_max_dispatch_tokens_per_rank
self.rank = dist.get_rank(group)
self._handle = None
self._pad_empty_combine = False
def _destroy_handle(self) -> None:
self._handle = None
def _get_buffer(self) -> ElasticBuffer:
return DeepEPv2Buffer.get_buffer(
self.group,
self.hidden_size,
self.router_topk,
self.num_max_dispatch_tokens_per_rank,
True,
)
def _validate_common(
self, hidden_states: torch.Tensor, topk_ids: torch.Tensor
) -> None:
if hidden_states.shape[0] > self.num_max_dispatch_tokens_per_rank:
raise ValueError(
f"DeepEP v2 dispatch input exceeds the per-rank buffer capacity "
f"{self.num_max_dispatch_tokens_per_rank}, got {hidden_states.shape[0]}. "
"Increase SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK or "
"lower the active prefill/decode batch limit."
)
if hidden_states.shape[1] != self.hidden_size:
raise ValueError(
f"DeepEP v2 hidden size mismatch: expected {self.hidden_size}, "
f"got {hidden_states.shape[1]}"
)
if self.hidden_size % _SCALE_BLOCK_SIZE != 0:
raise ValueError(
"DeepEP v2 FP8 dispatch requires hidden_size multiple of "
f"{_SCALE_BLOCK_SIZE}, got {self.hidden_size}"
)
if topk_ids.shape[1] != self.router_topk:
raise ValueError(
f"DeepEP v2 topk mismatch: expected {self.router_topk}, "
f"got {topk_ids.shape[1]}"
)
def dispatch(
self, hidden_states: torch.Tensor, topk_output: TopKOutput
) -> DeepEPv2DispatchOutput:
if self._handle is not None:
raise RuntimeError(
"DeepEP v2 dispatch called while the previous dispatch handle is "
"still unconsumed (missing combine)"
)
_ensure_deepep_v2_available()
topk_weights = topk_output.topk_weights
topk_ids = topk_output.topk_ids.to(torch.int64)
self._validate_common(hidden_states, topk_ids)
# Decode uses expanded/masked layout; extend uses contiguous in both modes.
use_expand_layout = not get_is_extend_in_batch()
use_masked = use_expand_layout
# CPU-synced dispatch needs a dummy token to notify from an idle rank.
self._pad_empty_combine = (not use_masked) and hidden_states.shape[0] == 0
if self._pad_empty_combine:
hidden_states = hidden_states.new_zeros((1, hidden_states.shape[-1]))
# Dummy routes need distinct expert ids; zero weights null the result.
topk_ids = torch.arange(
topk_ids.shape[-1], dtype=topk_ids.dtype, device=topk_ids.device
).unsqueeze(0)
topk_weights = topk_weights.new_zeros((1, topk_weights.shape[-1]))
_ensure_fp8_quant_available()
if use_masked:
_ue8m0 = self.scale_format.ue8m0
dispatch_x = sglang_per_token_group_quant_fp8(
hidden_states,
_SCALE_BLOCK_SIZE,
column_major_scales=_ue8m0,
scale_tma_aligned=_ue8m0,
scale_ue8m0=_ue8m0,
)
use_tma_aligned_col_major_sf = _ue8m0
else:
dispatch_x = _quantize_for_deepep_v2_dispatch(
hidden_states, self.scale_format
)
use_tma_aligned_col_major_sf = self.scale_format.tma_aligned
# This collective argument must not depend on a rank-local batch.
num_max_tokens = self.num_max_dispatch_tokens_per_rank
# Masked dispatch stays asynchronous for CUDA graph capture.
do_cpu_sync_val = True
if use_masked:
do_cpu_sync_val = False
buffer = self._get_buffer()
recv_x, recv_topk_idx, recv_topk_weights, handle, event = buffer.dispatch(
dispatch_x,
topk_idx=topk_ids,
topk_weights=topk_weights,
num_experts=self.num_experts,
num_max_tokens_per_rank=num_max_tokens,
expert_alignment=_EXPERT_ALIGNMENT,
num_sms=envs.SGLANG_DEEPEP_V2_NUM_SMS.get(),
use_tma_aligned_col_major_sf=use_tma_aligned_col_major_sf,
do_cpu_sync=do_cpu_sync_val,
do_expand=use_expand_layout,
)
self._handle = handle
local_tokens = hidden_states.shape[0]
if event.event is not None:
event.current_stream_wait()
if isinstance(recv_x, tuple):
recv_hidden_states, recv_hidden_states_scale = recv_x
else:
recv_hidden_states = recv_x
recv_hidden_states_scale = None
if use_expand_layout:
# Expanded combine uses handle metadata instead of recv_topk_idx.
local_topk_ids = None
else:
num_recv_tokens = int(
handle.psum_num_recv_tokens_per_scaleup_rank[-1].item()
)
recv_topk_idx = recv_topk_idx[:num_recv_tokens]
recv_topk_weights = recv_topk_weights[:num_recv_tokens]
recv_hidden_states = recv_hidden_states[:num_recv_tokens]
if recv_hidden_states_scale is not None:
recv_hidden_states_scale = recv_hidden_states_scale[:num_recv_tokens]
local_topk_ids = recv_topk_idx
expected_m = 0
masked_max_m = 0
total_expanded = 0
if use_masked:
# expected_m is only a schedule hint; masked_m is the actual bound.
ep_group_size = max(1, self.num_experts // self.num_local_experts)
expected_m = max(
1,
(local_tokens * ep_group_size * self.router_topk + self.num_experts)
// self.num_experts,
)
# Account for the worst case where every rank targets one local expert.
masked_max_m = self.num_max_dispatch_tokens_per_rank * ep_group_size
total_expanded = recv_hidden_states.shape[0]
return DeepEPv2DispatchOutput(
recv_hidden_states,
recv_hidden_states_scale,
local_topk_ids,
recv_topk_weights,
handle.psum_num_recv_tokens_per_expert,
use_expand_layout,
use_tma_aligned_col_major_sf,
use_masked,
expected_m,
masked_max_m,
total_expanded,
_EXPERT_ALIGNMENT,
)
def combine(self, combine_input: DeepEPv2CombineInput) -> torch.Tensor:
if self._handle is None:
raise RuntimeError(
"DeepEP v2 combine called without a valid dispatch handle"
)
# Release the single-use handle even when combine fails.
try:
buffer = self._get_buffer()
combined_x, _, event = buffer.combine(
combine_input.hidden_states,
handle=self._handle,
topk_weights=combine_input.topk_weights,
)
if event.event is not None:
event.current_stream_wait()
if self._pad_empty_combine:
combined_x = combined_x[:0]
return combined_x
finally:
self._pad_empty_combine = False
self._destroy_handle()
class DeepEPv2Dispatcher(BaseDispatcher):
def __init__(
self,
group: dist.ProcessGroup,
router_topk: int,
num_experts: int,
num_local_experts: int,
hidden_size: int,
params_dtype: torch.dtype,
):
super().__init__()
if params_dtype != torch.bfloat16:
raise NotImplementedError(
"DeepEP v2 dispatch adapter currently expects BF16 model activations, "
f"got {params_dtype}"
)
scale_format = get_deepep_v2_fp8_scale_format()
self.num_max_dispatch_tokens_per_rank = (
envs.SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK.get()
)
self._impl = _DeepEPv2Impl(
group=group,
router_topk=router_topk,
num_experts=num_experts,
num_local_experts=num_local_experts,
hidden_size=hidden_size,
scale_format=scale_format,
num_max_dispatch_tokens_per_rank=self.num_max_dispatch_tokens_per_rank,
)
def dispatch(
self, hidden_states: torch.Tensor, topk_output: TopKOutput
) -> DispatchOutput:
return self._impl.dispatch(hidden_states, topk_output)
def combine(self, combine_input: CombineInput) -> torch.Tensor:
if combine_input.format != CombineInputFormat.DEEPEP_V2:
raise TypeError(
f"Expected DeepEP v2 combine input, got {combine_input.format}"
)
return self._impl.combine(combine_input)
+33 -2
View File
@@ -4,6 +4,7 @@ import logging
import os import os
from contextlib import contextmanager from contextlib import contextmanager
from enum import Enum, IntEnum from enum import Enum, IntEnum
from typing import NamedTuple
import torch import torch
@@ -40,6 +41,7 @@ class MoeA2ABackend(Enum):
ASCEND_TP = "ascend_tp" ASCEND_TP = "ascend_tp"
FLASHINFER = "flashinfer" FLASHINFER = "flashinfer"
MEGAMOE = "megamoe" MEGAMOE = "megamoe"
DEEPEP_V2 = "deepep_v2"
PPLX = "pplx" PPLX = "pplx"
CUSTOMIZED = "customized" CUSTOMIZED = "customized"
@@ -79,6 +81,9 @@ class MoeA2ABackend(Enum):
def is_megamoe(self): def is_megamoe(self):
return self == MoeA2ABackend.MEGAMOE return self == MoeA2ABackend.MEGAMOE
def is_deepep_v2(self):
return self == MoeA2ABackend.DEEPEP_V2
def is_pplx(self): def is_pplx(self):
return self == MoeA2ABackend.PPLX return self == MoeA2ABackend.PPLX
@@ -178,6 +183,13 @@ class MoeRunnerBackend(Enum):
return self == MoeRunnerBackend.AITER return self == MoeRunnerBackend.AITER
class DeepEPv2Fp8ScaleFormat(NamedTuple):
"""DeepGEMM FP8 activation-scale layout expected from DeepEP v2."""
tma_aligned: bool
ue8m0: bool
class DeepEPMode(Enum): class DeepEPMode(Enum):
NORMAL = "normal" NORMAL = "normal"
@@ -311,6 +323,19 @@ def get_ascend_dispatcher_output_dtype(dispatcher):
return DispatcherOutputDtype.BF16 return DispatcherOutputDtype.BF16
def get_deepep_v2_fp8_scale_format() -> DeepEPv2Fp8ScaleFormat:
"""Resolve the FP8 scale layout DeepEP v2 must pre-quantize into."""
from sglang.srt.layers import deep_gemm_wrapper
return DeepEPv2Fp8ScaleFormat(
tma_aligned=(
deep_gemm_wrapper.DEEPGEMM_NEED_TMA_ALIGNED_SCALES
or deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0
),
ue8m0=deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0,
)
def initialize_moe_config(): def initialize_moe_config():
"""Seed the MoE runtime flags from the published configuration. """Seed the MoE runtime flags from the published configuration.
@@ -502,9 +527,15 @@ def is_sbo_enabled() -> bool:
def is_deepep_class_backend() -> bool: def is_deepep_class_backend() -> bool:
"""Check if the MoE backend is DeepEP-family (DeepEP, Mooncake, Mori, or PPLX).""" """Return whether A2A combine occurs inside a DeepEP-family dispatcher."""
b = get_moe_a2a_backend() b = get_moe_a2a_backend()
return b.is_deepep() or b.is_mooncake() or b.is_mori() or b.is_pplx() return (
b.is_deepep()
or b.is_deepep_v2()
or b.is_mooncake()
or b.is_mori()
or b.is_pplx()
)
def uses_per_rank_fused_shared_slots() -> bool: def uses_per_rank_fused_shared_slots() -> bool:
+8 -4
View File
@@ -744,6 +744,7 @@ class DeepseekV2MoE(nn.Module):
or get_moe_a2a_backend().is_ascend_fuseep() or get_moe_a2a_backend().is_ascend_fuseep()
or get_moe_a2a_backend().is_flashinfer() or get_moe_a2a_backend().is_flashinfer()
or get_moe_a2a_backend().is_megamoe() or get_moe_a2a_backend().is_megamoe()
or get_moe_a2a_backend().is_deepep_v2()
or should_use_flashinfer_cutlass_moe_fp4_allgather() or should_use_flashinfer_cutlass_moe_fp4_allgather()
or envs.SGLANG_SHARED_EXPERT_TP1.get() or envs.SGLANG_SHARED_EXPERT_TP1.get()
) )
@@ -833,6 +834,7 @@ class DeepseekV2MoE(nn.Module):
or get_moe_a2a_backend().is_nixl() or get_moe_a2a_backend().is_nixl()
or get_moe_a2a_backend().is_mori() or get_moe_a2a_backend().is_mori()
or get_moe_a2a_backend().is_ascend_fuseep() or get_moe_a2a_backend().is_ascend_fuseep()
or get_moe_a2a_backend().is_deepep_v2()
): ):
# TODO: we will support tp < ep in the future # TODO: we will support tp < ep in the future
self.ep_size = get_parallel().moe_ep_size self.ep_size = get_parallel().moe_ep_size
@@ -855,6 +857,7 @@ class DeepseekV2MoE(nn.Module):
or get_moe_a2a_backend().is_mori() or get_moe_a2a_backend().is_mori()
or get_moe_a2a_backend().is_ascend_fuseep() or get_moe_a2a_backend().is_ascend_fuseep()
or get_moe_a2a_backend().is_flashinfer() or get_moe_a2a_backend().is_flashinfer()
or get_moe_a2a_backend().is_deepep_v2()
) )
self._fuse_shared_experts_inside_sbo = SboFlags.fuse_shared_experts_inside_sbo() self._fuse_shared_experts_inside_sbo = SboFlags.fuse_shared_experts_inside_sbo()
# SGLANG_OPT_MOE_QUANT_ONCE eligibility, resolved lazily on first # SGLANG_OPT_MOE_QUANT_ONCE eligibility, resolved lazily on first
@@ -2757,10 +2760,11 @@ class DeepseekV2Model(nn.Module):
) )
) )
self.layers_to_capture = [] self.layers_to_capture = []
if get_moe_a2a_backend().is_deepep() or get_moe_a2a_backend().is_mooncake(): self.enable_a2a_moe = (
self.enable_a2a_moe = True get_moe_a2a_backend().is_deepep()
else: or get_moe_a2a_backend().is_mooncake()
self.enable_a2a_moe = False or get_moe_a2a_backend().is_deepep_v2()
)
# llama_4_scaling: for supporting Mistral-Large-3 model # llama_4_scaling: for supporting Mistral-Large-3 model
self.llama_4_scaling_config = getattr(config, "llama_4_scaling", None) self.llama_4_scaling_config = getattr(config, "llama_4_scaling", None)
+166
View File
@@ -303,10 +303,20 @@ MOE_A2A_BACKEND_CHOICES = [
"ascend_fuseep", "ascend_fuseep",
"flashinfer", "flashinfer",
"megamoe", "megamoe",
"deepep_v2",
"pplx", "pplx",
"ascend_tp", "ascend_tp",
] ]
# These architectures take the A2A MoE path and skip post-expert all-reduce.
_DEEPEP_V2_VALIDATED_ARCHITECTURES = frozenset(
{
"DeepseekV3ForCausalLM",
"DeepseekV4ForCausalLM",
"Qwen3MoeForCausalLM",
}
)
MXFP8_MOE_RUNNER_BACKEND_CHOICES = [ MXFP8_MOE_RUNNER_BACKEND_CHOICES = [
"cutlass", "cutlass",
"deep_gemm", "deep_gemm",
@@ -2443,6 +2453,8 @@ class ServerArgs:
"ascend_fuseep", "ascend_fuseep",
"flashinfer", "flashinfer",
"megamoe", "megamoe",
"deepep_v2",
"ascend_tp",
"pplx", "pplx",
], ],
Arg( Arg(
@@ -2459,6 +2471,15 @@ class ServerArgs:
"--moe-a2a-backend megamoe.", "--moe-a2a-backend megamoe.",
NS("exec.moe"), NS("exec.moe"),
] = False ] = False
deepep_v2_mode: A[
Literal["direct", "hybrid"],
"DeepEP v2 ElasticBuffer communication topology, fixed at server init: "
"`direct` (single-node NVLink) or `hybrid` (multi-node scale-out). "
"Layout/grouped-GEMM and the decode CUDA graph are chosen per batch by "
"inference phase, independent of this knob; not equivalent to DeepEP v1 "
"normal/low_latency.",
NS("exec.moe"),
] = "direct"
moe_runner_backend: A[ moe_runner_backend: A[
str, str,
Arg( Arg(
@@ -4020,6 +4041,10 @@ class ServerArgs:
# time; last declarations of the resolution, mirroring that order. # time; last declarations of the resolution, mirroring that order.
self._handle_model_capability_adjustments() self._handle_model_capability_adjustments()
# Validate after all batch-size declarations are visible.
self._validate_deepep_v2_speculative_draft()
self._validate_deepep_v2_dispatch_token_budget()
self._resolution_finished = True self._resolution_finished = True
def _handle_return_hidden_states_mode(self): def _handle_return_hidden_states_mode(self):
@@ -7415,6 +7440,93 @@ class ServerArgs:
f"(e.g. --max-prefill-tokens) to <= {max_cutedsl_tokens}." f"(e.g. --max-prefill-tokens) to <= {max_cutedsl_tokens}."
) )
def _validate_deepep_v2_dispatch_token_budget(self) -> None:
"""Check the configured prefill and decode-graph buffer bounds."""
view = resolved_view(self)
if view.moe_a2a_backend != "deepep_v2":
return
capacity = envs.SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK.get()
if view.disaggregation_mode != "decode":
prefill_tokens = self.max_prefill_buffer_tokens() or (
view.max_prefill_tokens or 0
)
if prefill_tokens > capacity:
raise ValueError(
"DeepEP v2 per-rank prefill budget exceeds "
"SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK: "
f"required={prefill_tokens}, capacity={capacity}. Raise the "
"environment value or lower --chunked-prefill-size/"
"--max-prefill-tokens."
)
if view.disaggregation_mode == "prefill":
return
decode_config = getattr(view.cuda_graph_config, "decode", None)
if decode_config is None or decode_config.backend == Backend.DISABLED:
return
graph_bs = decode_config.max_bs or 0
if view.max_running_requests is not None:
attn_dp_size = view.dp_size if view.enable_dp_attention else 1
per_rank_pool_bs = max(1, view.max_running_requests // attn_dp_size)
graph_bs = min(graph_bs, per_rank_pool_bs)
tokens_per_req = (
self.max_speculative_num_draft_tokens or 1
if view.speculative_algorithm
else 1
)
graph_tokens = graph_bs * tokens_per_req
if graph_tokens > capacity:
raise ValueError(
"DeepEP v2 per-rank decode CUDA graph exceeds "
"SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK: "
f"required={graph_tokens}, capacity={capacity} "
f"(requests={graph_bs}, tokens/request={tokens_per_req}). Raise "
"the environment value or lower --cuda-graph-max-bs."
)
def _validate_deepep_v2_model_architecture(self) -> None:
"""Allow DeepEP v2 only where its model workflow is validated."""
if (
parse_connector_type(resolved_view(self).model_path)
== ConnectorType.INSTANCE
):
raise ValueError(
"DeepEP v2 MoE cannot validate a model loaded through an instance "
"connector. Load it from a model path or use "
"--moe-a2a-backend deepep."
)
architectures = (
getattr(self.get_model_config().hf_config, "architectures", None) or []
)
architecture = architectures[0] if architectures else None
if architecture not in _DEEPEP_V2_VALIDATED_ARCHITECTURES:
raise ValueError(
f"DeepEP v2 MoE is not validated for {architecture!r}; supported "
f"architectures are {sorted(_DEEPEP_V2_VALIDATED_ARCHITECTURES)}. "
"Other model workflows may require an all-reduce after A2A "
"combine. Use --moe-a2a-backend deepep."
)
def _validate_deepep_v2_speculative_draft(self) -> None:
"""Reject an explicit or inherited DeepEP v2 draft backend."""
view = resolved_view(self)
draft_backend = view.speculative_moe_a2a_backend
if draft_backend is None and view.speculative_algorithm:
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
algorithm = SpeculativeAlgorithm.from_string(view.speculative_algorithm)
if not algorithm.is_ngram():
draft_backend = view.moe_a2a_backend
if draft_backend == "deepep_v2":
raise ValueError(
"DeepEP v2 MoE is not validated as a speculative draft backend. "
"Select another --speculative-moe-a2a-backend."
)
def _handle_a2a_moe(self): def _handle_a2a_moe(self):
# The backend overrides and the ep_size=tp_size adjustments moved to # The backend overrides and the ep_size=tp_size adjustments moved to
# the resolution pipeline (arg_groups/overrides.py: # the resolution pipeline (arg_groups/overrides.py:
@@ -7466,6 +7578,60 @@ class ServerArgs:
cfg.cuda_graph_config.decode.backend = Backend.DISABLED cfg.cuda_graph_config.decode.backend = Backend.DISABLED
cfg.cuda_graph_config.prefill.backend = Backend.DISABLED cfg.cuda_graph_config.prefill.backend = Backend.DISABLED
if a2a_backend == "deepep_v2":
self._validate_deepep_v2_model_architecture()
if resolved_view(self).enable_deterministic_inference:
raise ValueError(
"DeepEP v2 does not forward deterministic=True to "
"ElasticBuffer, so deterministic sorting remains disabled. "
"Disable --enable-deterministic-inference or use "
"--moe-a2a-backend deepep."
)
# ElasticBuffer requires CUMEM, but not NVLS or its preallocation.
os.environ.setdefault("NCCL_CUMEM_ENABLE", "1")
# Respect model-level runner declarations before resolving auto.
resolved_runner = resolved_view(self).moe_runner_backend
if resolved_runner == "auto":
self._declare("_handle_a2a_moe", moe_runner_backend="deep_gemm")
logger.warning(
"DeepEP v2 MoE: resolved --moe-runner-backend auto -> deep_gemm."
)
elif resolved_runner != "deep_gemm":
raise ValueError(
"DeepEP v2 MoE currently supports only "
f"--moe-runner-backend deep_gemm. Got {resolved_runner!r}. "
"Add a runner adapter before enabling DeepEP v2 with other "
"MoE runners."
)
if cfg.enable_two_batch_overlap or cfg.enable_single_batch_overlap:
raise ValueError(
"DeepEP v2 MoE has not implemented the TBO/SBO overlap hooks yet. "
"Disable --enable-two-batch-overlap and "
"--enable-single-batch-overlap when using --moe-a2a-backend deepep_v2."
)
if cfg.enforce_shared_experts_fusion:
raise ValueError(
"DeepEP v2 MoE has not validated fused shared experts yet. "
"Remove --enforce-shared-experts-fusion when using "
"--moe-a2a-backend deepep_v2."
)
# Prefill reads host counts and is not graph-capturable.
cfg.cuda_graph_config.prefill.backend = Backend.DISABLED
logger.warning(
f"DeepEP v2 MoE is enabled. The expert parallel size is adjusted to be the same as the tensor parallel size[{cfg.tp_size}]."
)
logger.warning(
"DeepEP v2 MoE is using deepep_v2_mode=%s. This controls "
"ElasticBuffer direct/hybrid mode and is independent from "
"--deepep-mode normal/low_latency. DeepEP v2 MoE enables the "
"decode CUDA graph on the masked decode path (any comm mode) "
"and disables shared expert fusion. "
"SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK is a "
"per-rank communication buffer capacity, not a model limit; "
"increase it for large prefill/chunked-prefill workloads.",
cfg.deepep_v2_mode,
)
# The resolving view, not the field: `_a2a_backend_overrides` may have # The resolving view, not the field: `_a2a_backend_overrides` may have
# moved this already (waterfill forces `deepep`). # moved this already (waterfill forces `deepep`).
a2a_now = resolved_view(self).moe_a2a_backend a2a_now = resolved_view(self).moe_a2a_backend
@@ -20,6 +20,12 @@ from sglang.srt.runtime_context import (
from sglang.srt.state_capturer.base import BaseTopkCapturer from sglang.srt.state_capturer.base import BaseTopkCapturer
def _is_scattered_a2a_backend() -> bool:
"""Return whether routed tokens are scattered across attention-TP ranks."""
backend = get_moe_a2a_backend()
return backend.is_deepep() or backend.is_deepep_v2()
class RoutedExpertsCapturer(BaseTopkCapturer): class RoutedExpertsCapturer(BaseTopkCapturer):
"""Capturer for routed experts with host buffer. """Capturer for routed experts with host buffer.
@@ -84,11 +90,8 @@ class RoutedExpertsCapturer(BaseTopkCapturer):
device_topk_size=topk_size + num_fused_shared_experts, device_topk_size=topk_size + num_fused_shared_experts,
) )
# DeepEP a2a path: each attn-TP rank only sees its scattered slice of # Rebuild the full token batch before routed-expert readback.
# topk_ids. All-gather across attn-TP at capture time so device_cache if _is_scattered_a2a_backend():
# holds the full batch and the existing _get_local_slice / D2H sync
# paths work unchanged. Pre-allocate the gather target.
if get_moe_a2a_backend().is_deepep():
attn_tp_size = ( attn_tp_size = (
get_parallel().attn_tp_size if is_dp_attention_enabled() else 1 get_parallel().attn_tp_size if is_dp_attention_enabled() else 1
) )
@@ -102,7 +105,7 @@ class RoutedExpertsCapturer(BaseTopkCapturer):
) )
def capture(self, layer_id: int, topk_indices: torch.Tensor): def capture(self, layer_id: int, topk_indices: torch.Tensor):
if get_moe_a2a_backend().is_deepep(): if _is_scattered_a2a_backend():
local_topk = topk_indices local_topk = topk_indices
topk_indices = self.gather_buffer[ topk_indices = self.gather_buffer[
: local_topk.size(0) * get_parallel().attn_tp_size : local_topk.size(0) * get_parallel().attn_tp_size
@@ -116,10 +119,8 @@ class RoutedExpertsCapturer(BaseTopkCapturer):
can_run_graph: bool, can_run_graph: bool,
cuda_graph_batch: Optional[int], cuda_graph_batch: Optional[int],
) -> torch.Tensor: ) -> torch.Tensor:
# Under DeepEP, capture() already attn_tp_all_gathered into the head of # Gathered rows start at buffer offset zero on every DP rank.
# the per-rank buffer, so the local DP rank's data lives at [0:N_local] if is_dp_attention_enabled() and not _is_scattered_a2a_backend():
# rather than at the global [start_pos:end_pos] offset.
if is_dp_attention_enabled() and not get_moe_a2a_backend().is_deepep():
# GPU->CPU sync would break overlap; operate on CPU directly. # GPU->CPU sync would break overlap; operate on CPU directly.
local_start_pos, local_num_tokens = get_dp_local_slice_cpu( local_start_pos, local_num_tokens = get_dp_local_slice_cpu(
forward_batch, can_run_graph, cuda_graph_batch forward_batch, can_run_graph, cuda_graph_batch
@@ -0,0 +1,194 @@
"""DP>1 routed-expert readback parity for DeepEP-family A2A backends."""
import concurrent.futures
import json
import os
import unittest
import numpy as np
import pybase64
import requests
import torch
from sglang.srt.utils import kill_process_tree
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import (
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
DEFAULT_URL_FOR_TEST,
CustomTestCase,
popen_launch_server,
)
register_cuda_ci(est_time=900, stage="base-c", runner_config="4-gpu-h100")
_MODEL = os.environ.get("SGLANG_ROUTED_EXPERTS_TEST_MODEL", "deepseek-ai/DeepSeek-V3")
_NUM_EXPERTS = 24
_NUM_LAYERS = 1
_TOPK = 8
_DUMMY_WEIGHT_ENV = {
"SGLANG_ENABLE_ASYNC_ASSERT": "0",
"SGLANG_SANITIZE_NAN_LOGITS": "1",
"SGLANG_CUDA_COREDUMP": "0",
"CUDA_ENABLE_COREDUMP_ON_EXCEPTION": "0",
"SGLANG_CUDA_COREDUMP_BEFORE_CRASH": "0",
}
def _deep_ep_has(attr: str) -> bool:
try:
import deep_ep # noqa: F401
except ImportError:
return False
return hasattr(deep_ep, attr)
def _deep_ep_nccl_compatible() -> bool:
try:
version = torch.cuda.nccl.version()
except (AttributeError, RuntimeError):
return False
return version is not None and version >= (2, 30, 7)
class _ReadbackMixin:
backend_args: list
@classmethod
def setUpClass(cls):
cls.base_url = DEFAULT_URL_FOR_TEST
other_args = [
"--trust-remote-code",
"--load-format",
"dummy",
"--json-model-override-args",
json.dumps(
{
"num_hidden_layers": _NUM_LAYERS,
"first_k_dense_replace": 0,
"n_routed_experts": _NUM_EXPERTS,
}
),
"--tp",
"2",
"--dp",
"2",
"--ep",
"2",
"--enable-dp-attention",
"--enable-return-routed-experts",
"--disable-cuda-graph",
"--disable-radix-cache",
# Keep the startup budget within the test's 256-token buffer.
"--chunked-prefill-size",
"256",
"--mem-fraction-static",
"0.5",
*cls.backend_args,
]
cls.process = popen_launch_server(
_MODEL,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=other_args,
env={
**os.environ,
**_DUMMY_WEIGHT_ENV,
"SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK": "256",
"SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK": "256",
},
)
@classmethod
def tearDownClass(cls):
if getattr(cls, "process", None):
kill_process_tree(cls.process.pid)
def _one_request(self, i: int):
resp = requests.post(
self.base_url + "/generate",
json={
"text": f"{self._WORDS[i]} is item number {i}. Describe it in detail.",
"sampling_params": {"max_new_tokens": 24, "temperature": 0},
"return_routed_experts": True,
},
timeout=300,
)
self.assertEqual(resp.status_code, 200)
meta = resp.json()["meta_info"]
self.assertIn("routed_experts", meta)
arr = np.frombuffer(pybase64.b64decode(meta["routed_experts"]), dtype=np.int32)
self.assertEqual(
arr.size % (_NUM_LAYERS * _TOPK),
0,
f"req{i}: payload size {arr.size} not a multiple of layers*topk",
)
rows = arr.reshape(-1, _NUM_LAYERS, _TOPK)
self.assertGreater(rows.shape[0], 0)
self.assertTrue(
bool(((rows >= 0) & (rows < _NUM_EXPERTS)).all()),
f"req{i}: expert id out of range [{rows.min()}, {rows.max()}]",
)
return rows
_WORDS = ["Alpha", "Bravo", "Charlie", "Delta", "Echo", "Foxtrot"]
_N_REQ = 6
def test_dp2_readback(self):
solo = [self._one_request(i) for i in range(self._N_REQ)]
with concurrent.futures.ThreadPoolExecutor(max_workers=self._N_REQ) as ex:
conc = list(ex.map(self._one_request, range(self._N_REQ)))
for i in range(self._N_REQ):
a, b = solo[i], conc[i]
n = min(a.shape[0], b.shape[0])
total = match = 0
for t in range(n):
for layer in range(_NUM_LAYERS):
total += 1
if set(a[t, layer].tolist()) == set(b[t, layer].tolist()):
match += 1
frac = match / max(1, total)
self.assertGreaterEqual(
frac,
0.9,
f"req{i}: only {frac:.1%} of per-token expert sets match the "
"solo baseline — the capturer is reading rows that belong to "
"other tokens (DeepEP-class backend misclassification)",
)
@unittest.skipUnless(_deep_ep_has("Buffer"), "DeepEP (v1 Buffer) not installed")
class TestRoutedExpertsReadbackDeepEP(_ReadbackMixin, CustomTestCase):
backend_args = [
"--moe-a2a-backend",
"deepep",
"--deepep-mode",
"low_latency",
"--deepep-dispatcher-output-dtype",
"fp8",
"--moe-runner-backend",
"deep_gemm",
]
@unittest.skipUnless(
_deep_ep_has("ElasticBuffer"), "DeepEP v2 (ElasticBuffer) not installed"
)
@unittest.skipUnless(
_deep_ep_nccl_compatible(), "DeepEP v2 requires NCCL runtime >= 2.30.7"
)
class TestRoutedExpertsReadbackDeepEPv2(_ReadbackMixin, CustomTestCase):
backend_args = [
"--moe-a2a-backend",
"deepep_v2",
"--deepep-v2-mode",
"direct",
"--moe-runner-backend",
"deep_gemm",
]
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,154 @@
"""CPU-only tests for the DeepEP v2 ElasticBuffer ownership facade."""
import unittest
from unittest.mock import patch
import torch
from sglang.srt.layers.moe.token_dispatcher import deepep_v2
from sglang.srt.runtime_context import get_resources, reset_context
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=3, suite="base-a-test-cpu")
class _FakeGroup:
pass
class _FakeBuffer:
instances = []
def __init__(self, *args, **kwargs):
self.args = args
self.kwargs = kwargs
self.num_bytes = 1 << 20
type(self).instances.append(self)
class TestDeepEPv2BufferLifecycle(CustomTestCase):
def setUp(self):
reset_context()
_FakeBuffer.instances = []
self._patches = [
patch.object(deepep_v2, "use_deepep_v2", True),
patch.object(deepep_v2, "ElasticBuffer", _FakeBuffer, create=True),
patch.object(deepep_v2.dist, "get_world_size", return_value=8),
]
for item in self._patches:
item.start()
def tearDown(self):
reset_context()
for item in reversed(self._patches):
item.stop()
def _get(self, group=None, **overrides):
kwargs = {
"group": group or _FakeGroup(),
"hidden_size": 4096,
"router_topk": 8,
"num_max_dispatch_tokens_per_rank": 128,
"use_fp8_dispatch": True,
"allow_hybrid_mode": False,
}
kwargs.update(overrides)
return deepep_v2.DeepEPv2Buffer.get_buffer(**kwargs)
def test_same_key_reuses_buffer(self):
group = _FakeGroup()
first = self._get(group)
second = self._get(group)
self.assertIs(first, second)
self.assertEqual(len(_FakeBuffer.instances), 1)
def test_constructor_inputs_participate_in_key(self):
group = _FakeGroup()
first = self._get(group)
second = self._get(group, num_max_dispatch_tokens_per_rank=256)
third = self._get(
group,
num_max_dispatch_tokens_per_rank=256,
allow_hybrid_mode=True,
)
self.assertIsNot(first, second)
self.assertIsNot(second, third)
self.assertEqual(len(_FakeBuffer.instances), 3)
def test_key_keeps_process_group_object(self):
group = _FakeGroup()
self._get(group)
state = get_resources().buffers[deepep_v2.DeepEPv2Buffer._STATE_KEY]
self.assertIs(state.key[0], group)
def test_distinct_process_group_rebuilds(self):
first = self._get(_FakeGroup())
second = self._get(_FakeGroup())
self.assertIsNot(first, second)
self.assertEqual(len(_FakeBuffer.instances), 2)
def test_state_lives_in_runtime_resources(self):
self._get()
self.assertIn(
deepep_v2.DeepEPv2Buffer._STATE_KEY,
get_resources().buffers,
)
def test_reset_context_drops_state_and_rebuilds(self):
group = _FakeGroup()
self._get(group)
reset_context()
self.assertNotIn(
deepep_v2.DeepEPv2Buffer._STATE_KEY,
get_resources().buffers,
)
self._get(group)
self.assertEqual(len(_FakeBuffer.instances), 2)
def test_failed_constructor_is_not_published(self):
class _FailingBuffer:
def __init__(self, *args, **kwargs):
raise RuntimeError("construct failed")
with patch.object(deepep_v2, "ElasticBuffer", _FailingBuffer):
with self.assertRaisesRegex(RuntimeError, "construct failed"):
self._get()
state = get_resources().buffers[deepep_v2.DeepEPv2Buffer._STATE_KEY]
self.assertIsNone(state.buffer)
self.assertIsNone(state.key)
self._get()
self.assertEqual(len(_FakeBuffer.instances), 1)
def test_destroy_clears_facade_state(self):
group = _FakeGroup()
first = self._get(group)
deepep_v2.DeepEPv2Buffer.destroy()
state = get_resources().buffers[deepep_v2.DeepEPv2Buffer._STATE_KEY]
self.assertIsNone(state.buffer)
self.assertIsNone(state.key)
second = self._get(group)
self.assertIsNot(first, second)
def test_unavailable_deepep_fails_before_state_creation(self):
with patch.object(deepep_v2, "use_deepep_v2", False):
with self.assertRaisesRegex(ImportError, "github.com/deepseek-ai/DeepEP"):
self._get()
self.assertNotIn(
deepep_v2.DeepEPv2Buffer._STATE_KEY,
get_resources().buffers,
)
def test_dispatch_capacity_guard_uses_actual_input_rows(self):
impl = object.__new__(deepep_v2._DeepEPv2Impl)
impl.num_max_dispatch_tokens_per_rank = 4
impl.hidden_size = 128
impl.router_topk = 2
impl._validate_common(torch.empty(4, 128), torch.zeros(4, 2))
with self.assertRaisesRegex(ValueError, "per-rank buffer capacity"):
impl._validate_common(torch.empty(5, 128), torch.zeros(5, 2))
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,244 @@
"""Tests for the DeepEP v2 expanded/masked repack kernels."""
import unittest
import torch
from sglang.kernels.ops.moe.ep_moe_kernels import (
expand_to_masked_slab,
masked_slab_to_expand,
)
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=10, stage="base-b", runner_config="1-gpu-large")
DEVICE = "cuda"
def _build_layout(counts, align, hidden, dtype, with_scale=False, scale_hidden=4):
"""Build synthetic expanded-layout buffers for per-expert counts."""
starts, psum = [], []
prev_end = 0
for c in counts:
start = ((prev_end + align - 1) // align) * align
end = start + c
starts.append(start)
psum.append(end)
prev_end = end
total = max(((prev_end + align - 1) // align) * align, 1)
# Vary rows and columns to expose broadcast or stride errors.
base = torch.zeros((total, hidden), dtype=torch.float32, device=DEVICE)
col_gain = 1.0 + (torch.arange(hidden, device=DEVICE) % 2).float()
for s, c in zip(starts, counts):
for j in range(c):
base[s + j] = float((s + j) % 200 + 1) * col_gain
recv_x = base.to(dtype)
scale = None
if with_scale:
scale = torch.zeros((total, scale_hidden), dtype=torch.float32, device=DEVICE)
# Vary scale columns to expose pack-dimension stride errors.
col = torch.arange(scale_hidden, dtype=torch.float32, device=DEVICE)
for s, c in zip(starts, counts):
for j in range(c):
scale[s + j] = float((s + j) % 50 + 1) * 0.5 + col
psum_t = torch.tensor(psum, dtype=torch.int32, device=DEVICE)
return recv_x, scale, psum_t, starts, total
def _real_rows(starts, counts):
rows = []
for s, c in zip(starts, counts):
rows.extend(range(s, s + c))
return rows
class TestDeepEPv2MaskedSlab(CustomTestCase):
ALIGN = 16
HIDDEN = 8
MAX_M = 32
def _check_expand_roundtrip(self, counts, dtype, with_scale, topk=False):
recv_x, scale, psum, starts, total = _build_layout(
counts, self.ALIGN, self.HIDDEN, dtype, with_scale=with_scale
)
E = len(counts)
masked_x, masked_x_scale, masked_m = expand_to_masked_slab(
recv_x, scale, psum, E, self.MAX_M, self.ALIGN
)
self.assertEqual(masked_m.tolist(), list(counts))
self.assertEqual(tuple(masked_x.shape), (E, self.MAX_M, self.HIDDEN))
for e, (s, c) in enumerate(zip(starts, counts)):
for j in range(c):
torch.testing.assert_close(
masked_x[e, j].float(), recv_x[s + j].float()
)
if with_scale:
torch.testing.assert_close(
masked_x_scale[e, j].float(), scale[s + j].float()
)
weights = None
if topk:
weights = torch.zeros(total, dtype=torch.float32, device=DEVICE)
for r in _real_rows(starts, counts):
weights[r] = 0.25 + (r % 7) * 0.1
out = masked_slab_to_expand(
masked_x, psum, total, self.ALIGN, topk_weights=weights
)
self.assertEqual(tuple(out.shape), (total, self.HIDDEN))
for e, (s, c) in enumerate(zip(starts, counts)):
for j in range(c):
expected = masked_x[e, j].float()
if topk:
expected = (expected * weights[s + j]).to(masked_x.dtype).float()
torch.testing.assert_close(out[s + j].float(), expected)
def test_roundtrip_bf16(self):
self._check_expand_roundtrip([3, 0, 5, 1], torch.bfloat16, with_scale=False)
def test_roundtrip_bf16_with_topk_weight(self):
self._check_expand_roundtrip(
[2, 4, 0, 7], torch.bfloat16, with_scale=False, topk=True
)
def test_roundtrip_fp8_with_scale(self):
self._check_expand_roundtrip([3, 1, 6, 2], torch.float8_e4m3fn, with_scale=True)
def test_empty_experts(self):
self._check_expand_roundtrip([0, 0, 0, 0], torch.bfloat16, with_scale=False)
def test_single_hot_expert(self):
self._check_expand_roundtrip(
[0, self.MAX_M, 0, 0], torch.bfloat16, with_scale=False, topk=True
)
def test_count_at_max_m_boundary(self):
self._check_expand_roundtrip(
[self.MAX_M, 1, self.MAX_M], torch.bfloat16, with_scale=False
)
def test_overflow_fails_fast(self):
counts = [self.MAX_M + 1, 2]
recv_x, scale, psum, starts, total = _build_layout(
counts, self.ALIGN, self.HIDDEN, torch.bfloat16
)
with self.assertRaises(RuntimeError):
expand_to_masked_slab(
recv_x, None, psum, len(counts), self.MAX_M, self.ALIGN
)
def _production_packed_ue8m0_layout(self, counts):
"""Build expanded rows with the production packed UE8M0 quantizer."""
from sglang.kernels.ops.quantization.fp8_kernel import (
sglang_per_token_group_quant_fp8,
)
# hidden=1024 ensures the packed scale has multiple columns.
hidden = 1024
raw, _, psum, starts, total = _build_layout(
counts, self.ALIGN, hidden, torch.bfloat16
)
recv_x, recv_x_scale = sglang_per_token_group_quant_fp8(
raw,
128,
column_major_scales=True,
scale_tma_aligned=True,
scale_ue8m0=True,
)
self.assertEqual(recv_x_scale.dtype, torch.int32)
self.assertGreater(recv_x_scale.shape[1], 1, "pack dim must be indexed")
self.assertNotEqual(recv_x_scale.stride(1), 1)
return recv_x, recv_x_scale, psum, starts, total, hidden
def test_fp8_packed_ue8m0_scale_from_production_quantizer(self):
counts = [3, 1, 6, 2]
recv_x, recv_x_scale, psum, starts, _, hidden = (
self._production_packed_ue8m0_layout(counts)
)
E = len(counts)
masked_x, masked_x_scale, masked_m = expand_to_masked_slab(
recv_x, recv_x_scale, psum, E, self.MAX_M, self.ALIGN
)
self.assertEqual(masked_m.tolist(), list(counts))
self.assertEqual(tuple(masked_x.shape), (E, self.MAX_M, hidden))
for e, (s, c) in enumerate(zip(starts, counts)):
for j in range(c):
torch.testing.assert_close(
masked_x[e, j].float(), recv_x[s + j].float()
)
torch.testing.assert_close(masked_x_scale[e, j], recv_x_scale[s + j])
def test_expand_under_cuda_graph_capture(self):
# Exercise replay with the production packed scale layout.
counts = [3, 1, 6, 2]
recv_x, recv_x_scale, psum, starts, _, _ = self._production_packed_ue8m0_layout(
counts
)
E = len(counts)
warm = torch.cuda.Stream()
warm.wait_stream(torch.cuda.current_stream())
with torch.cuda.stream(warm):
expand_to_masked_slab(recv_x, recv_x_scale, psum, E, self.MAX_M, self.ALIGN)
torch.cuda.current_stream().wait_stream(warm)
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
masked_x, masked_x_scale, masked_m = expand_to_masked_slab(
recv_x, recv_x_scale, psum, E, self.MAX_M, self.ALIGN
)
graph.replay()
torch.cuda.synchronize()
self.assertEqual(masked_m.tolist(), list(counts))
for e, (s, c) in enumerate(zip(starts, counts)):
for j in range(c):
torch.testing.assert_close(
masked_x[e, j].float(), recv_x[s + j].float()
)
torch.testing.assert_close(masked_x_scale[e, j], recv_x_scale[s + j])
class TestDeepEPv2HandleLifecycle(CustomTestCase):
"""CPU-only dispatch/combine handle guards."""
@staticmethod
def _bare_impl():
from sglang.srt.layers.moe.token_dispatcher.deepep_v2 import _DeepEPv2Impl
impl = object.__new__(_DeepEPv2Impl)
impl._handle = None
impl._pad_empty_combine = False
return impl
def test_combine_without_dispatch_raises(self):
impl = self._bare_impl()
with self.assertRaisesRegex(RuntimeError, "without a valid dispatch handle"):
impl.combine(None)
def test_dispatch_with_unconsumed_handle_raises(self):
impl = self._bare_impl()
impl._handle = object()
with self.assertRaisesRegex(RuntimeError, "unconsumed"):
impl.dispatch(None, None)
def test_handle_cleared_when_combine_fails(self):
impl = self._bare_impl()
impl._handle = object()
impl._pad_empty_combine = True
def _boom():
raise RuntimeError("boom")
impl._get_buffer = _boom
with self.assertRaisesRegex(RuntimeError, "boom"):
impl.combine(None)
self.assertIsNone(impl._handle)
self.assertFalse(impl._pad_empty_combine)
if __name__ == "__main__":
unittest.main()
@@ -1,16 +1,13 @@
"""The hpc_ops MoE runner backend makes the standard dispatcher keep global """Startup guards for MoE runner and dispatcher quantization contracts."""
expert ids, so a quant method that silently falls back to another runner
(e.g. an unquantized MoE) would misroute tokens under EP>1. MoeRunner must
reject that combination loudly at startup.
"""
import sys import sys
from types import SimpleNamespace
import pytest import pytest
from sglang.srt.layers.moe.moe_runner.base import MoeRunnerConfig from sglang.srt.layers.moe.moe_runner.base import MoeRunnerConfig
from sglang.srt.layers.moe.moe_runner.runner import MoeRunner from sglang.srt.layers.moe.moe_runner.runner import MoeRunner
from sglang.srt.layers.moe.utils import MoeRunnerBackend from sglang.srt.layers.moe.utils import MoeA2ABackend, MoeRunnerBackend
from sglang.srt.runtime_context import get_flags from sglang.srt.runtime_context import get_flags
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
@@ -18,27 +15,27 @@ register_cpu_ci(est_time=6, suite="base-c-test-cpu")
@pytest.fixture @pytest.fixture
def _runner_backend_flag(): def _moe_flags():
moe = get_flags().moe moe = get_flags().moe
saved = moe.runner_backend saved = (moe.runner_backend, moe.a2a_backend)
yield moe yield moe
moe.runner_backend = saved moe.runner_backend, moe.a2a_backend = saved
def test_non_hpc_runner_rejected_when_hpc_ops_requested(_runner_backend_flag): def test_non_hpc_runner_rejected_when_hpc_ops_requested(_moe_flags):
_runner_backend_flag.runner_backend = MoeRunnerBackend.HPC_OPS _moe_flags.runner_backend = MoeRunnerBackend.HPC_OPS
with pytest.raises(ValueError, match="hpc_ops"): with pytest.raises(ValueError, match="hpc_ops"):
MoeRunner(MoeRunnerBackend.TRITON, MoeRunnerConfig()) MoeRunner(MoeRunnerBackend.TRITON, MoeRunnerConfig())
def test_triton_runner_allowed_without_hpc_ops(_runner_backend_flag): def test_triton_runner_allowed_without_hpc_ops(_moe_flags):
_runner_backend_flag.runner_backend = MoeRunnerBackend.TRITON _moe_flags.runner_backend = MoeRunnerBackend.TRITON
runner = MoeRunner(MoeRunnerBackend.TRITON, MoeRunnerConfig()) runner = MoeRunner(MoeRunnerBackend.TRITON, MoeRunnerConfig())
assert runner.runner_core is not None assert runner.runner_core is not None
def test_direct_kernel_quant_method_rejected_when_hpc_ops_requested( def test_direct_kernel_quant_method_rejected_when_hpc_ops_requested(
_runner_backend_flag, _moe_flags,
): ):
# W4AFp8MoEMethod never constructs a MoeRunner (apply() calls its kernel # W4AFp8MoEMethod never constructs a MoeRunner (apply() calls its kernel
# directly), so it bypasses the MoeRunner-level guard; the layer-level # directly), so it bypasses the MoeRunner-level guard; the layer-level
@@ -49,15 +46,95 @@ def test_direct_kernel_quant_method_rejected_when_hpc_ops_requested(
from sglang.srt.layers.quantization.fp8 import Fp8MoEMethod from sglang.srt.layers.quantization.fp8 import Fp8MoEMethod
from sglang.srt.layers.quantization.w4afp8 import W4AFp8MoEMethod from sglang.srt.layers.quantization.w4afp8 import W4AFp8MoEMethod
_runner_backend_flag.runner_backend = MoeRunnerBackend.HPC_OPS _moe_flags.runner_backend = MoeRunnerBackend.HPC_OPS
with pytest.raises(ValueError, match="hpc_ops"): with pytest.raises(ValueError, match="hpc_ops"):
_validate_hpc_ops_quant_method(object.__new__(W4AFp8MoEMethod)) _validate_hpc_ops_quant_method(object.__new__(W4AFp8MoEMethod))
# The FP8 method (the one the hpc_ops runner supports) passes. # The FP8 method (the one the hpc_ops runner supports) passes.
_validate_hpc_ops_quant_method(object.__new__(Fp8MoEMethod)) _validate_hpc_ops_quant_method(object.__new__(Fp8MoEMethod))
# Without hpc_ops requested, any quant method passes. # Without hpc_ops requested, any quant method passes.
_runner_backend_flag.runner_backend = MoeRunnerBackend.TRITON _moe_flags.runner_backend = MoeRunnerBackend.TRITON
_validate_hpc_ops_quant_method(object.__new__(W4AFp8MoEMethod)) _validate_hpc_ops_quant_method(object.__new__(W4AFp8MoEMethod))
def _fp8_method(**overrides):
from sglang.srt.layers.quantization.fp8 import Fp8MoEMethod
values = {
"activation_scheme": "dynamic",
"weight_block_size": (128, 128),
"use_mxfp8": False,
"is_fp4_expert": False,
}
values.update(overrides)
method = object.__new__(Fp8MoEMethod)
method.quant_config = SimpleNamespace(
activation_scheme=values["activation_scheme"],
)
method.weight_block_size = values["weight_block_size"]
method.use_mxfp8 = values["use_mxfp8"]
method.is_fp4_expert = values["is_fp4_expert"]
return method
def test_deepep_v2_quant_contract_accepts_blockwise_fp8(_moe_flags):
from sglang.srt.layers.moe.fused_moe_triton.layer import (
_validate_deepep_v2_quant_method,
)
_moe_flags.a2a_backend = MoeA2ABackend.DEEPEP_V2
_validate_deepep_v2_quant_method(_fp8_method(weight_block_size=[128, 128]))
@pytest.mark.parametrize(
("overrides", "expected"),
[
({"activation_scheme": "static"}, "activation_scheme"),
({"weight_block_size": None}, "weight_block_size"),
({"weight_block_size": (1, 32), "use_mxfp8": True}, "MXFP8"),
({"is_fp4_expert": True}, "FP4 experts"),
],
)
def test_deepep_v2_quant_contract_rejects_incompatible_fp8(
_moe_flags, overrides, expected
):
from sglang.srt.layers.moe.fused_moe_triton.layer import (
_validate_deepep_v2_quant_method,
)
_moe_flags.a2a_backend = MoeA2ABackend.DEEPEP_V2
with pytest.raises(ValueError, match=expected):
_validate_deepep_v2_quant_method(_fp8_method(**overrides))
def test_deepep_v2_quant_contract_rejects_incompatible_methods(_moe_flags):
from sglang.srt.layers.moe.fused_moe_triton.layer import (
_validate_deepep_v2_quant_method,
)
from sglang.srt.layers.quantization.unquant import UnquantizedFusedMoEMethod
from sglang.srt.layers.quantization.w4afp8 import W4AFp8MoEMethod
_moe_flags.a2a_backend = MoeA2ABackend.DEEPEP_V2
for method_type in (UnquantizedFusedMoEMethod, W4AFp8MoEMethod):
with pytest.raises(ValueError, match=method_type.__name__):
_validate_deepep_v2_quant_method(object.__new__(method_type))
def test_deepep_v2_quant_contract_does_not_affect_other_backends(_moe_flags):
from sglang.srt.layers.moe.fused_moe_triton.layer import (
_validate_deepep_v2_quant_method,
)
from sglang.srt.layers.quantization.unquant import UnquantizedFusedMoEMethod
_moe_flags.a2a_backend = MoeA2ABackend.DEEPEP
_validate_deepep_v2_quant_method(object.__new__(UnquantizedFusedMoEMethod))
def test_deepep_v2_runner_backstop(_moe_flags):
_moe_flags.a2a_backend = MoeA2ABackend.DEEPEP_V2
with pytest.raises(ValueError, match="deep_gemm"):
MoeRunner(MoeRunnerBackend.TRITON, MoeRunnerConfig())
assert MoeRunner(MoeRunnerBackend.DEEP_GEMM, MoeRunnerConfig()).runner_core
if __name__ == "__main__": if __name__ == "__main__":
sys.exit(pytest.main([__file__, "-v"])) sys.exit(pytest.main([__file__, "-v"]))
@@ -2106,6 +2106,328 @@ class TestSamplingBackendTokenOracleEnvGate(CustomTestCase):
self.assertEqual(parsed.sampling_backend, "token_oracle") self.assertEqual(parsed.sampling_backend, "token_oracle")
class TestDeepEPv2Args(CustomTestCase):
"""DeepEP v2 server-argument resolution and validation."""
def _args(self, **overrides):
server_args = ServerArgs(model_path="dummy", moe_a2a_backend="deepep_v2")
server_args.model_config = SimpleNamespace(
hf_config=SimpleNamespace(architectures=["DeepseekV4ForCausalLM"])
)
# The dummy path does not initialize phase configs.
server_args.cuda_graph_config = CudaGraphConfig(
decode=PhaseConfig(backend=Backend.FULL, max_bs=512),
prefill=PhaseConfig(backend=Backend.FULL, max_bs=512),
)
server_args._resolved_overrides = []
valid = {f.name for f in dataclasses.fields(ServerArgs)}
for key, value in overrides.items():
# Reject stale field names before setattr silently accepts them.
assert key in valid, f"{key} is not a ServerArgs field"
setattr(server_args, key, value)
return server_args
def test_validated_architectures_allowed(self):
for architecture in (
"DeepseekV3ForCausalLM",
"DeepseekV4ForCausalLM",
"Qwen3MoeForCausalLM",
):
args = self._args(moe_runner_backend="deep_gemm")
args.model_config.hf_config.architectures = [architecture]
args._handle_a2a_moe()
def test_unvalidated_and_missing_architectures_rejected(self):
for architectures in (
["Qwen2MoeForCausalLM"],
["Qwen3_5MoeForCausalLM"],
[],
None,
):
args = self._args(moe_runner_backend="deep_gemm")
args.model_config.hf_config.architectures = architectures
with self.assertRaisesRegex(ValueError, "not validated"):
args._handle_a2a_moe()
def test_instance_connector_rejected(self):
args = self._args(
model_path="instance://worker/model",
moe_runner_backend="deep_gemm",
)
with self.assertRaisesRegex(ValueError, "instance connector"):
args._handle_a2a_moe()
def test_deterministic_inference_rejected(self):
args = self._args(
moe_runner_backend="deep_gemm",
enable_deterministic_inference=True,
)
with self.assertRaisesRegex(ValueError, "deterministic sorting"):
args._handle_a2a_moe()
def test_rl_on_policy_deterministic_inference_rejected(self):
args = self._args(
moe_runner_backend="deep_gemm",
rl_on_policy_target="fsdp",
)
args.model_config.hf_config.architectures = ["Qwen3MoeForCausalLM"]
with (
envs.SGLANG_VLM_CACHE_SIZE_MB.override(envs.SGLANG_VLM_CACHE_SIZE_MB.get()),
envs.SGLANG_ENABLE_DETERMINISTIC_INFERENCE.override(
envs.SGLANG_ENABLE_DETERMINISTIC_INFERENCE.get()
),
):
args._handle_deterministic_inference()
with self.assertRaisesRegex(ValueError, "deterministic sorting"):
args._handle_a2a_moe()
def test_deterministic_inference_does_not_affect_legacy_deepep(self):
args = self._args(
moe_a2a_backend="deepep",
moe_runner_backend="deep_gemm",
enable_deterministic_inference=True,
)
args._handle_a2a_moe()
def test_runner_restored_by_declaration_fails_fast(self):
# Validate the declaration-resolved runner rather than the raw field.
args = self._args(moe_runner_backend="auto")
args._resolved_overrides = [
("test_mxfp8", {"moe_runner_backend": "flashinfer_trtllm"})
]
with self.assertRaises(ValueError):
args._handle_a2a_moe()
def test_declarations_resolve_ep_size_and_fusion(self):
from sglang.srt.arg_groups.overrides import resolved_view
args = self._args(moe_runner_backend="auto", tp_size=2)
args._handle_a2a_moe()
self.assertEqual(resolved_view(args).ep_size, args.tp_size)
self.assertTrue(resolved_view(args).disable_shared_experts_fusion)
def test_auto_runner_defaults_to_deep_gemm(self):
from sglang.srt.arg_groups.overrides import resolved_view
args = self._args(moe_runner_backend="auto")
args._handle_a2a_moe()
self.assertEqual(resolved_view(args).moe_runner_backend, "deep_gemm")
def test_unsupported_runner_rejected(self):
args = self._args(moe_runner_backend="flashinfer_trtllm")
with self.assertRaises(ValueError):
args._handle_a2a_moe()
def test_triton_runner_rejected(self):
args = self._args(moe_runner_backend="triton")
with self.assertRaises(ValueError):
args._handle_a2a_moe()
def test_decode_graph_stays_enabled_in_both_comm_modes(self):
for mode in ("direct", "hybrid"):
args = self._args(moe_runner_backend="deep_gemm", deepep_v2_mode=mode)
args._handle_a2a_moe()
self.assertEqual(args.cuda_graph_config.decode.backend, Backend.FULL)
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.DISABLED)
def test_two_batch_overlap_rejected(self):
args = self._args(moe_runner_backend="deep_gemm", enable_two_batch_overlap=True)
with self.assertRaises(ValueError):
args._handle_a2a_moe()
def test_speculative_draft_backend_rejected(self):
for main_backend in ("none", "deepep", "deepep_v2"):
args = self._args(
moe_a2a_backend=main_backend,
moe_runner_backend="deep_gemm",
speculative_moe_a2a_backend="deepep_v2",
)
with self.assertRaisesRegex(ValueError, "speculative draft backend"):
args._validate_deepep_v2_speculative_draft()
def test_inherited_speculative_draft_backend_rejected(self):
args = self._args(
moe_runner_backend="deep_gemm",
speculative_algorithm="EAGLE",
)
with self.assertRaisesRegex(ValueError, "speculative draft backend"):
args._validate_deepep_v2_speculative_draft()
def test_ngram_does_not_inherit_a_draft_backend(self):
args = self._args(
moe_runner_backend="deep_gemm",
speculative_algorithm="NGRAM",
)
args._validate_deepep_v2_speculative_draft()
def test_explicit_legacy_speculative_backend_allowed(self):
args = self._args(
moe_runner_backend="deep_gemm",
speculative_algorithm="EAGLE",
speculative_moe_a2a_backend="deepep",
)
args._validate_deepep_v2_speculative_draft()
def test_resolved_legacy_speculative_backend_allowed(self):
args = self._args(
moe_runner_backend="deep_gemm",
speculative_algorithm="EAGLE",
)
args._resolved_overrides = [
(
"test_speculative_backend",
{"speculative_moe_a2a_backend": "deepep"},
)
]
args._validate_deepep_v2_speculative_draft()
def test_prefill_chunk_exceeding_cap_rejected(self):
args = self._args(moe_runner_backend="deep_gemm", chunked_prefill_size=2048)
with envs.SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK.override(1024):
with self.assertRaisesRegex(ValueError, "NUM_MAX_DISPATCH_TOKENS_PER_RANK"):
args._validate_deepep_v2_dispatch_token_budget()
def test_prefill_chunk_at_cap_boundary_accepted(self):
args = self._args(moe_runner_backend="deep_gemm", chunked_prefill_size=1024)
with envs.SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK.override(1024):
args._validate_deepep_v2_dispatch_token_budget()
def test_dynamic_chunking_probe_is_included(self):
args = self._args(
chunked_prefill_size=1024,
max_prefill_tokens=1024,
enable_dynamic_chunking=True,
pp_size=2,
disaggregation_mode="prefill",
)
with envs.SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK.override(1024):
with self.assertRaisesRegex(ValueError, "required=1280"):
args._validate_deepep_v2_dispatch_token_budget()
def test_disabled_chunking_uses_max_prefill_tokens(self):
for disabled in (None, 0, -1):
args = self._args(
chunked_prefill_size=disabled,
max_prefill_tokens=1024,
disaggregation_mode="prefill",
)
with envs.SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK.override(128):
with self.assertRaisesRegex(ValueError, "required=1024"):
args._validate_deepep_v2_dispatch_token_budget()
def test_decode_role_skips_prefill_capacity(self):
args = self._args(
chunked_prefill_size=4096,
disaggregation_mode="decode",
max_running_requests=32,
dp_size=1,
)
with envs.SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK.override(128):
args._validate_deepep_v2_dispatch_token_budget()
def test_decode_graph_capacity_boundaries(self):
for max_bs, raises in ((128, False), (129, True)):
args = self._args(
disaggregation_mode="decode",
max_running_requests=None,
)
args.cuda_graph_config.decode.max_bs = max_bs
with envs.SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK.override(128):
if raises:
with self.assertRaisesRegex(ValueError, "decode CUDA graph"):
args._validate_deepep_v2_dispatch_token_budget()
else:
args._validate_deepep_v2_dispatch_token_budget()
def test_dp_attention_divides_max_running_requests_per_rank(self):
args = self._args(
disaggregation_mode="decode",
max_running_requests=256,
tp_size=8,
dp_size=8,
enable_dp_attention=True,
)
with envs.SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK.override(128):
args._validate_deepep_v2_dispatch_token_budget()
def test_tp_only_max_running_requests_is_not_divided(self):
args = self._args(
disaggregation_mode="decode",
max_running_requests=256,
tp_size=8,
dp_size=1,
enable_dp_attention=False,
)
with envs.SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK.override(128):
with self.assertRaisesRegex(ValueError, "decode CUDA graph"):
args._validate_deepep_v2_dispatch_token_budget()
def test_memory_derived_eager_pool_remains_runtime_validated(self):
args = self._args(
disaggregation_mode="decode",
max_running_requests=None,
)
args.cuda_graph_config.decode.backend = Backend.DISABLED
with envs.SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK.override(1):
args._validate_deepep_v2_dispatch_token_budget()
def test_speculative_decode_width_is_included(self):
args = self._args(
disaggregation_mode="decode",
speculative_algorithm="EAGLE",
speculative_num_draft_tokens=8,
max_running_requests=256,
dp_size=8,
enable_dp_attention=True,
)
with envs.SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK.override(128):
with self.assertRaisesRegex(ValueError, "tokens/request=8"):
args._validate_deepep_v2_dispatch_token_budget()
def test_adaptive_speculative_uses_widest_candidate(self):
args = self._args(
disaggregation_mode="decode",
speculative_algorithm="EAGLE",
speculative_num_draft_tokens=4,
speculative_adaptive=True,
max_running_requests=128,
dp_size=8,
enable_dp_attention=True,
)
with patch.object(
ServerArgs,
"max_speculative_num_draft_tokens",
new=property(lambda _self: 16),
):
with envs.SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK.override(128):
with self.assertRaisesRegex(ValueError, "tokens/request=16"):
args._validate_deepep_v2_dispatch_token_budget()
def test_prefill_role_skips_decode_capacity(self):
args = self._args(
disaggregation_mode="prefill",
chunked_prefill_size=64,
max_running_requests=8192,
)
with envs.SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK.override(128):
args._validate_deepep_v2_dispatch_token_budget()
def test_other_backend_skips_capacity_validation(self):
args = self._args(
moe_a2a_backend="deepep",
chunked_prefill_size=4096,
max_running_requests=4096,
)
with envs.SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK.override(1):
args._validate_deepep_v2_dispatch_token_budget()
def test_capacity_validation_uses_resolved_backend(self):
args = self._args(chunked_prefill_size=4096)
args._resolved_overrides = [("test", {"moe_a2a_backend": "deepep"})]
with envs.SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK.override(1):
args._validate_deepep_v2_dispatch_token_budget()
class TestHandleCrashDumpEnv(CustomTestCase): class TestHandleCrashDumpEnv(CustomTestCase):
_COREDUMP_ENV_KEYS = ( _COREDUMP_ENV_KEYS = (
"CUDA_ENABLE_COREDUMP_ON_EXCEPTION", "CUDA_ENABLE_COREDUMP_ON_EXCEPTION",
@@ -0,0 +1,86 @@
"""DeepEP-family backend recognition in RoutedExpertsCapturer."""
import unittest
from types import SimpleNamespace
from unittest import mock
import torch
from sglang.srt.layers.moe.utils import MoeA2ABackend
from sglang.srt.state_capturer import routed_experts as re_mod
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
class TestScatteredA2ABackendHelper(CustomTestCase):
def test_classification(self):
expected = {
"deepep": True,
"deepep_v2": True,
"none": False,
"mooncake": False,
}
for value, exp in expected.items():
with mock.patch.object(
re_mod, "get_moe_a2a_backend", return_value=MoeA2ABackend(value)
):
self.assertEqual(
re_mod._is_scattered_a2a_backend(), exp, f"backend={value}"
)
class TestGetLocalSliceBackendBranch(CustomTestCase):
T, L, K = 16, 3, 4
def _capturer(self):
cap = object.__new__(re_mod.RoutedExpertsCapturer)
buf = torch.arange(self.T * self.L * self.K, dtype=torch.int32).reshape(
self.T, self.L, self.K
)
cap.device_cache = SimpleNamespace(buffer=buf)
cap.topk_size = self.K
return cap, buf
def _slice(self, cap, n_local):
fb = SimpleNamespace(out_cache_loc=torch.empty(n_local))
return cap._get_local_slice(fb, can_run_graph=False, cuda_graph_batch=None)
def test_deepep_v2_reads_buffer_head(self):
cap, buf = self._capturer()
with mock.patch.object(
re_mod, "is_dp_attention_enabled", return_value=True
), mock.patch.object(
re_mod, "get_moe_a2a_backend", return_value=MoeA2ABackend("deepep_v2")
):
out = self._slice(cap, n_local=5)
self.assertTrue(torch.equal(out, buf[0:5, :, : self.K]))
def test_deepep_v2_matches_deepep(self):
cap, _ = self._capturer()
outs = []
for backend in ("deepep", "deepep_v2"):
with mock.patch.object(
re_mod, "is_dp_attention_enabled", return_value=True
), mock.patch.object(
re_mod, "get_moe_a2a_backend", return_value=MoeA2ABackend(backend)
):
outs.append(self._slice(cap, n_local=7))
self.assertTrue(torch.equal(outs[0], outs[1]))
def test_tp_moe_reads_global_offset(self):
cap, buf = self._capturer()
with mock.patch.object(
re_mod, "is_dp_attention_enabled", return_value=True
), mock.patch.object(
re_mod, "get_moe_a2a_backend", return_value=MoeA2ABackend("none")
), mock.patch.object(
re_mod, "get_dp_local_slice_cpu", return_value=(6, 4)
):
out = self._slice(cap, n_local=999)
self.assertTrue(torch.equal(out, buf[6:10, :, : self.K]))
if __name__ == "__main__":
unittest.main()