[refactor] Move the EP dispatcher and fusion-workspace manager state onto ctx.resources (#30489)
This commit is contained in:
@@ -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."""
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user