[Feature] Integrate Elastic NIXL-EP into SGLang (#19248)

Signed-off-by: Barak Biber <bbiber@nvidia.com>
Signed-off-by: Yoray Zack <yorayz@nvidia.com>
Signed-off-by: Itay Alroy <ialroy@nvidia.com>
Co-authored-by: Barak Biber <bbiber@nvidia.com>
This commit is contained in:
Yoray Zack
2026-03-11 11:37:43 +02:00
committed by GitHub
parent 680d9d98e4
commit 9991debde3
18 changed files with 752 additions and 17 deletions

View File

@@ -747,7 +747,11 @@ 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():
if (
get_moe_a2a_backend().is_deepep()
or get_moe_a2a_backend().is_mooncake()
or get_moe_a2a_backend().is_nixl()
):
return DeepEPMoE
if get_moe_a2a_backend().is_ascend_fuseep():
return NpuFuseEPMoE

View File

@@ -95,7 +95,12 @@ def create_moe_dispatcher(moe_runner_config: MoeRunnerConfig) -> BaseDispatcher:
a2a_backend = get_moe_a2a_backend()
if a2a_backend.is_none():
return StandardDispatcher(moe_runner_config)
elif a2a_backend.is_deepep() or a2a_backend.is_mooncake() or a2a_backend.is_mori():
elif (
a2a_backend.is_deepep()
or a2a_backend.is_mooncake()
or a2a_backend.is_mori()
or a2a_backend.is_nixl()
):
return MaybeTboDeepEPDispatcher(
group=(
get_tp_group().device_group

View File

@@ -33,6 +33,11 @@ from sglang.srt.layers.moe.token_dispatcher.moriep import (
MoriEPNormalCombineInput,
MoriEPNormalDispatchOutput,
)
from sglang.srt.layers.moe.token_dispatcher.nixl import (
NixlEPCombineInput,
NixlEPDispatcher,
NixlEPDispatchOutput,
)
from sglang.srt.layers.moe.token_dispatcher.standard import (
StandardCombineInput,
StandardDispatcher,
@@ -58,6 +63,9 @@ __all__ = [
"MoriEPLLDispatchOutput",
"MoriEPLLCombineInput",
"MoriEPDispatcher",
"NixlEPCombineInput",
"NixlEPDispatchOutput",
"NixlEPDispatcher",
"StandardDispatcher",
"StandardDispatchOutput",
"StandardCombineInput",

View File

@@ -0,0 +1,465 @@
from __future__ import annotations
import logging
from enum import Enum, auto
from typing import Optional
import torch
import torch.distributed as dist
from sglang.srt.distributed.utils import get_global_tcp_store
from sglang.srt.elastic_ep.elastic_ep import ElasticEPStateManager
from sglang.srt.environ import envs
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
from sglang.srt.layers import deep_gemm_wrapper
from sglang.srt.layers.dp_attention import get_is_extend_in_batch
from sglang.srt.layers.moe.token_dispatcher.base import (
BaseDispatcher,
CombineInput,
DispatchOutput,
)
from sglang.srt.layers.moe.token_dispatcher.deepep import (
DeepEPLLCombineInput,
DeepEPLLDispatchOutput,
)
from sglang.srt.layers.moe.topk import TopKOutput
from sglang.srt.layers.moe.utils import DeepEPMode
try:
from nixl_ep import Buffer
use_nixl = True
except ImportError:
use_nixl = False
logger = logging.getLogger(__name__)
NixlEPDispatchOutput = DeepEPLLDispatchOutput
NixlEPCombineInput = DeepEPLLCombineInput
class NixlEPBuffer:
_buffer = None
_hidden_size: Optional[int] = None
_num_max_dispatch_tokens_per_rank: Optional[int] = None
_num_experts: Optional[int] = None
_num_local_experts: Optional[int] = None
@classmethod
def get_nixl_buffer(
cls,
group: dist.ProcessGroup,
hidden_size: int,
deepep_mode: DeepEPMode,
num_max_dispatch_tokens_per_rank: int = -1,
num_experts: int = -1,
num_local_experts: int = -1,
):
if cls._buffer is not None:
return cls._buffer
cls._hidden_size = hidden_size
cls._num_max_dispatch_tokens_per_rank = num_max_dispatch_tokens_per_rank
cls._num_experts = num_experts
cls._num_local_experts = num_local_experts
num_rdma_bytes = 0
if deepep_mode.enable_normal():
raise NotImplementedError("Normal mode is not supported for Nixl EP yet.")
if deepep_mode.enable_low_latency():
assert num_max_dispatch_tokens_per_rank != -1
assert num_experts != -1 and num_experts % group.size() == 0
num_rdma_bytes = Buffer.get_rdma_size_hint(
num_max_dispatch_tokens_per_rank,
hidden_size,
group.size(),
num_experts,
)
rank = dist.get_rank(group)
world_size = dist.get_world_size(group)
# Get the global TCPStore for coordination
tcp_store = get_global_tcp_store()
if tcp_store is None:
raise RuntimeError(
"Global TCPStore is not initialized. "
"Make sure init_distributed_environment was called before using NIXL EP."
)
logger.info(
f"Using NIXL EP (world_size={world_size}, rank={rank}, "
f"num_experts={cls._num_experts}, num_experts_per_rank={cls._num_local_experts}) "
)
cls._buffer = Buffer(
rank=rank,
tcp_store_group=tcp_store,
)
cls._buffer.update_memory_buffers(
num_ranks=world_size,
num_experts_per_rank=cls._num_local_experts,
num_rdma_bytes=num_rdma_bytes,
)
all_ranks = list(range(world_size))
cls._buffer.connect_ranks(all_ranks)
return cls._buffer
@classmethod
def clean_buffer(cls):
cls._buffer.clean_buffer(
cls._num_max_dispatch_tokens_per_rank,
cls._hidden_size,
cls._num_experts,
)
class _NixlEPDispatcherImplBase:
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,
deepep_mode: DeepEPMode,
):
if not use_nixl:
raise ImportError(
"NixlEP is not installed. Please install NixlEP package from "
"https://github.com/ai-dynamo/nixl."
)
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.deepep_mode = deepep_mode
self.num_max_dispatch_tokens_per_rank = (
envs.SGLANG_NIXL_EP_NUM_MAX_DISPATCH_TOKENS_PER_RANK.get()
)
# NixlEP internode_ll dispatch uses FINISHED_SUM_TAG=1024
# and the logic requires num-tokens-sent-from-one-rank-to-another-rank less than it
assert self.num_max_dispatch_tokens_per_rank <= 1024
elastic_state = ElasticEPStateManager.instance()
self.active_ranks = (
elastic_state.active_ranks if elastic_state is not None else None
)
self._mask_buffer = (
torch.zeros_like(self.active_ranks)
if self.active_ranks is not None
else None
)
self.handle = None
self.quant_config = None
self.overlap_args = None
self.meta_overlap_args = None
def set_quant_config(self, quant_config: dict) -> None:
self.quant_config = quant_config
def set_overlap_args(self, combine_overlap_args, meta_overlap_args) -> None:
self.overlap_args = combine_overlap_args
self.meta_overlap_args = meta_overlap_args
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,
):
raise NotImplementedError
def combine_b(self, *args, **kwargs):
raise NotImplementedError
def _get_buffer(self):
raise NotImplementedError
class _NixlEPDispatcherImpl(_NixlEPDispatcherImplBase):
def __init__(self, return_recv_hook: bool, **kwargs):
super().__init__(**kwargs)
"""
num_max_dispatch_tokens_per_rank: the actual batch size in the decoding engine should be less than 256
https://github.com/ai-dynamo/nixl
"""
self.return_recv_hook = return_recv_hook
self.device_module = torch.get_device_module()
def dispatch_a(
self,
hidden_states: torch.Tensor,
topk_output: TopKOutput,
):
buffer = self._get_buffer()
topk_weights, topk_ids = topk_output.topk_weights, topk_output.topk_ids
topk_ids = topk_ids.to(torch.int64)
expected_m = (
hidden_states.shape[0] * buffer.group_size * topk_ids.shape[1]
+ self.num_experts
) // self.num_experts
hidden_states, masked_m, event, hook = self._dispatch_core(
hidden_states,
topk_ids,
)
return (
hidden_states,
topk_ids,
topk_weights,
masked_m,
expected_m,
event,
hook,
)
def dispatch_b(
self,
hidden_states,
topk_ids,
topk_weights,
masked_m,
expected_m,
event,
hook,
):
hook() if self.return_recv_hook else event.current_stream_wait()
get_global_expert_distribution_recorder().on_deepep_dispatch_low_latency(
masked_m
)
if isinstance(hidden_states, tuple):
hidden_states, hidden_states_scale = hidden_states
else:
hidden_states_scale = None
nixl_output = NixlEPDispatchOutput(
hidden_states,
hidden_states_scale,
topk_ids,
topk_weights,
masked_m,
expected_m,
)
return nixl_output
def _dispatch_core(
self,
hidden_states: torch.Tensor,
topk_idx: torch.Tensor,
):
use_fp8 = not envs.SGLANG_NIXL_EP_BF16_DISPATCH.get()
buffer = self._get_buffer()
packed_recv_hidden, self.packed_recv_count, self.handle, event, hook = (
buffer.dispatch(
hidden_states,
topk_idx,
self.num_max_dispatch_tokens_per_rank,
self.num_experts,
use_fp8=use_fp8,
async_finish=not self.return_recv_hook,
return_recv_hook=self.return_recv_hook,
round_scale=deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM
and deep_gemm_wrapper.DEEPGEMM_BLACKWELL,
use_ue8m0=deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM
and deep_gemm_wrapper.DEEPGEMM_BLACKWELL,
)
)
return packed_recv_hidden, self.packed_recv_count, event, hook
def combine_a(
self,
hidden_states: torch.Tensor,
topk_ids: torch.Tensor,
topk_weights: torch.Tensor,
):
hidden_states, event, hook = self._combine_core(
hidden_states,
topk_ids,
topk_weights,
)
return hidden_states, event, hook
def combine_b(self, hidden_states, event, hook):
hook() if self.return_recv_hook else event.current_stream_wait()
return hidden_states
def _combine_core(
self,
hidden_states: torch.Tensor,
topk_ids: torch.Tensor,
topk_weights: torch.Tensor,
):
buffer = self._get_buffer()
combined_hidden_states, event, hook = buffer.combine(
x=hidden_states,
topk_idx=topk_ids,
topk_weights=topk_weights,
handle=self.handle,
async_finish=not self.return_recv_hook,
return_recv_hook=self.return_recv_hook,
)
if self._mask_buffer is not None:
buffer.query_mask_buffer(self._mask_buffer)
self.active_ranks.copy_(1 - self._mask_buffer)
self.packed_recv_count = self.handle = None
return combined_hidden_states, event, hook
def _get_buffer(self):
return NixlEPBuffer.get_nixl_buffer(
self.group,
self.hidden_size,
self.deepep_mode,
self.num_max_dispatch_tokens_per_rank,
self.num_experts,
self.num_local_experts,
)
class _Stage(Enum):
INITIAL = auto()
AFTER_DISPATCH_A = auto()
AFTER_DISPATCH_B = auto()
AFTER_COMBINE_A = auto()
class NixlEPDispatcher(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.LOW_LATENCY,
async_finish: bool = False,
return_recv_hook: bool = False,
):
self.deepep_mode = deepep_mode
common_kwargs = dict(
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,
deepep_mode=deepep_mode,
)
if self.deepep_mode.enable_low_latency():
self._low_latency_dispatcher = _NixlEPDispatcherImpl(
return_recv_hook=return_recv_hook,
**common_kwargs,
)
if self.deepep_mode.enable_normal():
raise NotImplementedError("Normal mode is not supported for Nixl EP yet.")
self._stage = _Stage.INITIAL
def dispatch(
self,
hidden_states: torch.Tensor,
topk_output: TopKOutput,
) -> DispatchOutput:
self.dispatch_a(hidden_states=hidden_states, topk_output=topk_output)
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,
) -> torch.Tensor:
self.combine_a(combine_input)
ret = self.combine_b()
return ret
def combine_a(
self,
combine_input: CombineInput,
):
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,
)
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) -> _NixlEPDispatcherImplBase:
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:
raise NotImplementedError("Normal mode is not supported for Nixl EP yet.")
elif resolved_deepep_mode == DeepEPMode.LOW_LATENCY:
return self._low_latency_dispatcher
else:
raise ValueError(f"Invalid deepep_mode: {self.deepep_mode}")
def set_quant_config(self, quant_config: dict):
super().set_quant_config(quant_config)
if self.deepep_mode.enable_low_latency():
self._low_latency_dispatcher.set_quant_config(quant_config)
def set_overlap_args(self, combine_overlap_args, meta_overlap_args):
super().set_overlap_args(combine_overlap_args, meta_overlap_args)
if self.deepep_mode.enable_low_latency():
self._low_latency_dispatcher.set_overlap_args(
combine_overlap_args, meta_overlap_args
)
def _update_stage(self, old_stage, new_stage):
assert self._stage == old_stage
self._stage = new_stage

View File

@@ -22,6 +22,7 @@ class MoeA2ABackend(Enum):
NONE = "none"
DEEPEP = "deepep"
MOONCAKE = "mooncake"
NIXL = "nixl"
MORI = "mori"
ASCEND_FUSEEP = "ascend_fuseep"
FLASHINFER = "flashinfer"
@@ -44,6 +45,9 @@ class MoeA2ABackend(Enum):
def is_mooncake(self):
return self == MoeA2ABackend.MOONCAKE
def is_nixl(self):
return self == MoeA2ABackend.NIXL
def is_flashinfer(self):
return self == MoeA2ABackend.FLASHINFER

View File

@@ -748,7 +748,9 @@ class Fp8MoEMethod(FusedMoEMethodBase):
return True
if moe_runner_backend.is_auto():
return deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM and (
get_moe_a2a_backend().is_deepep() or get_moe_a2a_backend().is_mooncake()
get_moe_a2a_backend().is_deepep()
or get_moe_a2a_backend().is_mooncake()
or get_moe_a2a_backend().is_nixl()
)
return False