[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:
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:
@@ -580,14 +580,22 @@ class FlashInferWorkspaceManager:
self._logged_init = False
_attn_tp_workspace_manager = FlashInferWorkspaceManager()
_moe_tp_workspace_manager = FlashInferWorkspaceManager()
def _get_workspace_manager(use_attn_tp_group: bool) -> FlashInferWorkspaceManager:
return (
_attn_tp_workspace_manager if use_attn_tp_group else _moe_tp_workspace_manager
"""The per-group fusion workspace manager; the instances live on
``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():
@@ -853,11 +861,13 @@ def pre_initialize_workspaces(
def cleanup_flashinfer_workspace():
global _attn_tp_workspace_manager, _moe_tp_workspace_manager
if _attn_tp_workspace_manager is not None:
_attn_tp_workspace_manager.cleanup()
if (
_moe_tp_workspace_manager is not None
and _moe_tp_workspace_manager is not _attn_tp_workspace_manager
from sglang.srt.runtime_context import get_resources
buffers = get_resources().buffers
for name in (
"flashinfer_fusion_attn_tp_workspace",
"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:
_buffer = None
_dispatch_mode: Optional[DeepEPDispatchMode] = None
_hidden_size: Optional[int] = None
_num_max_dispatch_tokens_per_rank: Optional[int] = None
_num_experts: Optional[int] = None
"""Managing facade for the process-wide DeepEP comm buffer; the state
itself lives on ``ctx.resources`` (one entry per process)."""
@classmethod
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
def get_deepep_buffer(
@@ -175,12 +191,13 @@ class DeepEPBuffer:
num_max_dispatch_tokens_per_rank: int = -1,
num_experts: int = -1,
):
if cls._buffer is not None:
return cls._buffer
state = cls._state()
if state.buffer is not None:
return state.buffer
cls._hidden_size = hidden_size
cls._num_max_dispatch_tokens_per_rank = num_max_dispatch_tokens_per_rank
cls._num_experts = num_experts
state.hidden_size = hidden_size
state.num_max_dispatch_tokens_per_rank = num_max_dispatch_tokens_per_rank
state.num_experts = num_experts
num_nvl_bytes, num_rdma_bytes = 0, 0
if deepep_mode.enable_normal():
@@ -263,28 +280,30 @@ class DeepEPBuffer:
if not is_cu12 and use_mnnvl_fabric:
buffer_kwargs["use_fabric"] = True
cls._buffer = Buffer(group, num_nvl_bytes, num_rdma_bytes, **buffer_kwargs)
return cls._buffer
state.buffer = Buffer(group, num_nvl_bytes, num_rdma_bytes, **buffer_kwargs)
return state.buffer
@classmethod
def clean_buffer(cls):
if not cls._buffer.low_latency_mode:
state = cls._state()
if not state.buffer.low_latency_mode:
return
cls._buffer.clean_low_latency_buffer(
cls._num_max_dispatch_tokens_per_rank,
cls._hidden_size,
cls._num_experts,
state.buffer.clean_low_latency_buffer(
state.num_max_dispatch_tokens_per_rank,
state.hidden_size,
state.num_experts,
)
@classmethod
def set_dispatch_mode_as_normal(cls):
cls._dispatch_mode = DeepEPDispatchMode.NORMAL
cls._state().dispatch_mode = DeepEPDispatchMode.NORMAL
@classmethod
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._dispatch_mode = DeepEPDispatchMode.LOW_LATENCY
state.dispatch_mode = DeepEPDispatchMode.LOW_LATENCY
@classmethod
def set_dispatch_mode(cls, mode: DeepEPMode):
@@ -57,10 +57,31 @@ assert isinstance(MooncakeCombineInput, CombineInput)
class EPBuffer:
_buffer = None
_hidden_size: Optional[int] = None
_num_max_dispatch_tokens_per_rank: Optional[int] = None
_num_experts: Optional[int] = None
"""Managing facade for the process-wide Mooncake EP buffer; the state
itself lives on ``ctx.resources``."""
@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
def get_ep_buffer(
@@ -72,15 +93,16 @@ class EPBuffer:
num_max_dispatch_tokens_per_rank: int = -1,
num_experts: int = -1,
):
if cls._buffer is not None:
return cls._buffer
state = cls._state()
if state.buffer is not None:
return state.buffer
# Lazy import Buffer to avoid creating CUDA context at module import time
from mooncake.mooncake_ep_buffer import Buffer
cls._hidden_size = hidden_size
cls._num_max_dispatch_tokens_per_rank = num_max_dispatch_tokens_per_rank
cls._num_experts = num_experts
state.hidden_size = hidden_size
state.num_max_dispatch_tokens_per_rank = num_max_dispatch_tokens_per_rank
state.num_experts = num_experts
num_ep_buffer_bytes = 0
if deepep_mode.enable_normal():
@@ -97,8 +119,8 @@ class EPBuffer:
num_experts,
)
cls._buffer = Buffer(group, num_ep_buffer_bytes)
return cls._buffer
state.buffer = Buffer(group, num_ep_buffer_bytes)
return state.buffer
class _MooncakeEPDispatcherImpl:
@@ -2,7 +2,6 @@ from __future__ import annotations
import logging
from enum import Enum, auto
from typing import Optional
import torch
import torch.distributed as dist
@@ -39,11 +38,27 @@ 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
"""Managing facade for the process-wide NIXL EP buffer; the state itself
lives on ``ctx.resources``."""
@classmethod
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
def get_nixl_buffer(
@@ -55,13 +70,14 @@ class NixlEPBuffer:
num_experts: int = -1,
num_local_experts: int = -1,
):
if cls._buffer is not None:
return cls._buffer
state = cls._state()
if state.buffer is not None:
return state.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
state.hidden_size = hidden_size
state.num_max_dispatch_tokens_per_rank = num_max_dispatch_tokens_per_rank
state.num_experts = num_experts
state.num_local_experts = num_local_experts
num_rdma_bytes = 0
if deepep_mode.enable_normal():
@@ -89,30 +105,31 @@ class NixlEPBuffer:
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}) "
f"num_experts={state.num_experts}, num_experts_per_rank={state.num_local_experts}) "
)
cls._buffer = Buffer(
state.buffer = Buffer(
rank=rank,
tcp_store_group=tcp_store,
)
cls._buffer.update_memory_buffers(
state.buffer.update_memory_buffers(
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,
)
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
def clean_buffer(cls):
cls._buffer.clean_buffer(
cls._num_max_dispatch_tokens_per_rank,
cls._hidden_size,
cls._num_experts,
state = cls._state()
state.buffer.clean_buffer(
state.num_max_dispatch_tokens_per_rank,
state.hidden_size,
state.num_experts,
)