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:
co-authored by
Duyi-Wang
billishyahao
HaiShaw
parent
673dc09d9b
commit
f1384f5293
@@ -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
|
||||
|
||||
|
||||
@@ -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 -->
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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."
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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"
|
||||
)
|
||||
|
||||
@@ -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()
|
||||
@@ -64,6 +64,7 @@ suites = {
|
||||
TestFile(
|
||||
"models/test_qwen3_next_models_pcg.py"
|
||||
), # Disabled: intermittent failures, see #17039
|
||||
TestFile("ep/test_moriep_small.py"),
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user