[PD] Bound transfer engine init with SGLANG_DISAGGREGATION_ENGINE_INIT_TIMEOUT (#37874)
This commit is contained in:
@@ -8,6 +8,8 @@ from sglang.srt.disaggregation.utils import DisaggregationMode
|
|||||||
from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import (
|
from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import (
|
||||||
MooncakeTransferEngine,
|
MooncakeTransferEngine,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.environ import envs
|
||||||
|
from sglang.srt.utils.common import run_with_deadline
|
||||||
from sglang.srt.utils.network import NetworkAddress
|
from sglang.srt.utils.network import NetworkAddress
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -75,8 +77,12 @@ class AscendTransferEngine(MooncakeTransferEngine):
|
|||||||
output_tensor_list, tmp_tensor, group=get_world_group().device_group
|
output_tensor_list, tmp_tensor, group=get_world_group().device_group
|
||||||
)
|
)
|
||||||
"""Initialize the ascend transfer instance."""
|
"""Initialize the ascend transfer instance."""
|
||||||
ret_value = self.engine.initialize(
|
ret_value = run_with_deadline(
|
||||||
self.store_url, self.session_id, self.role, self.npu_id, trans_op_type
|
lambda: self.engine.initialize(
|
||||||
|
self.store_url, self.session_id, self.role, self.npu_id, trans_op_type
|
||||||
|
),
|
||||||
|
timeout_s=envs.SGLANG_DISAGGREGATION_ENGINE_INIT_TIMEOUT.get(),
|
||||||
|
what=f"Ascend TransferEngine.initialize({self.store_url!r}, {self.session_id!r})",
|
||||||
)
|
)
|
||||||
if ret_value != 0:
|
if ret_value != 0:
|
||||||
logger.error("Ascend Transfer Engine initialization failed.")
|
logger.error("Ascend Transfer Engine initialization failed.")
|
||||||
|
|||||||
@@ -45,6 +45,7 @@ from sglang.srt.disaggregation.common.utils import (
|
|||||||
from sglang.srt.disaggregation.utils import DisaggregationMode
|
from sglang.srt.disaggregation.utils import DisaggregationMode
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
|
from sglang.srt.utils.common import run_with_deadline
|
||||||
from sglang.srt.utils.network import NetworkAddress, get_local_ip_auto
|
from sglang.srt.utils.network import NetworkAddress, get_local_ip_auto
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -350,7 +351,11 @@ class MoriKVManager(CommonKVManager):
|
|||||||
f"{uuid.uuid4().hex[:8]}"
|
f"{uuid.uuid4().hex[:8]}"
|
||||||
)
|
)
|
||||||
|
|
||||||
engine = IOEngine(engine_key, config)
|
engine = run_with_deadline(
|
||||||
|
lambda: IOEngine(engine_key, config),
|
||||||
|
timeout_s=envs.SGLANG_DISAGGREGATION_ENGINE_INIT_TIMEOUT.get(),
|
||||||
|
what=f"Mori IOEngine({engine_key!r}, host={self.local_ip!r})",
|
||||||
|
)
|
||||||
poll_mode = PollCqMode.POLLING
|
poll_mode = PollCqMode.POLLING
|
||||||
|
|
||||||
qp_per_transfer = envs.SGLANG_MORI_QP_PER_TRANSFER.get()
|
qp_per_transfer = envs.SGLANG_MORI_QP_PER_TRANSFER.get()
|
||||||
|
|||||||
@@ -48,6 +48,7 @@ from sglang.srt.disaggregation.utils import (
|
|||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.runtime_context import get_parallel, get_schedule
|
from sglang.srt.runtime_context import get_parallel, get_schedule
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
|
from sglang.srt.utils.common import run_with_deadline
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from nixl._bindings import (
|
from nixl._bindings import (
|
||||||
@@ -456,7 +457,11 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
|
|||||||
backend_params.setdefault("thread_count", str(num_threads))
|
backend_params.setdefault("thread_count", str(num_threads))
|
||||||
elif backend == "UCCL":
|
elif backend == "UCCL":
|
||||||
backend_params.setdefault("num_cpus", str(num_threads))
|
backend_params.setdefault("num_cpus", str(num_threads))
|
||||||
self.agent.create_backend(backend, backend_params)
|
run_with_deadline(
|
||||||
|
lambda: self.agent.create_backend(backend, backend_params),
|
||||||
|
timeout_s=envs.SGLANG_DISAGGREGATION_ENGINE_INIT_TIMEOUT.get(),
|
||||||
|
what=f"NIXL create_backend({backend!r}, {backend_params})",
|
||||||
|
)
|
||||||
|
|
||||||
available_plugins = self.agent.get_plugin_list()
|
available_plugins = self.agent.get_plugin_list()
|
||||||
if backend not in available_plugins:
|
if backend not in available_plugins:
|
||||||
|
|||||||
@@ -11,6 +11,7 @@ from sglang.srt.runtime_context import (
|
|||||||
get_exec,
|
get_exec,
|
||||||
get_memory,
|
get_memory,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.utils.common import run_with_deadline
|
||||||
from sglang.srt.utils.network import NetworkAddress, get_free_port, get_local_ip_auto
|
from sglang.srt.utils.network import NetworkAddress, get_free_port, get_local_ip_auto
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -207,11 +208,15 @@ class MooncakeTransferEngine:
|
|||||||
# Default is "rdma"; set MOONCAKE_PROTOCOL=efa on AWS EFA hardware.
|
# Default is "rdma"; set MOONCAKE_PROTOCOL=efa on AWS EFA hardware.
|
||||||
protocol = envs.MOONCAKE_PROTOCOL.get()
|
protocol = envs.MOONCAKE_PROTOCOL.get()
|
||||||
|
|
||||||
ret_value = self.engine.initialize(
|
ret_value = run_with_deadline(
|
||||||
hostname,
|
lambda: self.engine.initialize(
|
||||||
"P2PHANDSHAKE",
|
hostname,
|
||||||
protocol,
|
"P2PHANDSHAKE",
|
||||||
device_name if device_name is not None else "",
|
protocol,
|
||||||
|
device_name if device_name is not None else "",
|
||||||
|
),
|
||||||
|
timeout_s=envs.SGLANG_DISAGGREGATION_ENGINE_INIT_TIMEOUT.get(),
|
||||||
|
what=f"Mooncake TransferEngine.initialize({hostname!r}, {protocol!r}, {device_name!r})",
|
||||||
)
|
)
|
||||||
if ret_value != 0:
|
if ret_value != 0:
|
||||||
logger.error("Mooncake Transfer Engine initialization failed.")
|
logger.error("Mooncake Transfer Engine initialization failed.")
|
||||||
|
|||||||
@@ -658,6 +658,8 @@ class Envs:
|
|||||||
SGLANG_DISAGGREGATION_HEARTBEAT_INTERVAL = EnvFloat(5.0)
|
SGLANG_DISAGGREGATION_HEARTBEAT_INTERVAL = EnvFloat(5.0)
|
||||||
SGLANG_DISAGGREGATION_HEARTBEAT_MAX_FAILURE = EnvInt(2)
|
SGLANG_DISAGGREGATION_HEARTBEAT_MAX_FAILURE = EnvInt(2)
|
||||||
SGLANG_DISAGGREGATION_WAITING_TIMEOUT = EnvInt(300)
|
SGLANG_DISAGGREGATION_WAITING_TIMEOUT = EnvInt(300)
|
||||||
|
# A wedged RDMA stack fails startup here instead of at the scheduler watchdog.
|
||||||
|
SGLANG_DISAGGREGATION_ENGINE_INIT_TIMEOUT = EnvInt(60)
|
||||||
SGLANG_DISAGGREGATION_NIXL_BACKEND = EnvStr("UCX")
|
SGLANG_DISAGGREGATION_NIXL_BACKEND = EnvStr("UCX")
|
||||||
SGLANG_DISAGGREGATION_NIXL_BACKEND_PARAMS = EnvStr("{}")
|
SGLANG_DISAGGREGATION_NIXL_BACKEND_PARAMS = EnvStr("{}")
|
||||||
SGLANG_DISAGG_PREFILL_EARLY_SEND_CACHED_PREFIX = EnvBool(True)
|
SGLANG_DISAGG_PREFILL_EARLY_SEND_CACHED_PREFIX = EnvBool(True)
|
||||||
|
|||||||
@@ -38,6 +38,7 @@ import re
|
|||||||
import resource
|
import resource
|
||||||
import shutil
|
import shutil
|
||||||
import signal
|
import signal
|
||||||
|
import socket
|
||||||
import subprocess
|
import subprocess
|
||||||
import sys
|
import sys
|
||||||
import tempfile
|
import tempfile
|
||||||
@@ -3460,6 +3461,29 @@ def parse_connector_type(url: str) -> str:
|
|||||||
return m.group(1)
|
return m.group(1)
|
||||||
|
|
||||||
|
|
||||||
|
def run_with_deadline(fn: Callable[[], Any], *, timeout_s: float, what: str) -> Any:
|
||||||
|
result: list = []
|
||||||
|
error: list = []
|
||||||
|
|
||||||
|
def _target():
|
||||||
|
try:
|
||||||
|
result.append(fn())
|
||||||
|
except BaseException as e:
|
||||||
|
error.append(e)
|
||||||
|
|
||||||
|
# An overrunning fn cannot be cancelled; only process exit reaps the daemon thread.
|
||||||
|
thread = threading.Thread(target=_target, daemon=True)
|
||||||
|
thread.start()
|
||||||
|
thread.join(timeout_s)
|
||||||
|
if thread.is_alive():
|
||||||
|
raise RuntimeError(
|
||||||
|
f"{what} did not return within {timeout_s}s on {socket.gethostname()}"
|
||||||
|
)
|
||||||
|
if error:
|
||||||
|
raise error[0]
|
||||||
|
return result[0]
|
||||||
|
|
||||||
|
|
||||||
def retry(
|
def retry(
|
||||||
fn,
|
fn,
|
||||||
max_retry: int,
|
max_retry: int,
|
||||||
|
|||||||
@@ -93,11 +93,13 @@ class PDDisaggregationServerBase(CustomTestCase):
|
|||||||
|
|
||||||
# config transfer backend and rdma devices
|
# config transfer backend and rdma devices
|
||||||
cls._mc_gid_index_set = False
|
cls._mc_gid_index_set = False
|
||||||
|
cls._ucx_net_devices_set = False
|
||||||
if is_in_ci():
|
if is_in_ci():
|
||||||
cls.transfer_backend = ["--disaggregation-transfer-backend", "mooncake"]
|
cls.transfer_backend = ["--disaggregation-transfer-backend", "mooncake"]
|
||||||
ib_devices = get_rdma_devices_args()
|
ib_devices = get_rdma_devices_args()
|
||||||
cls.rdma_devices = ["--disaggregation-ib-device", ib_devices]
|
cls.rdma_devices = ["--disaggregation-ib-device", ib_devices]
|
||||||
cls._mc_gid_index_set = _maybe_set_roce_gid_index(ib_devices)
|
cls._mc_gid_index_set = _maybe_set_roce_gid_index(ib_devices)
|
||||||
|
cls._ucx_net_devices_set = _maybe_set_ucx_net_devices(ib_devices)
|
||||||
else:
|
else:
|
||||||
cls.transfer_backend = [
|
cls.transfer_backend = [
|
||||||
"--disaggregation-transfer-backend",
|
"--disaggregation-transfer-backend",
|
||||||
@@ -226,6 +228,8 @@ class PDDisaggregationServerBase(CustomTestCase):
|
|||||||
os.environ.pop("MC_TCP_ENABLE_CONNECTION_POOL")
|
os.environ.pop("MC_TCP_ENABLE_CONNECTION_POOL")
|
||||||
if getattr(cls, "_mc_gid_index_set", False):
|
if getattr(cls, "_mc_gid_index_set", False):
|
||||||
os.environ.pop("MC_GID_INDEX", None)
|
os.environ.pop("MC_GID_INDEX", None)
|
||||||
|
if getattr(cls, "_ucx_net_devices_set", False):
|
||||||
|
os.environ.pop("UCX_NET_DEVICES", None)
|
||||||
# The LB holds no device state, and popen_with_error_check only stays
|
# The LB holds no device state, and popen_with_error_check only stays
|
||||||
# quiet for a SIGKILL rc, so hard-kill it rather than SIGTERM first.
|
# quiet for a SIGKILL rc, so hard-kill it rather than SIGTERM first.
|
||||||
if cls.process_lb:
|
if cls.process_lb:
|
||||||
@@ -501,3 +505,20 @@ def _maybe_set_roce_gid_index(ib_devices) -> bool:
|
|||||||
os.environ["MC_GID_INDEX"] = str(gid_index)
|
os.environ["MC_GID_INDEX"] = str(gid_index)
|
||||||
logger.warning("RoCE fabric detected; set MC_GID_INDEX=%d for mooncake", gid_index)
|
logger.warning("RoCE fabric detected; set MC_GID_INDEX=%d for mooncake", gid_index)
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
def _maybe_set_ucx_net_devices(ib_devices) -> bool:
|
||||||
|
if not ib_devices or os.environ.get("UCX_NET_DEVICES"):
|
||||||
|
return False
|
||||||
|
if ib_devices.lstrip().startswith("{"):
|
||||||
|
# Per-GPU JSON mapping; UCX_NET_DEVICES cannot express it.
|
||||||
|
return False
|
||||||
|
devices = [d.strip() for d in ib_devices.split(",") if d.strip()]
|
||||||
|
if not devices:
|
||||||
|
return False
|
||||||
|
net_devices = ",".join(f"{d}:1" for d in devices)
|
||||||
|
# NIXL ignores --disaggregation-ib-device; without this UCX opens every RDMA
|
||||||
|
# device on the host, and that full-device init can stall inside the driver.
|
||||||
|
os.environ["UCX_NET_DEVICES"] = net_devices
|
||||||
|
logger.warning("Set UCX_NET_DEVICES=%s for NIXL/UCX", net_devices)
|
||||||
|
return True
|
||||||
|
|||||||
Reference in New Issue
Block a user