Absorb capturer setup and extract the shared-mooncake gate (#31160)

This commit is contained in:
fzyzcjy
2026-07-14 16:00:51 +08:00
committed by GitHub
parent 08798dba0d
commit 1dc48c2c3b
5 changed files with 103 additions and 78 deletions
@@ -1,15 +1,20 @@
from __future__ import annotations
import json
import logging
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.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__)
# 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(
@@ -295,3 +300,46 @@ def init_mooncake_transfer_engine(
def get_mooncake_transfer_engine() -> Optional[MooncakeTransferEngine]:
"""Return the shared MooncakeTransferEngine if initialized, else None."""
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,
ModelConfig,
ModelImpl,
get_num_indexer_layers,
)
from sglang.srt.configs.update_config import adjust_config_with_unaligned_cpu_tp
from sglang.srt.debug_utils.dumper import dumper
@@ -38,6 +37,9 @@ from sglang.srt.distributed import (
bootstrap,
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 (
prealloc_symmetric_memory_pool,
)
@@ -195,7 +197,6 @@ from sglang.srt.utils import (
set_cuda_arch,
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_utils import profile_range
from sglang.srt.utils.offloader import (
@@ -751,18 +752,10 @@ class ModelRunner(ModelRunnerKVCacheMixin):
# would overwrite the target's process-global one.
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(
RoutedExpertsCapturer.create(
enable=get_server_args().enable_return_routed_experts,
model=self.model,
model_config=self.model_config,
num_fused_shared_experts=num_fused_shared_experts,
num_tokens=self.max_total_num_tokens + self.page_size,
max_running_requests=self.max_running_requests,
device=self.device,
@@ -770,26 +763,9 @@ class ModelRunner(ModelRunnerKVCacheMixin):
)
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(
create_indexer_capturer(
enable=enable,
num_indexer_layers=num_indexer_layers,
index_topk=index_topk,
model_config=self.model_config,
num_tokens=self.max_total_num_tokens + self.page_size,
max_running_requests=self.max_running_requests,
device=self.device,
@@ -857,50 +833,10 @@ class ModelRunner(ModelRunnerKVCacheMixin):
self.pre_model_load_memory = result.pre_model_load_memory
def init_shared_mooncake_transfer_engine(self):
"""
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 = (
(
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
)
maybe_init_shared_mooncake_transfer_engine(
server_args=self.server_args, gpu_id=self.gpu_id
)
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):
tic_total = time.perf_counter()
before_avail_memory = get_available_gpu_memory(self.device, self.gpu_id)
@@ -5,6 +5,7 @@ import numpy as np
import pybase64
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.state_capturer.base import BaseTopkCapturer
@@ -82,6 +83,39 @@ def extract_indexer_topk_from_meta_info(data):
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,
num_indexer_layers: int,
index_topk: int,
@@ -28,15 +28,22 @@ class RoutedExpertsCapturer(BaseTopkCapturer):
@staticmethod
def create(
enable: bool,
*,
model: torch.nn.Module,
model_config: ModelConfig,
num_fused_shared_experts: int,
num_tokens: int,
max_running_requests: int,
device: str,
) -> Optional["RoutedExpertsCapturer"]:
if not enable:
server_args = get_server_args()
if not server_args.enable_return_routed_experts:
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(
model_config,
num_tokens=num_tokens,