[Elastic EP] Centralize Mooncake PG configuration (#31708)
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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, (
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user