Integration mori backend for EP a2a data communication (#17012)

Co-authored-by: Duyi-Wang <duyi.wang@amd.com>
Co-authored-by: billishyahao <bill.he@amd.com>
Co-authored-by: HaiShaw <hixiao@gmail.com>
This commit is contained in:
kk
2026-01-28 19:07:34 -08:00
committed by GitHub
co-authored by Duyi-Wang billishyahao HaiShaw
parent 673dc09d9b
commit f1384f5293
15 changed files with 934 additions and 10 deletions
+3 -2
View File
@@ -15,12 +15,13 @@ SGLang's EP integrates diverse, highly efficient backends for different use case
| **`none` (default)** | Disables all-to-all for EP. Uses All-Reduce or All-Gather for token dispatch. | Hybrid EP and TP setups. |
| `deepep` | DeepEP, a communication library for efficient token shuffling in MoE models. | Large-scale EP deployments. |
| `mooncake` | An extension of DeepEP for elastic inference, leveraging RDMA for high-performance data transfers. | Elastic EP serving. |
| `mori` | MORI-EP, AMD's native all-to-all communication implementation optimized for ROCm. | AMD GPU deployments. |
| `flashinfer` | Flashinfer implementation of all-to-all. | Large-scale EP deployments. |
| `ascend_fuseep` | Ascend NPU native fused all-to-all communication. | Ascend NPU deployments. |
DeepEP and Mooncake backends support two modes for token dispatch: `normal` mode (optimized for prefill workloads with high throughput) and `low_latency` mode (optimized for decode workloads with low latency and CUDA Graph compatibility). Users are recommended to set `--deepep-mode auto` to enable automatic dispatch mode switching during runtime. Setting `--deepep-mode normal` or `--deepep-mode low_latency` is useful for debugging or development purposes.
DeepEP and Mooncake backends support two modes for token dispatch: `normal` mode (optimized for prefill workloads with high throughput) and `low_latency` mode (optimized for decode workloads with low latency and CUDA Graph compatibility). MORI backend only supports `normal` mode now. Users are recommended to set `--deepep-mode auto` to enable automatic dispatch mode switching during runtime. Setting `--deepep-mode normal` or `--deepep-mode low_latency` is useful for debugging or development purposes.
Currently, DeepEP and Mooncake only support cases where `ep_size = tp_size`. For hybrid EP and TP (i.e., `ep_size < tp_size`), only the `none` backend (All-Reduce or All-Gather-based dispatching) is supported.
Currently, DeepEP, Mooncake and MORI only support cases where `ep_size = tp_size`. For hybrid EP and TP (i.e., `ep_size < tp_size`), only the `none` backend (All-Reduce or All-Gather-based dispatching) is supported.
### Backends for MoE Computation
+10
View File
@@ -66,6 +66,16 @@ SGLang supports various environment variables that can be used to configure its
| `SGLANG_DEEPEP_LL_COMBINE_SEND_NUM_SMS` | Number of SMs used for DeepEP combine when single batch overlap is enabled | `"32"` |
| `SGLANG_BLACKWELL_OVERLAP_SHARED_EXPERTS_OUTSIDE_SBO` | Run shared experts on an alternate stream when single batch overlap is enabled on GB200. When not setting this flag, shared experts and down gemm will be overlapped with DeepEP combine together. | `"false"` |
## MORI Configuration
| Environment Variable | Description | Default Value |
| --- | --- | --- |
| `SGLANG_MORI_FP8_DISP` | Use FP8 for dispatch | `"false"` |
| `SGLANG_MORI_NUM_MAX_DISPATCH_TOKENS_PER_RANK` | Maximum number of dispatch tokens per rank for MORI-EP buffer allocation | `4096` |
| `SGLANG_MORI_QP_PER_TRANSFER` | Number of RDMA Queue Pairs (QPs) used per transfer operation | `1` |
| `SGLANG_MORI_POST_BATCH_SIZE` | Number of RDMA work requests posted in a single batch to each QP | `-1` |
| `SGLANG_MORI_NUM_WORKERS` | Number of worker threads in the RDMA executor thread pool | `1` |
## NSA Backend Configuration (For DeepSeek V3.2)
<!-- # Environment variable to control mtp precomputing of metadata for multi-step speculative decoding -->
+1 -1
View File
@@ -255,8 +255,8 @@ def pad_sequence_with_mask(
dtype=torch.bool,
)
BLOCK_M = 32
BLOCK_D = triton.next_power_of_2(hidden_dim)
BLOCK_M = triton.next_power_of_2(max_len)
grid = (
B,
+143 -1
View File
@@ -19,17 +19,20 @@ from sglang.srt.layers.moe.fused_moe_triton.layer import (
FusedMoE,
moe_forward_piecewise_cuda_graph_impl,
)
from sglang.srt.layers.moe.rocm_moe_utils import upscale
from sglang.srt.layers.moe.token_dispatcher.deepep import (
DeepEPLLCombineInput,
DeepEPNormalCombineInput,
)
from sglang.srt.layers.moe.token_dispatcher.moriep import MoriEPNormalCombineInput
from sglang.srt.layers.moe.topk import TopKOutput, TopKOutputChecker
from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.layers.quantization.compressed_tensors.compressed_tensors_moe import (
NPUCompressedTensorsW4A16Int4DynamicMoEMethod,
)
from sglang.srt.layers.quantization.fp8 import Fp8Config
from sglang.srt.layers.quantization.fp8 import Fp8Config, Fp8MoEMethod
from sglang.srt.layers.quantization.fp8_kernel import is_fp8_fnuz
from sglang.srt.layers.quantization.quark.quark_moe import QuarkW4A4MXFp4MoEMethod
from sglang.srt.layers.quantization.w4afp8 import W4AFp8Config, W4AFp8MoEMethod
from sglang.srt.utils import get_bool_env_var, is_hip, is_npu
@@ -545,7 +548,146 @@ class NpuFuseEPMoE(DeepEPMoE):
)
class MoriEPMoE(DeepEPMoE):
def __init__(
self,
num_experts: int,
top_k: int,
hidden_size: int,
intermediate_size: int,
layer_id: int,
num_fused_shared_experts: int = 0,
params_dtype: Optional[torch.dtype] = None,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
activation: str = "silu",
routed_scaling_factor: Optional[float] = None,
**kwargs,
):
super().__init__(
num_experts=num_experts,
top_k=top_k,
hidden_size=hidden_size,
intermediate_size=intermediate_size,
layer_id=layer_id,
num_fused_shared_experts=num_fused_shared_experts,
params_dtype=params_dtype,
quant_config=quant_config,
prefix=prefix,
activation=activation,
routed_scaling_factor=routed_scaling_factor,
**kwargs,
)
assert _use_aiter, "Mori need to be used together with aiter as of now"
self.expert_mask = torch.zeros(
(self.num_experts),
device=torch.cuda.current_device(),
dtype=torch.int32,
)
expert_start_idx = self.moe_ep_rank * self.num_local_experts
expert_end_idx = expert_start_idx + self.num_local_experts
self.expert_mask[expert_start_idx:expert_end_idx] = 1
def forward(
self,
hidden_states: torch.Tensor,
topk_output: TopKOutput,
forward_shared_experts=None,
alt_stream=None,
disable_sbo=False,
):
num_token = hidden_states.shape[0]
output_dtype = hidden_states.dtype
scale = None
is_fp8_quant = isinstance(self.quant_method, Fp8MoEMethod)
is_quark_w4a4 = isinstance(self.quant_method, QuarkW4A4MXFp4MoEMethod)
# dispatch
dispatch_output = self.dispatcher.dispatch(
hidden_states, topk_output
) # , scale=scale)
(
dispatch_a1,
dispatch_scale,
dispatch_ids,
dispatch_weights,
dispatch_recv_token_num,
) = dispatch_output
w13_weight = self.w13_weight
w2_weight = self.w2_weight
w13_scale = None
w2_scale = None
quant_type = QuantType.No
if not is_fp8_quant and dispatch_scale is not None:
dispatch_a1 = upscale(
dispatch_a1, dispatch_scale, dispatch_recv_token_num, output_dtype
)
dispatch_scale = None
if is_quark_w4a4:
if hasattr(torch, "float4_e2m1fn_x2"):
w13_weight = self.w13_weight.view(torch.float4_e2m1fn_x2)
w2_weight = self.w2_weight.view(torch.float4_e2m1fn_x2)
w13_scale = self.w13_weight_scale
w2_scale = self.w2_weight_scale
quant_type = QuantType.per_1x32
if hasattr(self.w13_weight, "is_shuffled"):
w13_weight.is_shuffled = True
w2_weight.is_shuffled = True
elif is_fp8_quant:
if hasattr(self, "w13_weight_scale_inv"):
w13_scale = self.w13_weight_scale_inv
if hasattr(self, "w2_weight_scale_inv"):
w2_scale = self.w2_weight_scale_inv
quant_type = QuantType.per_128x128
# [KK TODO] should to call the apply of quant method to handle fused moe
hidden_states = fused_moe(
hidden_states=dispatch_a1,
w1=w13_weight,
w2=w2_weight,
w1_scale=w13_scale,
w2_scale=w2_scale,
a1_scale=dispatch_scale,
topk_weight=dispatch_weights,
topk_ids=dispatch_ids,
quant_type=quant_type,
activation=(
ActivationType.Silu
if self.moe_runner_config.activation == "silu"
else ActivationType.Gelu
),
expert_mask=self.expert_mask,
num_local_tokens=dispatch_recv_token_num,
dtype=output_dtype,
)
combine_input_wrapper = MoriEPNormalCombineInput
combine_input = combine_input_wrapper(
hidden_states=hidden_states,
topk_ids=topk_output.topk_ids,
topk_weights=topk_output.topk_weights,
)
# combine
result = self.dispatcher.combine(combine_input)
return result[:num_token]
def get_moe_impl_class(quant_config: Optional[QuantizationConfig]):
# [TODO] kk, temporary solution
if get_moe_a2a_backend().is_mori():
return MoriEPMoE
if get_moe_a2a_backend().is_deepep() or get_moe_a2a_backend().is_mooncake():
return DeepEPMoE
if get_moe_a2a_backend().is_ascend_fuseep():
@@ -121,6 +121,19 @@ def create_moe_dispatcher(moe_runner_config: MoeRunnerConfig) -> BaseDispatcher:
hidden_size=moe_runner_config.hidden_size,
params_dtype=moe_runner_config.params_dtype,
)
elif a2a_backend.is_mori():
from sglang.srt.layers.moe.token_dispatcher import MoriEPDispatcher
return MoriEPDispatcher(
group=get_tp_group(),
router_topk=moe_runner_config.top_k,
permute_fusion=True,
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,
deepep_mode=get_deepep_mode(),
)
elif a2a_backend.is_flashinfer():
return FlashinferDispatcher(
group=get_tp_group().device_group,
@@ -5,6 +5,8 @@ from enum import IntEnum
from typing import Optional
import torch
import triton
import triton.language as tl
from sglang.srt.utils import get_bool_env_var, is_hip
from sglang.srt.utils.custom_op import register_custom_op
@@ -114,3 +116,74 @@ def rocm_fused_experts_tkw1(
)
else:
assert False, "This should not be called."
@triton.jit
def upscale_kernel(
A_ptr, # *fp16 / *fp32
scale_ptr, # *fp16 / *fp32
Out_ptr, # *fp16 / *fp32
M,
N,
recv_token_num,
stride_am,
stride_an,
stride_sm,
stride_sn,
stride_om,
stride_on,
BLOCK_N: tl.constexpr,
):
pid_m = tl.program_id(0) # row id
pid_n = tl.program_id(1) # block id along N
recv_token_num_val = tl.load(recv_token_num)
if pid_m >= recv_token_num_val:
return
# column offsets
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
mask = offs_n < N
# A[m, n]
a_ptrs = A_ptr + pid_m * stride_am + offs_n * stride_an
a = tl.load(a_ptrs, mask=mask, other=0.0)
# scale index: n // 128
scale_idx = offs_n // 128
s_ptrs = scale_ptr + pid_m * stride_sm + scale_idx * stride_sn
s = tl.load(s_ptrs, mask=mask, other=1.0)
out = a * s
out_ptrs = Out_ptr + pid_m * stride_om + offs_n * stride_on
tl.store(out_ptrs, out, mask=mask)
def upscale(hidden_state, hidden_state_scale, recv_token_num, output_dtype):
M, N = hidden_state.shape
Out = torch.empty_like(hidden_state, dtype=output_dtype)
BLOCK_N = 256
grid = (M, triton.cdiv(N, BLOCK_N))
upscale_kernel[grid](
hidden_state,
hidden_state_scale,
Out,
M,
N,
recv_token_num,
hidden_state.stride(0),
hidden_state.stride(1),
hidden_state_scale.stride(0),
hidden_state_scale.stride(1),
Out.stride(0),
Out.stride(1),
BLOCK_N=BLOCK_N,
)
return Out
@@ -26,6 +26,11 @@ from sglang.srt.layers.moe.token_dispatcher.mooncake import (
MooncakeDispatchOutput,
MooncakeEPDispatcher,
)
from sglang.srt.layers.moe.token_dispatcher.moriep import (
MoriEPDispatcher,
MoriEPNormalCombineInput,
MoriEPNormalDispatchOutput,
)
from sglang.srt.layers.moe.token_dispatcher.standard import (
StandardCombineInput,
StandardDispatcher,
@@ -46,6 +51,9 @@ __all__ = [
"MooncakeCombineInput",
"MooncakeDispatchOutput",
"MooncakeEPDispatcher",
"MoriEPNormalDispatchOutput",
"MoriEPNormalCombineInput",
"MoriEPDispatcher",
"StandardDispatcher",
"StandardDispatchOutput",
"StandardCombineInput",
@@ -0,0 +1,461 @@
from __future__ import annotations
import logging
from dataclasses import dataclass
from typing import TYPE_CHECKING, List, NamedTuple, Optional, Tuple
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 DeepEPMode
from sglang.srt.utils import get_bool_env_var, get_int_env_var, is_hip
if TYPE_CHECKING:
from sglang.srt.single_batch_overlap import CombineOverlapArgs
import mori
from enum import Enum, auto
from functools import lru_cache
import torch
from sglang.srt.distributed import (
get_moe_expert_parallel_rank,
get_moe_expert_parallel_world_size,
)
from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype
_is_hip = is_hip()
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
if _use_aiter:
from aiter import QuantType, get_hip_quant
logger = logging.getLogger(__name__)
class MoriEPNormalDispatchOutput(NamedTuple):
"""Mori EP dispatch output."""
hidden_states: torch.Tensor
hidden_states_scale: Optional[torch.Tensor]
topk_ids: torch.Tensor
topk_weights: torch.Tensor
num_recv_tokens_per_expert: List[int]
@property
def format(self) -> DispatchOutputFormat:
return DispatchOutputFormat.DEEPEP_NORMAL
assert isinstance(MoriEPNormalDispatchOutput, DispatchOutput)
class MoriEPNormalCombineInput(NamedTuple):
"""Mori EP combine input."""
hidden_states: torch.Tensor
topk_ids: torch.Tensor
topk_weights: torch.Tensor
@property
def format(self) -> CombineInputFormat:
return CombineInputFormat.DEEPEP_NORMAL
assert isinstance(MoriEPNormalCombineInput, CombineInput)
class EpMode(Enum):
INTRA_NODE = "intra_node"
INTER_NODE = "inter_node"
@dataclass(frozen=True)
class EpDispatchConfig:
kernel_type: mori.ops.EpDispatchCombineKernelType
warp_num_per_block: int
block_num: int
rdma_block_num: int
def get_ep_dispatch_configs():
import mori
return {
EpMode.INTRA_NODE: EpDispatchConfig(
kernel_type=mori.ops.EpDispatchCombineKernelType.IntraNode,
warp_num_per_block=16,
block_num=80,
rdma_block_num=0,
),
EpMode.INTER_NODE: EpDispatchConfig(
kernel_type=mori.ops.EpDispatchCombineKernelType.InterNodeV1,
warp_num_per_block=8,
block_num=64,
rdma_block_num=32,
),
}
# init_mori_op only needs do once in model initial stage
# use lru_cache to reuse the same mori_op instance to avoid the init overhead for mori
@lru_cache(maxsize=1)
def init_mori_op(
group,
router_topk,
num_experts,
num_local_experts,
hidden_size,
params_dtype,
num_max_dispatch_tokens_per_rank,
):
import mori
world_size = get_moe_expert_parallel_world_size()
rank = get_moe_expert_parallel_rank()
cpu_group = group.cpu_group
torch._C._distributed_c10d._register_process_group("mori", cpu_group)
mori.shmem.shmem_torch_process_group_init("mori")
logger.info(
f"[MORI init] {world_size=} {rank=} {hidden_size=} {params_dtype=} {num_max_dispatch_tokens_per_rank=} {num_local_experts=} {router_topk=}"
)
mode = EpMode.INTRA_NODE if world_size <= 8 else EpMode.INTER_NODE
cfg = get_ep_dispatch_configs()[mode]
kernel_type = cfg.kernel_type
warp_num_per_block = cfg.warp_num_per_block
block_num = cfg.block_num
rdma_block_num = cfg.rdma_block_num
mori_config = mori.ops.EpDispatchCombineConfig(
rank=rank,
world_size=world_size,
data_type=fp8_dtype,
hidden_dim=hidden_size,
scale_dim=(
hidden_size // 128
if get_bool_env_var("SGLANG_MORI_FP8_DISP", "False")
else 1
),
scale_type_size=torch.float32.itemsize,
max_token_type_size=params_dtype.itemsize,
max_num_inp_token_per_rank=num_max_dispatch_tokens_per_rank,
num_experts_per_rank=num_local_experts,
num_experts_per_token=router_topk,
warp_num_per_block=warp_num_per_block,
block_num=block_num,
kernel_type=kernel_type,
rdma_block_num=rdma_block_num,
num_qp_per_pe=2,
)
mori_op = mori.ops.EpDispatchCombineOp(mori_config)
return mori_op
class _MoriEPDispatcherImplBase:
def __init__(
self,
group: torch.distributed.ProcessGroup,
router_topk: int,
permute_fusion: bool,
num_experts: int,
num_local_experts: int,
hidden_size: int,
params_dtype: torch.dtype,
return_recv_hook: bool,
deepep_mode: DeepEPMode,
):
try:
import mori # noqa: F401
except ImportError:
raise ImportError("Mori EP is not installed. Please install.")
self.group = group
self.router_topk = router_topk
self.permute_fusion = permute_fusion
self.num_experts = num_experts
self.num_local_experts = num_local_experts
self.hidden_size = hidden_size
self.params_dtype = params_dtype
self.return_recv_hook = return_recv_hook
self.deepep_mode = deepep_mode
self.num_max_dispatch_tokens_per_rank = get_int_env_var(
"SGLANG_MORI_NUM_MAX_DISPATCH_TOKENS_PER_RANK", 4096
)
self.mori_op = init_mori_op(
self.group,
self.router_topk,
self.num_experts,
self.num_local_experts,
self.hidden_size,
self.params_dtype,
num_max_dispatch_tokens_per_rank=self.num_max_dispatch_tokens_per_rank,
)
def dispatch_a(
self,
hidden_states: torch.Tensor,
topk_output: TopKOutput,
):
raise NotImplementedError
def dispatch_b(self, *args, **kwargs):
raise NotImplementedError
def combine_a(
self,
hidden_states: torch.Tensor,
topk_ids: torch.Tensor,
topk_weights: torch.Tensor,
overlap_args: Optional[CombineOverlapArgs] = None,
):
raise NotImplementedError
def combine_b(self, *args, **kwargs):
raise NotImplementedError
def _get_buffer(self):
raise NotImplementedError
class _MoriEPDispatcherImplNormal(_MoriEPDispatcherImplBase):
def __init__(self, **kwargs):
super().__init__(**kwargs)
self.quant_config = {}
# [kk TODO] need to support mxfp4 type
self.quant_func = get_hip_quant(QuantType.per_1x128)
def dispatch_a(
self,
hidden_states: torch.Tensor,
topk_output: TopKOutput,
):
topk_weights, topk_ids = topk_output.topk_weights, topk_output.topk_ids
return (
hidden_states,
topk_weights,
topk_ids,
)
def dispatch_b(
self,
hidden_states,
topk_weights,
topk_ids,
):
num_token = hidden_states.shape[0]
scale = None
fp8_dispatch = get_bool_env_var("SGLANG_MORI_FP8_DISP", "False")
if fp8_dispatch:
# FP8 quant
if num_token > 0:
# NOTE: aiter is able to handle token=0 case in UT. But for some reason it failed at e2e case. Root cause TBD.
hidden_states, scale = self.quant_func(
hidden_states, quant_dtype=fp8_dtype
)
else:
hidden_states = torch.empty(
hidden_states.shape, dtype=fp8_dtype, device=hidden_states.device
)
scale = torch.empty(
(0, self.hidden_size // 128),
dtype=torch.float32,
device=hidden_states.device,
)
(
packed_recv_hidden,
recv_topk_weights,
recv_scales,
recv_topk_ids,
packed_recv_count,
) = self._dispatch_core(hidden_states, topk_weights, topk_ids, scale)
return MoriEPNormalDispatchOutput(
packed_recv_hidden,
recv_scales,
recv_topk_ids,
recv_topk_weights,
packed_recv_count,
)
def _dispatch_core(
self,
hidden_states: torch.Tensor,
topk_weights: torch.Tensor,
topk_ids: torch.Tensor,
scale: Optional[torch.Tensor] = None,
):
(
packed_recv_hidden,
recv_topk_weights,
recv_scales,
recv_topk_ids,
packed_recv_count,
) = self.mori_op.dispatch(hidden_states, topk_weights, scale, topk_ids)
# TODO(billishyahao): EPLB
# get_global_expert_distribution_recorder().on_deepep_dispatch_normal(
return (
packed_recv_hidden,
recv_topk_weights,
recv_scales,
recv_topk_ids,
packed_recv_count,
)
def combine_a(
self,
hidden_states: torch.Tensor,
topk_ids: torch.Tensor,
topk_weights: torch.Tensor,
overlap_args: Optional[CombineOverlapArgs] = None,
):
previous_event = None
return hidden_states, topk_ids, topk_weights, previous_event
def combine_b(self, hidden_states, topk_ids, topk_weights, previous_event):
hidden_states = self._combine_core(hidden_states, topk_ids, topk_weights)
return hidden_states
def _combine_core(
self,
hidden_states: torch.Tensor,
topk_ids: torch.Tensor,
topk_weights: torch.Tensor,
):
combined_hidden_states = self.mori_op.combine(hidden_states, None, topk_ids)
return combined_hidden_states[0]
def set_quant_config(self, quant_config: dict):
self.quant_config = quant_config
@dataclass
class _Stage(Enum):
INITIAL = auto()
AFTER_DISPATCH_A = auto()
AFTER_DISPATCH_B = auto()
AFTER_COMBINE_A = auto()
class MoriEPDispatcher(BaseDispatcher):
def __init__(
self,
group: torch.distributed.ProcessGroup,
router_topk: int,
permute_fusion: bool = False,
num_experts: int = None,
num_local_experts: int = None,
hidden_size: int = None,
params_dtype: torch.dtype = None,
deepep_mode: DeepEPMode = DeepEPMode.AUTO,
async_finish: bool = False,
return_recv_hook: bool = False,
):
self.deepep_mode = deepep_mode
if self.deepep_mode.enable_normal():
self._normal_dispatcher = _MoriEPDispatcherImplNormal(
group=group,
router_topk=router_topk,
permute_fusion=permute_fusion,
num_experts=num_experts,
num_local_experts=num_local_experts,
hidden_size=hidden_size,
params_dtype=params_dtype,
return_recv_hook=return_recv_hook,
deepep_mode=deepep_mode,
)
if self.deepep_mode.enable_low_latency():
raise NotImplementedError
self._stage = _Stage.INITIAL
def dispatch(self, *args, **kwargs) -> DispatchOutput:
self.dispatch_a(*args, **kwargs)
ret = self.dispatch_b()
return ret
def dispatch_a(
self,
hidden_states: torch.Tensor,
topk_output: TopKOutput,
):
self._update_stage(_Stage.INITIAL, _Stage.AFTER_DISPATCH_A)
inner_state = self._get_impl().dispatch_a(
hidden_states=hidden_states,
topk_output=topk_output,
)
self._dispatch_intermediate_state = inner_state
def dispatch_b(self):
self._update_stage(_Stage.AFTER_DISPATCH_A, _Stage.AFTER_DISPATCH_B)
inner_state = self._dispatch_intermediate_state
del self._dispatch_intermediate_state
return self._get_impl().dispatch_b(*inner_state)
def combine(
self,
combine_input: CombineInput,
overlap_args: Optional[CombineOverlapArgs] = None,
) -> Tuple:
self.combine_a(combine_input, overlap_args)
ret = self.combine_b()
return ret
def combine_a(
self,
combine_input: CombineInput,
overlap_args: Optional[CombineOverlapArgs] = None,
):
hidden_states, topk_ids, topk_weights = combine_input
self._update_stage(_Stage.AFTER_DISPATCH_B, _Stage.AFTER_COMBINE_A)
inner_state = self._get_impl().combine_a(
hidden_states=hidden_states,
topk_ids=topk_ids,
topk_weights=topk_weights,
overlap_args=overlap_args,
)
self._combine_intermediate_state = inner_state
def combine_b(self):
self._update_stage(_Stage.AFTER_COMBINE_A, _Stage.INITIAL)
inner_state = self._combine_intermediate_state
del self._combine_intermediate_state
return self._get_impl().combine_b(*inner_state)
def _get_impl(self) -> _MoriEPDispatcherImplBase:
is_extend_in_batch = get_is_extend_in_batch()
resolved_deepep_mode = self.deepep_mode.resolve(is_extend_in_batch)
if resolved_deepep_mode == DeepEPMode.NORMAL:
return self._normal_dispatcher
elif resolved_deepep_mode == DeepEPMode.LOW_LATENCY:
raise NotImplementedError
else:
raise ValueError(f"Invalid deepep_mode: {self.deepep_mode}")
def _update_stage(self, old_stage, new_stage):
assert self._stage == old_stage
self._stage = new_stage
def set_quant_config(self, quant_config: dict):
if self.deepep_mode.enable_low_latency():
raise NotImplementedError
if self.deepep_mode.enable_normal():
self._normal_dispatcher.set_quant_config(quant_config)
+4
View File
@@ -23,6 +23,7 @@ class MoeA2ABackend(Enum):
NONE = "none"
DEEPEP = "deepep"
MOONCAKE = "mooncake"
MORI = "mori"
ASCEND_FUSEEP = "ascend_fuseep"
FLASHINFER = "flashinfer"
@@ -50,6 +51,9 @@ class MoeA2ABackend(Enum):
def is_ascend_fuseep(self):
return self == MoeA2ABackend.ASCEND_FUSEEP
def is_mori(self):
return self == MoeA2ABackend.MORI
class MoeRunnerBackend(Enum):
@@ -218,6 +218,8 @@ class SchedulerMetricsMixin:
f += f"#prealloc-req: {len(self.disagg_prefill_bootstrap_queue.queue)}, "
f += f"#inflight-req: {len(self.disagg_prefill_inflight_queue)}, "
f += f"input throughput (token/s): {self.last_input_throughput:.2f}, "
else:
f += f"input throughput (token/s): {self.last_input_throughput:.2f}, "
logger.info(f)
+22 -4
View File
@@ -467,6 +467,7 @@ class DeepseekV2MoE(nn.Module):
dict(tp_rank=0, tp_size=1)
if get_moe_a2a_backend().is_deepep()
or get_moe_a2a_backend().is_mooncake()
or get_moe_a2a_backend().is_mori()
or get_moe_a2a_backend().is_ascend_fuseep()
or get_moe_a2a_backend().is_flashinfer()
or should_use_flashinfer_cutlass_moe_fp4_allgather()
@@ -510,6 +511,7 @@ class DeepseekV2MoE(nn.Module):
if (
get_moe_a2a_backend().is_deepep()
or get_moe_a2a_backend().is_mooncake()
or get_moe_a2a_backend().is_mori()
or get_moe_a2a_backend().is_ascend_fuseep()
):
# TODO: we will support tp < ep in the future
@@ -530,6 +532,7 @@ class DeepseekV2MoE(nn.Module):
self._enable_a2a_moe = (
get_moe_a2a_backend().is_deepep()
or get_moe_a2a_backend().is_mooncake()
or get_moe_a2a_backend().is_mori()
or get_moe_a2a_backend().is_ascend_fuseep()
or get_moe_a2a_backend().is_flashinfer()
)
@@ -945,13 +948,17 @@ class DeepseekV2MoE(nn.Module):
torch.cuda.current_stream().wait_event(shared_event)
if shared_output is not None:
x = shared_output
if self.experts.should_fuse_routed_scaling_factor_in_topk:
# aiter moe call will handle routed_scaling_factor in the function
# so add _use_aiter condition to eliminate to use self.routed_scaling_factor in add_ call
if self.experts.should_fuse_routed_scaling_factor_in_topk or _use_aiter:
x.add_(final_hidden_states)
else:
x.add_(final_hidden_states, alpha=self.routed_scaling_factor)
final_hidden_states = x
else:
if not self.experts.should_fuse_routed_scaling_factor_in_topk:
if not (
self.experts.should_fuse_routed_scaling_factor_in_topk or _use_aiter
):
final_hidden_states *= self.routed_scaling_factor
return final_hidden_states
@@ -2576,7 +2583,16 @@ class DeepseekV2Model(nn.Module):
allocate_size = 0
for i in range(len(self.layers)):
if isinstance(self.layers[i].mlp, DeepseekV2MoE):
tp_size = get_tensor_model_parallel_world_size()
# tp_size = get_tensor_model_parallel_world_size()
a2a_backend = get_moe_a2a_backend()
is_a2a_moe = (
a2a_backend.is_deepep()
or a2a_backend.is_mori()
or a2a_backend.is_mooncake()
)
tp_size = (
1 if is_a2a_moe else get_tensor_model_parallel_world_size()
)
intermediate_size = (
config.moe_intermediate_size * config.n_shared_experts
)
@@ -2842,7 +2858,9 @@ class DeepseekV2ForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin):
not _is_hip or torch.cuda.get_device_capability("cuda") < (9, 4)
):
disable_reason = "Only Deepseek V3/R1 on AMD-platform with capability >= gfx942(MI30x) can use shared experts fusion optimization under expert parallelism."
elif disable_reason is None and get_moe_a2a_backend().is_deepep():
elif disable_reason is None and (
get_moe_a2a_backend().is_deepep() or get_moe_a2a_backend().is_mori()
):
disable_reason = "Deepseek V3/R1 can not use shared experts fusion optimization under deepep expert parallelism."
elif self.quant_config and self.quant_config.get_name() == "w4afp8":
disable_reason = "Deepseek V3/R1 W4AFP8 model uses different quant method for routed experts and shared experts."
+22 -2
View File
@@ -43,6 +43,7 @@ from sglang.srt.utils.common import (
get_device_memory_capacity,
get_device_name,
get_device_sm,
get_int_env_var,
get_quantization_config,
is_blackwell_supported,
is_cuda,
@@ -182,7 +183,14 @@ MOE_RUNNER_BACKEND_CHOICES = [
"cutlass",
]
MOE_A2A_BACKEND_CHOICES = ["none", "deepep", "mooncake", "ascend_fuseep", "flashinfer"]
MOE_A2A_BACKEND_CHOICES = [
"none",
"deepep",
"mooncake",
"mori",
"ascend_fuseep",
"flashinfer",
]
FP8_GEMM_RUNNER_BACKEND_CHOICES = [
"auto",
@@ -475,7 +483,7 @@ class ServerArgs:
# Expert parallelism
ep_size: int = 1
moe_a2a_backend: Literal[
"none", "deepep", "mooncake", "ascend_fuseep", "flashinfer"
"none", "deepep", "mooncake", "mori", "ascend_fuseep", "flashinfer"
] = "none"
moe_runner_backend: str = "auto"
flashinfer_mxfp4_moe_precision: Literal["default", "bf16"] = "default"
@@ -2076,6 +2084,18 @@ class ServerArgs:
"flashinfer_cutlass"
], "Flashinfer MoE A2A is only supported with flashinfer_cutlass moe runner backend"
if self.moe_a2a_backend == "mori":
self.ep_size = self.tp_size
self.deepep_mode = "normal"
logger.warning("auto set deepep_mode=`normal` for MORI EP")
logger.warning(
f"MoRI MoE is enabled. The expert parallel size is adjusted to be the same as the tensor parallel size[{self.tp_size}]."
)
assert (self.chunked_prefill_size) <= get_int_env_var(
"SGLANG_MORI_NUM_MAX_DISPATCH_TOKENS_PER_RANK", 4096
), "SGLANG_MORI_NUM_MAX_DISPATCH_TOKENS_PER_RANK (default 4096) must be larger or equal to chunked_prefill_size"
def _handle_eplb_and_dispatch(self):
if self.enable_eplb and (self.expert_distribution_recorder_mode is None):
self.expert_distribution_recorder_mode = "stat"
+1
View File
@@ -126,6 +126,7 @@ DEFAULT_MODEL_NAME_FOR_TEST_LOCAL_ATTENTION = (
DEFAULT_SMALL_EMBEDDING_MODEL_NAME_FOR_TEST = "Alibaba-NLP/gte-Qwen2-1.5B-instruct"
DEFAULT_REASONING_MODEL_NAME_FOR_TEST = "deepseek-ai/DeepSeek-R1-Distill-Qwen-7B"
DEFAULT_DEEPEP_MODEL_NAME_FOR_TEST = "deepseek-ai/DeepSeek-V3-0324"
DEFAULT_DEEPEP_MODEL_NAME_FOR_TEST_NEXTN = "lmsys/DeepSeek-V3-NextN"
DEFAULT_AWQ_MOE_MODEL_NAME_FOR_TEST = (
"hugging-quants/Mixtral-8x7B-Instruct-v0.1-AWQ-INT4"
)
+170
View File
@@ -0,0 +1,170 @@
import os
import unittest
from types import SimpleNamespace
from sglang.srt.utils import kill_process_tree
from sglang.test.few_shot_gsm8k import run_eval as run_eval_few_shot_gsm8k
from sglang.test.test_utils import (
DEFAULT_DEEPEP_MODEL_NAME_FOR_TEST,
DEFAULT_DEEPEP_MODEL_NAME_FOR_TEST_NEXTN,
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
DEFAULT_URL_FOR_TEST,
CustomTestCase,
popen_launch_server,
)
class TestPureDP(CustomTestCase):
@classmethod
def setUpClass(cls):
cls.model = DEFAULT_DEEPEP_MODEL_NAME_FOR_TEST
cls.base_url = DEFAULT_URL_FOR_TEST
other_args = [
"--tp-size",
"8",
"--ep-size",
"8",
"--dp-size",
"8",
"--enable-dp-attention",
"--moe-a2a-backend",
"mori",
"--trust-remote-code",
"--load-balance-method",
"round_robin",
"--moe-dense-tp-size",
"1",
"--enable-dp-lm-head",
"--mem-fraction-static",
"0.6",
"--chunked-prefill-size",
"131072",
"--max-running-requests",
"128",
"--context-length",
"12288",
"--attention-backend",
"aiter",
]
env = dict(os.environ)
env["SGLANG_USE_AITER"] = "1"
env["SGLANG_MORI_FP8_DISP"] = "True"
env["SGLANG_MORI_NUM_MAX_DISPATCH_TOKENS_PER_RANK"] = "16384"
env["MORI_SHMEM_MODE"] = "ISOLATION" # avoid out of symmetric heap memory
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH * 5,
other_args=other_args,
env=env,
)
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
def test_gsm8k(
self,
):
args = SimpleNamespace(
num_shots=5,
data_path=None,
num_questions=200,
max_new_tokens=512,
parallel=128,
host="http://127.0.0.1",
port=int(self.base_url.split(":")[-1]),
)
metrics = run_eval_few_shot_gsm8k(args)
print(f"{metrics=}")
self.assertGreater(metrics["accuracy"], 0.935)
class TestMTP(CustomTestCase):
@classmethod
def setUpClass(cls):
cls.model = DEFAULT_DEEPEP_MODEL_NAME_FOR_TEST
cls.base_url = DEFAULT_URL_FOR_TEST
other_args = [
"--tp-size",
"8",
"--ep-size",
"8",
"--dp-size",
"8",
"--enable-dp-attention",
"--moe-a2a-backend",
"mori",
"--trust-remote-code",
"--load-balance-method",
"round_robin",
"--moe-dense-tp-size",
"1",
"--enable-dp-lm-head",
"--mem-fraction-static",
"0.6",
"--chunked-prefill-size",
"131072",
"--max-running-requests",
"128",
"--context-length",
"12288",
"--attention-backend",
"aiter",
"--speculative-algo",
"EAGLE",
"--speculative-draft-model-path",
DEFAULT_DEEPEP_MODEL_NAME_FOR_TEST_NEXTN,
"--speculative-num-steps",
"1",
"--speculative-eagle-topk",
"1",
"--speculative-num-draft-tokens",
"2",
"--cuda-graph-max-bs",
"32",
]
env = dict(os.environ)
env["SGLANG_USE_AITER"] = "1"
env["SGLANG_MORI_FP8_DISP"] = "True"
env["SGLANG_MORI_NUM_MAX_DISPATCH_TOKENS_PER_RANK"] = "16384"
env["MORI_SHMEM_MODE"] = "ISOLATION" # avoid out of symmetric heap memory
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH * 5,
other_args=other_args,
env=env,
)
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
def test_gsm8k(
self,
):
args = SimpleNamespace(
num_shots=5,
data_path=None,
num_questions=200,
max_new_tokens=512,
parallel=128,
host="http://127.0.0.1",
port=int(self.base_url.split(":")[-1]),
)
metrics = run_eval_few_shot_gsm8k(args)
print(f"{metrics=}")
self.assertGreater(metrics["accuracy"], 0.935)
if __name__ == "__main__":
unittest.main()
+1
View File
@@ -64,6 +64,7 @@ suites = {
TestFile(
"models/test_qwen3_next_models_pcg.py"
), # Disabled: intermittent failures, see #17039
TestFile("ep/test_moriep_small.py"),
],
}