diff --git a/python/sglang/srt/distributed/bootstrap.py b/python/sglang/srt/distributed/bootstrap.py index ab7be1f9e..e0c37793d 100644 --- a/python/sglang/srt/distributed/bootstrap.py +++ b/python/sglang/srt/distributed/bootstrap.py @@ -61,15 +61,7 @@ def init_torch_distributed( tic = time.perf_counter() logger.info("Init torch distributed begin.") - try: - torch.get_device_module(device).set_device(ps.gpu_id) - except Exception: - logger.warning( - f"Context: {device=} {ps.gpu_id=} {os.environ.get('CUDA_VISIBLE_DEVICES')=} {ps.tp_rank=} {ps.tp_size=}" - ) - raise - - backend = _resolve_backend(device=device, server_args=server_args, gpu_id=ps.gpu_id) + backend = _resolve_backend(device=device, server_args=server_args) before_avail_memory = get_available_gpu_memory(device, ps.gpu_id) if not server_args.enable_p2p_check: @@ -143,27 +135,10 @@ def init_torch_distributed( ) -def _resolve_backend(*, device: str, server_args: ServerArgs, gpu_id: int) -> str: +def _resolve_backend(*, device: str, server_args: ServerArgs) -> str: backend = get_default_distributed_backend(device) if device == "cuda" and server_args.elastic_ep_backend == "mooncake": backend = "mooncake" - if server_args.mooncake_ib_device: - from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import ( - get_ib_devices_for_gpu, - ) - - ib_device_for_gpu = get_ib_devices_for_gpu( - server_args.mooncake_ib_device, gpu_id - ) - mooncake_ib_device = ( - ib_device_for_gpu.split(",") if ib_device_for_gpu else [] - ) - try: - from mooncake import ep as mooncake_ep - - mooncake_ep.set_device_filter(mooncake_ib_device) - except: - pass # A warning will be raised in `init_distributed_environment` return backend diff --git a/python/sglang/srt/distributed/device_communicators/mooncake_transfer_engine.py b/python/sglang/srt/distributed/device_communicators/mooncake_transfer_engine.py index ce08c6fbc..de7c011ec 100644 --- a/python/sglang/srt/distributed/device_communicators/mooncake_transfer_engine.py +++ b/python/sglang/srt/distributed/device_communicators/mooncake_transfer_engine.py @@ -333,6 +333,7 @@ def maybe_init_shared_mooncake_transfer_engine( server_args.enable_elastic_expert_backup and server_args.elastic_ep_backend is not None ) + or server_args.elastic_ep_backend == "mooncake" ) if use_mooncake_te: @@ -343,3 +344,14 @@ def maybe_init_shared_mooncake_transfer_engine( server_args.disaggregation_ib_device or server_args.mooncake_ib_device ), ) + + if server_args.elastic_ep_backend == "mooncake": + try: + from mooncake.pg import set_transfer_engine + except ImportError as e: + raise ImportError( + "Failed to import 'set_transfer_engine' from 'mooncake.pg'. " + "Please upgrade your 'mooncake-transfer-engine' " + "installation to 0.3.11 or above." + ) from e + set_transfer_engine(_mooncake_transfer_engine.engine) diff --git a/python/sglang/srt/distributed/parallel_state.py b/python/sglang/srt/distributed/parallel_state.py index df1e90b93..241be9026 100644 --- a/python/sglang/srt/distributed/parallel_state.py +++ b/python/sglang/srt/distributed/parallel_state.py @@ -1960,17 +1960,6 @@ def init_distributed_environment( distributed_init_method, backend, ) - if "mooncake" in backend: - try: - from mooncake import ep as mooncake_ep - except ImportError as e: - raise ImportError( - "Please install mooncake by following the instructions at " - "https://github.com/kvcache-ai/Mooncake/blob/main/doc/en/build.md " # noqa: E501 - "to run SGLang with Mooncake Backend." - ) from e - mooncake_ep.set_host_ip(get_local_ip_auto()) - if not torch.distributed.is_initialized(): global _MODEL_PARALLEL_GROUP_TIMEOUT assert distributed_init_method is not None, ( diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 67a65b473..232919230 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -322,13 +322,27 @@ class ModelRunner: if get_server_args().enable_tf32_matmul: torch.set_float32_matmul_precision("high") + # Set device early so that TransferEngine init (e.g. Ascend NPU) + # can access the device context. + try: + torch.get_device_module(self.device).set_device(ps.gpu_id) + except Exception: + import os + + logger.warning( + f"Context: {self.device=} {ps.gpu_id=} {os.environ.get('CUDA_VISIBLE_DEVICES')=} {ps.tp_rank=} {ps.tp_size=}" + ) + raise + + # Initialize MooncakeTransferEngine BEFORE init_torch_distributed so + # that the shared TE can be passed to the Mooncake PG backend (avoids + # creating duplicate TransferEngines). + self.init_shared_mooncake_transfer_engine() + # Get available memory before model loading. # Stored for later use by alloc_memory_pool(). self.init_torch_distributed() - # Initialize MooncakeTransferEngine - self.init_shared_mooncake_transfer_engine() - # Init forward stream for overlap schedule self.forward_stream = torch.get_device_module(self.device).Stream()