[refactor] Move the EP dispatcher and fusion-workspace manager state onto ctx.resources (#30489)

This commit is contained in:
Cheng Wan
2026-07-09 02:08:35 -07:00
committed by GitHub
parent bc5d376c2c
commit 65b14881c5
7 changed files with 182 additions and 70 deletions
+1 -1
View File
@@ -141,7 +141,7 @@ def _maybe_create_message_queue(group) -> None:
def _refresh_ep_members() -> None: def _refresh_ep_members() -> None:
from sglang.srt.layers.moe.token_dispatcher.mooncake import EPBuffer from sglang.srt.layers.moe.token_dispatcher.mooncake import EPBuffer
EPBuffer._buffer.update_ep_member() EPBuffer.get_existing_buffer().update_ep_member()
def try_recover_ranks(global_ranks: List[int]) -> bool: def try_recover_ranks(global_ranks: List[int]) -> bool:
@@ -580,14 +580,22 @@ class FlashInferWorkspaceManager:
self._logged_init = False self._logged_init = False
_attn_tp_workspace_manager = FlashInferWorkspaceManager()
_moe_tp_workspace_manager = FlashInferWorkspaceManager()
def _get_workspace_manager(use_attn_tp_group: bool) -> FlashInferWorkspaceManager: def _get_workspace_manager(use_attn_tp_group: bool) -> FlashInferWorkspaceManager:
return ( """The per-group fusion workspace manager; the instances live on
_attn_tp_workspace_manager if use_attn_tp_group else _moe_tp_workspace_manager ``ctx.resources`` (one per comm group, created lazily)."""
from sglang.srt.runtime_context import get_resources
buffers = get_resources().buffers
name = (
"flashinfer_fusion_attn_tp_workspace"
if use_attn_tp_group
else "flashinfer_fusion_moe_tp_workspace"
) )
manager = buffers.get(name)
if manager is None:
manager = FlashInferWorkspaceManager()
buffers[name] = manager
return manager
def _sync_allreduce_unavailable_across_tp(): def _sync_allreduce_unavailable_across_tp():
@@ -853,11 +861,13 @@ def pre_initialize_workspaces(
def cleanup_flashinfer_workspace(): def cleanup_flashinfer_workspace():
global _attn_tp_workspace_manager, _moe_tp_workspace_manager from sglang.srt.runtime_context import get_resources
if _attn_tp_workspace_manager is not None:
_attn_tp_workspace_manager.cleanup() buffers = get_resources().buffers
if ( for name in (
_moe_tp_workspace_manager is not None "flashinfer_fusion_attn_tp_workspace",
and _moe_tp_workspace_manager is not _attn_tp_workspace_manager "flashinfer_fusion_moe_tp_workspace",
): ):
_moe_tp_workspace_manager.cleanup() manager = buffers.get(name)
if manager is not None:
manager.cleanup()
@@ -159,11 +159,27 @@ class DeepEPDispatchMode(IntEnum):
class DeepEPBuffer: class DeepEPBuffer:
_buffer = None """Managing facade for the process-wide DeepEP comm buffer; the state
_dispatch_mode: Optional[DeepEPDispatchMode] = None itself lives on ``ctx.resources`` (one entry per process)."""
_hidden_size: Optional[int] = None
_num_max_dispatch_tokens_per_rank: Optional[int] = None @classmethod
_num_experts: Optional[int] = None def _state(cls):
from types import SimpleNamespace
from sglang.srt.runtime_context import get_resources
buffers = get_resources().buffers
state = buffers.get("deepep_ep_state")
if state is None:
state = SimpleNamespace(
buffer=None,
dispatch_mode=None,
hidden_size=None,
num_max_dispatch_tokens_per_rank=None,
num_experts=None,
)
buffers["deepep_ep_state"] = state
return state
@classmethod @classmethod
def get_deepep_buffer( def get_deepep_buffer(
@@ -175,12 +191,13 @@ class DeepEPBuffer:
num_max_dispatch_tokens_per_rank: int = -1, num_max_dispatch_tokens_per_rank: int = -1,
num_experts: int = -1, num_experts: int = -1,
): ):
if cls._buffer is not None: state = cls._state()
return cls._buffer if state.buffer is not None:
return state.buffer
cls._hidden_size = hidden_size state.hidden_size = hidden_size
cls._num_max_dispatch_tokens_per_rank = num_max_dispatch_tokens_per_rank state.num_max_dispatch_tokens_per_rank = num_max_dispatch_tokens_per_rank
cls._num_experts = num_experts state.num_experts = num_experts
num_nvl_bytes, num_rdma_bytes = 0, 0 num_nvl_bytes, num_rdma_bytes = 0, 0
if deepep_mode.enable_normal(): if deepep_mode.enable_normal():
@@ -263,28 +280,30 @@ class DeepEPBuffer:
if not is_cu12 and use_mnnvl_fabric: if not is_cu12 and use_mnnvl_fabric:
buffer_kwargs["use_fabric"] = True buffer_kwargs["use_fabric"] = True
cls._buffer = Buffer(group, num_nvl_bytes, num_rdma_bytes, **buffer_kwargs) state.buffer = Buffer(group, num_nvl_bytes, num_rdma_bytes, **buffer_kwargs)
return cls._buffer return state.buffer
@classmethod @classmethod
def clean_buffer(cls): def clean_buffer(cls):
if not cls._buffer.low_latency_mode: state = cls._state()
if not state.buffer.low_latency_mode:
return return
cls._buffer.clean_low_latency_buffer( state.buffer.clean_low_latency_buffer(
cls._num_max_dispatch_tokens_per_rank, state.num_max_dispatch_tokens_per_rank,
cls._hidden_size, state.hidden_size,
cls._num_experts, state.num_experts,
) )
@classmethod @classmethod
def set_dispatch_mode_as_normal(cls): def set_dispatch_mode_as_normal(cls):
cls._dispatch_mode = DeepEPDispatchMode.NORMAL cls._state().dispatch_mode = DeepEPDispatchMode.NORMAL
@classmethod @classmethod
def set_dispatch_mode_as_low_latency(cls): def set_dispatch_mode_as_low_latency(cls):
if cls._dispatch_mode == DeepEPDispatchMode.NORMAL: state = cls._state()
if state.dispatch_mode == DeepEPDispatchMode.NORMAL:
cls.clean_buffer() cls.clean_buffer()
cls._dispatch_mode = DeepEPDispatchMode.LOW_LATENCY state.dispatch_mode = DeepEPDispatchMode.LOW_LATENCY
@classmethod @classmethod
def set_dispatch_mode(cls, mode: DeepEPMode): def set_dispatch_mode(cls, mode: DeepEPMode):
@@ -57,10 +57,31 @@ assert isinstance(MooncakeCombineInput, CombineInput)
class EPBuffer: class EPBuffer:
_buffer = None """Managing facade for the process-wide Mooncake EP buffer; the state
_hidden_size: Optional[int] = None itself lives on ``ctx.resources``."""
_num_max_dispatch_tokens_per_rank: Optional[int] = None
_num_experts: Optional[int] = None @classmethod
def _state(cls):
from types import SimpleNamespace
from sglang.srt.runtime_context import get_resources
buffers = get_resources().buffers
state = buffers.get("mooncake_ep_state")
if state is None:
state = SimpleNamespace(
buffer=None,
hidden_size=None,
num_max_dispatch_tokens_per_rank=None,
num_experts=None,
)
buffers["mooncake_ep_state"] = state
return state
@classmethod
def get_existing_buffer(cls):
"""The already-created buffer (elastic-EP membership refresh)."""
return cls._state().buffer
@classmethod @classmethod
def get_ep_buffer( def get_ep_buffer(
@@ -72,15 +93,16 @@ class EPBuffer:
num_max_dispatch_tokens_per_rank: int = -1, num_max_dispatch_tokens_per_rank: int = -1,
num_experts: int = -1, num_experts: int = -1,
): ):
if cls._buffer is not None: state = cls._state()
return cls._buffer if state.buffer is not None:
return state.buffer
# Lazy import Buffer to avoid creating CUDA context at module import time # Lazy import Buffer to avoid creating CUDA context at module import time
from mooncake.mooncake_ep_buffer import Buffer from mooncake.mooncake_ep_buffer import Buffer
cls._hidden_size = hidden_size state.hidden_size = hidden_size
cls._num_max_dispatch_tokens_per_rank = num_max_dispatch_tokens_per_rank state.num_max_dispatch_tokens_per_rank = num_max_dispatch_tokens_per_rank
cls._num_experts = num_experts state.num_experts = num_experts
num_ep_buffer_bytes = 0 num_ep_buffer_bytes = 0
if deepep_mode.enable_normal(): if deepep_mode.enable_normal():
@@ -97,8 +119,8 @@ class EPBuffer:
num_experts, num_experts,
) )
cls._buffer = Buffer(group, num_ep_buffer_bytes) state.buffer = Buffer(group, num_ep_buffer_bytes)
return cls._buffer return state.buffer
class _MooncakeEPDispatcherImpl: class _MooncakeEPDispatcherImpl:
@@ -2,7 +2,6 @@ from __future__ import annotations
import logging import logging
from enum import Enum, auto from enum import Enum, auto
from typing import Optional
import torch import torch
import torch.distributed as dist import torch.distributed as dist
@@ -39,11 +38,27 @@ NixlEPCombineInput = DeepEPLLCombineInput
class NixlEPBuffer: class NixlEPBuffer:
_buffer = None """Managing facade for the process-wide NIXL EP buffer; the state itself
_hidden_size: Optional[int] = None lives on ``ctx.resources``."""
_num_max_dispatch_tokens_per_rank: Optional[int] = None
_num_experts: Optional[int] = None @classmethod
_num_local_experts: Optional[int] = None def _state(cls):
from types import SimpleNamespace
from sglang.srt.runtime_context import get_resources
buffers = get_resources().buffers
state = buffers.get("nixl_ep_state")
if state is None:
state = SimpleNamespace(
buffer=None,
hidden_size=None,
num_max_dispatch_tokens_per_rank=None,
num_experts=None,
num_local_experts=None,
)
buffers["nixl_ep_state"] = state
return state
@classmethod @classmethod
def get_nixl_buffer( def get_nixl_buffer(
@@ -55,13 +70,14 @@ class NixlEPBuffer:
num_experts: int = -1, num_experts: int = -1,
num_local_experts: int = -1, num_local_experts: int = -1,
): ):
if cls._buffer is not None: state = cls._state()
return cls._buffer if state.buffer is not None:
return state.buffer
cls._hidden_size = hidden_size state.hidden_size = hidden_size
cls._num_max_dispatch_tokens_per_rank = num_max_dispatch_tokens_per_rank state.num_max_dispatch_tokens_per_rank = num_max_dispatch_tokens_per_rank
cls._num_experts = num_experts state.num_experts = num_experts
cls._num_local_experts = num_local_experts state.num_local_experts = num_local_experts
num_rdma_bytes = 0 num_rdma_bytes = 0
if deepep_mode.enable_normal(): if deepep_mode.enable_normal():
@@ -89,30 +105,31 @@ class NixlEPBuffer:
logger.info( logger.info(
f"Using NIXL EP (world_size={world_size}, rank={rank}, " f"Using NIXL EP (world_size={world_size}, rank={rank}, "
f"num_experts={cls._num_experts}, num_experts_per_rank={cls._num_local_experts}) " f"num_experts={state.num_experts}, num_experts_per_rank={state.num_local_experts}) "
) )
cls._buffer = Buffer( state.buffer = Buffer(
rank=rank, rank=rank,
tcp_store_group=tcp_store, tcp_store_group=tcp_store,
) )
cls._buffer.update_memory_buffers( state.buffer.update_memory_buffers(
num_ranks=world_size, num_ranks=world_size,
num_experts_per_rank=cls._num_local_experts, num_experts_per_rank=state.num_local_experts,
num_rdma_bytes=num_rdma_bytes, num_rdma_bytes=num_rdma_bytes,
) )
all_ranks = list(range(world_size)) all_ranks = list(range(world_size))
cls._buffer.connect_ranks(all_ranks) state.buffer.connect_ranks(all_ranks)
return cls._buffer return state.buffer
@classmethod @classmethod
def clean_buffer(cls): def clean_buffer(cls):
cls._buffer.clean_buffer( state = cls._state()
cls._num_max_dispatch_tokens_per_rank, state.buffer.clean_buffer(
cls._hidden_size, state.num_max_dispatch_tokens_per_rank,
cls._num_experts, state.hidden_size,
state.num_experts,
) )
@@ -174,8 +174,12 @@ class TestFlashInferCommFusion(unittest.TestCase):
fake_comm = _FakeFlashInferComm() fake_comm = _FakeFlashInferComm()
original_comm = fusion._flashinfer_comm original_comm = fusion._flashinfer_comm
original_create = fusion._create_allreduce_fusion_workspace original_create = fusion._create_allreduce_fusion_workspace
original_manager = fusion._attn_tp_workspace_manager
original_unavailable = fusion._flashinfer_allreduce_unavailable original_unavailable = fusion._flashinfer_allreduce_unavailable
from sglang.srt.runtime_context import get_resources
buffers = get_resources().buffers
manager_key = "flashinfer_fusion_attn_tp_workspace"
original_manager = buffers.get(manager_key)
try: try:
fusion._flashinfer_comm = fake_comm fusion._flashinfer_comm = fake_comm
fusion._create_allreduce_fusion_workspace = ( fusion._create_allreduce_fusion_workspace = (
@@ -189,7 +193,7 @@ class TestFlashInferCommFusion(unittest.TestCase):
manager = fusion.FlashInferWorkspaceManager() manager = fusion.FlashInferWorkspaceManager()
manager.workspace = _FakeWorkspace(backend, world_size) manager.workspace = _FakeWorkspace(backend, world_size)
manager.initialized = True manager.initialized = True
fusion._attn_tp_workspace_manager = manager buffers[manager_key] = manager
if not torch.cuda.is_available(): if not torch.cuda.is_available():
self.skipTest("FlashInfer allreduce custom op is CUDA-only") self.skipTest("FlashInfer allreduce custom op is CUDA-only")
device = torch.device("cuda") device = torch.device("cuda")
@@ -229,7 +233,10 @@ class TestFlashInferCommFusion(unittest.TestCase):
finally: finally:
fusion._flashinfer_comm = original_comm fusion._flashinfer_comm = original_comm
fusion._create_allreduce_fusion_workspace = original_create fusion._create_allreduce_fusion_workspace = original_create
fusion._attn_tp_workspace_manager = original_manager if original_manager is None:
buffers.pop(manager_key, None)
else:
buffers[manager_key] = original_manager
fusion._flashinfer_allreduce_unavailable = original_unavailable fusion._flashinfer_allreduce_unavailable = original_unavailable
@@ -463,6 +463,43 @@ class TestNamedStreams(_IsolatedServerArgs):
self.assertEqual(get_context().resources.streams, {}) self.assertEqual(get_context().resources.streams, {})
class TestEpBufferState(_IsolatedServerArgs):
"""EP dispatcher buffer managers: state lives on ctx.resources; the
facade keeps the mode-transition and clean semantics."""
def test_deepep_dispatch_mode_transitions_and_reset(self):
try:
from sglang.srt.layers.moe.token_dispatcher.deepep import DeepEPBuffer
except ImportError:
self.skipTest("deep_ep not installed")
reset_context()
cleans = []
class _FakeBuffer:
low_latency_mode = True
def clean_low_latency_buffer(self, *args):
cleans.append(args)
state = DeepEPBuffer._state()
state.buffer = _FakeBuffer()
state.hidden_size = 7168
state.num_max_dispatch_tokens_per_rank = 128
state.num_experts = 256
DeepEPBuffer.set_dispatch_mode_as_normal()
# NORMAL -> LOW_LATENCY must clean the low-latency buffer once.
DeepEPBuffer.set_dispatch_mode_as_low_latency()
self.assertEqual(cleans, [(128, 7168, 256)])
# LOW_LATENCY -> LOW_LATENCY must not clean again.
DeepEPBuffer.set_dispatch_mode_as_low_latency()
self.assertEqual(len(cleans), 1)
reset_context()
self.assertIsNone(DeepEPBuffer._state().buffer)
class TestPublishLifecycle(_IsolatedServerArgs): class TestPublishLifecycle(_IsolatedServerArgs):
"""Publish installs the resolved server_args and seeds the capture tier.""" """Publish installs the resolved server_args and seeds the capture tier."""