Absorb capturer setup and extract the shared-mooncake gate (#31160)
This commit is contained in:
@@ -1,15 +1,20 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
from typing import Dict, List, Optional, Union
|
from typing import TYPE_CHECKING, Dict, List, Optional, Union
|
||||||
|
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.utils.network import NetworkAddress, get_free_port
|
from sglang.srt.utils.network import NetworkAddress, get_free_port, get_local_ip_auto
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from sglang.srt.server_args import ServerArgs
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
# Module-level shared engine instance, set by init_mooncake_transfer_engine().
|
# Module-level shared engine instance, set by init_mooncake_transfer_engine().
|
||||||
_mooncake_transfer_engine: Optional["MooncakeTransferEngine"] = None
|
_mooncake_transfer_engine: Optional[MooncakeTransferEngine] = None
|
||||||
|
|
||||||
|
|
||||||
def parse_ib_device_config(
|
def parse_ib_device_config(
|
||||||
@@ -295,3 +300,46 @@ def init_mooncake_transfer_engine(
|
|||||||
def get_mooncake_transfer_engine() -> Optional[MooncakeTransferEngine]:
|
def get_mooncake_transfer_engine() -> Optional[MooncakeTransferEngine]:
|
||||||
"""Return the shared MooncakeTransferEngine if initialized, else None."""
|
"""Return the shared MooncakeTransferEngine if initialized, else None."""
|
||||||
return _mooncake_transfer_engine
|
return _mooncake_transfer_engine
|
||||||
|
|
||||||
|
|
||||||
|
def maybe_init_shared_mooncake_transfer_engine(
|
||||||
|
*, server_args: ServerArgs, gpu_id: int
|
||||||
|
) -> None:
|
||||||
|
"""
|
||||||
|
Need MooncakeTransferEngine when:
|
||||||
|
1) PD disaggregation uses mooncake for KV transfer (prefill/decode)
|
||||||
|
2) HiCache uses mooncake storage backend
|
||||||
|
3) Encoder disaggregation uses mooncake
|
||||||
|
"""
|
||||||
|
use_mooncake_te = (
|
||||||
|
(
|
||||||
|
server_args.disaggregation_mode != "null"
|
||||||
|
and server_args.disaggregation_transfer_backend == "mooncake"
|
||||||
|
)
|
||||||
|
or (
|
||||||
|
server_args.enable_hierarchical_cache
|
||||||
|
and server_args.hicache_storage_backend == "mooncake"
|
||||||
|
and envs.SGLANG_HICACHE_MOONCAKE_REUSE_TE.get()
|
||||||
|
)
|
||||||
|
or (
|
||||||
|
server_args.encoder_only
|
||||||
|
and server_args.encoder_transfer_backend == "mooncake"
|
||||||
|
)
|
||||||
|
or (
|
||||||
|
server_args.language_only
|
||||||
|
and server_args.encoder_transfer_backend == "mooncake"
|
||||||
|
)
|
||||||
|
or (
|
||||||
|
server_args.enable_elastic_expert_backup
|
||||||
|
and server_args.elastic_ep_backend is not None
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
if use_mooncake_te:
|
||||||
|
init_mooncake_transfer_engine(
|
||||||
|
hostname=get_local_ip_auto(),
|
||||||
|
gpu_id=gpu_id,
|
||||||
|
ib_device=(
|
||||||
|
server_args.disaggregation_ib_device or server_args.mooncake_ib_device
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|||||||
@@ -30,7 +30,6 @@ from sglang.srt.configs.model_config import (
|
|||||||
AttentionArch,
|
AttentionArch,
|
||||||
ModelConfig,
|
ModelConfig,
|
||||||
ModelImpl,
|
ModelImpl,
|
||||||
get_num_indexer_layers,
|
|
||||||
)
|
)
|
||||||
from sglang.srt.configs.update_config import adjust_config_with_unaligned_cpu_tp
|
from sglang.srt.configs.update_config import adjust_config_with_unaligned_cpu_tp
|
||||||
from sglang.srt.debug_utils.dumper import dumper
|
from sglang.srt.debug_utils.dumper import dumper
|
||||||
@@ -38,6 +37,9 @@ from sglang.srt.distributed import (
|
|||||||
bootstrap,
|
bootstrap,
|
||||||
get_world_group,
|
get_world_group,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import (
|
||||||
|
maybe_init_shared_mooncake_transfer_engine,
|
||||||
|
)
|
||||||
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||||
prealloc_symmetric_memory_pool,
|
prealloc_symmetric_memory_pool,
|
||||||
)
|
)
|
||||||
@@ -195,7 +197,6 @@ from sglang.srt.utils import (
|
|||||||
set_cuda_arch,
|
set_cuda_arch,
|
||||||
slow_rank_detector,
|
slow_rank_detector,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils.network import get_local_ip_auto
|
|
||||||
from sglang.srt.utils.nvtx_pytorch_hooks import PytHooks
|
from sglang.srt.utils.nvtx_pytorch_hooks import PytHooks
|
||||||
from sglang.srt.utils.nvtx_utils import profile_range
|
from sglang.srt.utils.nvtx_utils import profile_range
|
||||||
from sglang.srt.utils.offloader import (
|
from sglang.srt.utils.offloader import (
|
||||||
@@ -751,18 +752,10 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
# would overwrite the target's process-global one.
|
# would overwrite the target's process-global one.
|
||||||
return
|
return
|
||||||
|
|
||||||
if not self.server_args.disable_shared_experts_fusion and hasattr(
|
|
||||||
self.model, "num_fused_shared_experts"
|
|
||||||
):
|
|
||||||
num_fused_shared_experts = self.model.num_fused_shared_experts
|
|
||||||
else:
|
|
||||||
num_fused_shared_experts = 0
|
|
||||||
|
|
||||||
set_global_experts_capturer(
|
set_global_experts_capturer(
|
||||||
RoutedExpertsCapturer.create(
|
RoutedExpertsCapturer.create(
|
||||||
enable=get_server_args().enable_return_routed_experts,
|
model=self.model,
|
||||||
model_config=self.model_config,
|
model_config=self.model_config,
|
||||||
num_fused_shared_experts=num_fused_shared_experts,
|
|
||||||
num_tokens=self.max_total_num_tokens + self.page_size,
|
num_tokens=self.max_total_num_tokens + self.page_size,
|
||||||
max_running_requests=self.max_running_requests,
|
max_running_requests=self.max_running_requests,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
@@ -770,26 +763,9 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def init_indexer_capturer(self):
|
def init_indexer_capturer(self):
|
||||||
enable = get_server_args().enable_return_indexer_topk
|
|
||||||
# Producer wiring is CUDA-only (Indexer.forward_cuda + MLA skip_topk
|
|
||||||
# path); other backends would create a capturer but never feed it.
|
|
||||||
if enable and self.device != "cuda":
|
|
||||||
logger.warning(
|
|
||||||
"indexer-topk capture is CUDA-only; %s backend not yet wired. "
|
|
||||||
"Disabling capturer.",
|
|
||||||
self.device,
|
|
||||||
)
|
|
||||||
set_global_indexer_capturer(None)
|
|
||||||
return
|
|
||||||
|
|
||||||
hf_text_config = self.model_config.hf_text_config
|
|
||||||
num_indexer_layers = get_num_indexer_layers(hf_text_config)
|
|
||||||
index_topk = getattr(hf_text_config, "index_topk", 0)
|
|
||||||
set_global_indexer_capturer(
|
set_global_indexer_capturer(
|
||||||
create_indexer_capturer(
|
create_indexer_capturer(
|
||||||
enable=enable,
|
model_config=self.model_config,
|
||||||
num_indexer_layers=num_indexer_layers,
|
|
||||||
index_topk=index_topk,
|
|
||||||
num_tokens=self.max_total_num_tokens + self.page_size,
|
num_tokens=self.max_total_num_tokens + self.page_size,
|
||||||
max_running_requests=self.max_running_requests,
|
max_running_requests=self.max_running_requests,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
@@ -857,48 +833,8 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
self.pre_model_load_memory = result.pre_model_load_memory
|
self.pre_model_load_memory = result.pre_model_load_memory
|
||||||
|
|
||||||
def init_shared_mooncake_transfer_engine(self):
|
def init_shared_mooncake_transfer_engine(self):
|
||||||
"""
|
maybe_init_shared_mooncake_transfer_engine(
|
||||||
Need MooncakeTransferEngine when:
|
server_args=self.server_args, gpu_id=self.gpu_id
|
||||||
1) PD disaggregation uses mooncake for KV transfer (prefill/decode)
|
|
||||||
2) HiCache uses mooncake storage backend
|
|
||||||
3) Encoder disaggregation uses mooncake
|
|
||||||
"""
|
|
||||||
use_mooncake_te = (
|
|
||||||
(
|
|
||||||
self.server_args.disaggregation_mode != "null"
|
|
||||||
and self.server_args.disaggregation_transfer_backend == "mooncake"
|
|
||||||
)
|
|
||||||
or (
|
|
||||||
self.server_args.enable_hierarchical_cache
|
|
||||||
and self.server_args.hicache_storage_backend == "mooncake"
|
|
||||||
and envs.SGLANG_HICACHE_MOONCAKE_REUSE_TE.get()
|
|
||||||
)
|
|
||||||
or (
|
|
||||||
self.server_args.encoder_only
|
|
||||||
and self.server_args.encoder_transfer_backend == "mooncake"
|
|
||||||
)
|
|
||||||
or (
|
|
||||||
self.server_args.language_only
|
|
||||||
and self.server_args.encoder_transfer_backend == "mooncake"
|
|
||||||
)
|
|
||||||
or (
|
|
||||||
self.server_args.enable_elastic_expert_backup
|
|
||||||
and self.server_args.elastic_ep_backend is not None
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
if use_mooncake_te:
|
|
||||||
from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import (
|
|
||||||
init_mooncake_transfer_engine,
|
|
||||||
)
|
|
||||||
|
|
||||||
init_mooncake_transfer_engine(
|
|
||||||
hostname=get_local_ip_auto(),
|
|
||||||
gpu_id=self.gpu_id,
|
|
||||||
ib_device=(
|
|
||||||
self.server_args.disaggregation_ib_device
|
|
||||||
or self.server_args.mooncake_ib_device
|
|
||||||
),
|
|
||||||
)
|
)
|
||||||
|
|
||||||
def load_model(self):
|
def load_model(self):
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ import numpy as np
|
|||||||
import pybase64
|
import pybase64
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.configs.model_config import ModelConfig, get_num_indexer_layers
|
||||||
from sglang.srt.runtime_context import get_parallel
|
from sglang.srt.runtime_context import get_parallel
|
||||||
from sglang.srt.state_capturer.base import BaseTopkCapturer
|
from sglang.srt.state_capturer.base import BaseTopkCapturer
|
||||||
|
|
||||||
@@ -82,6 +83,39 @@ def extract_indexer_topk_from_meta_info(data):
|
|||||||
|
|
||||||
|
|
||||||
def create_indexer_capturer(
|
def create_indexer_capturer(
|
||||||
|
*,
|
||||||
|
model_config: ModelConfig,
|
||||||
|
num_tokens: int,
|
||||||
|
max_running_requests: int,
|
||||||
|
device: str,
|
||||||
|
) -> Optional[IndexerTopkCapturer]:
|
||||||
|
from sglang.srt.runtime_context import get_server_args
|
||||||
|
|
||||||
|
enable = get_server_args().enable_return_indexer_topk
|
||||||
|
# Producer wiring is CUDA-only (Indexer.forward_cuda + MLA skip_topk
|
||||||
|
# path); other backends would create a capturer but never feed it.
|
||||||
|
if enable and device != "cuda":
|
||||||
|
logger.warning(
|
||||||
|
"indexer-topk capture is CUDA-only; %s backend not yet wired. "
|
||||||
|
"Disabling capturer.",
|
||||||
|
device,
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
|
||||||
|
hf_text_config = model_config.hf_text_config
|
||||||
|
num_indexer_layers = get_num_indexer_layers(hf_text_config)
|
||||||
|
index_topk = getattr(hf_text_config, "index_topk", 0)
|
||||||
|
return _create_indexer_capturer_raw(
|
||||||
|
enable=enable,
|
||||||
|
num_indexer_layers=num_indexer_layers,
|
||||||
|
index_topk=index_topk,
|
||||||
|
num_tokens=num_tokens,
|
||||||
|
max_running_requests=max_running_requests,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _create_indexer_capturer_raw(
|
||||||
enable: bool,
|
enable: bool,
|
||||||
num_indexer_layers: int,
|
num_indexer_layers: int,
|
||||||
index_topk: int,
|
index_topk: int,
|
||||||
|
|||||||
@@ -28,15 +28,22 @@ class RoutedExpertsCapturer(BaseTopkCapturer):
|
|||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def create(
|
def create(
|
||||||
enable: bool,
|
*,
|
||||||
|
model: torch.nn.Module,
|
||||||
model_config: ModelConfig,
|
model_config: ModelConfig,
|
||||||
num_fused_shared_experts: int,
|
|
||||||
num_tokens: int,
|
num_tokens: int,
|
||||||
max_running_requests: int,
|
max_running_requests: int,
|
||||||
device: str,
|
device: str,
|
||||||
) -> Optional["RoutedExpertsCapturer"]:
|
) -> Optional["RoutedExpertsCapturer"]:
|
||||||
if not enable:
|
server_args = get_server_args()
|
||||||
|
if not server_args.enable_return_routed_experts:
|
||||||
return None
|
return None
|
||||||
|
if not server_args.disable_shared_experts_fusion and hasattr(
|
||||||
|
model, "num_fused_shared_experts"
|
||||||
|
):
|
||||||
|
num_fused_shared_experts = model.num_fused_shared_experts
|
||||||
|
else:
|
||||||
|
num_fused_shared_experts = 0
|
||||||
return RoutedExpertsCapturer(
|
return RoutedExpertsCapturer(
|
||||||
model_config,
|
model_config,
|
||||||
num_tokens=num_tokens,
|
num_tokens=num_tokens,
|
||||||
|
|||||||
@@ -65,7 +65,7 @@ def test_mooncake_te_condition(server_args: ServerArgs) -> bool:
|
|||||||
side_effect=_fake_init_mooncake_transfer_engine,
|
side_effect=_fake_init_mooncake_transfer_engine,
|
||||||
),
|
),
|
||||||
patch(
|
patch(
|
||||||
"sglang.srt.model_executor.model_runner.get_local_ip_auto",
|
"sglang.srt.distributed.device_communicators.mooncake_transfer_engine.get_local_ip_auto",
|
||||||
return_value="127.0.0.1",
|
return_value="127.0.0.1",
|
||||||
),
|
),
|
||||||
):
|
):
|
||||||
|
|||||||
Reference in New Issue
Block a user