[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:
|
||||
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,
|
||||
)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user