diff --git a/python/sglang/srt/disaggregation/nixl/conn.py b/python/sglang/srt/disaggregation/nixl/conn.py index 3705c1661..d4e2b3d1b 100644 --- a/python/sglang/srt/disaggregation/nixl/conn.py +++ b/python/sglang/srt/disaggregation/nixl/conn.py @@ -12,6 +12,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Set, Tuple import numpy as np import numpy.typing as npt +import torch import zmq if TYPE_CHECKING: @@ -48,7 +49,7 @@ from sglang.srt.disaggregation.utils import ( slice_dsa_tail_dst_ptrs_for_pp, ) from sglang.srt.environ import envs -from sglang.srt.runtime_context import get_parallel, get_schedule +from sglang.srt.runtime_context import get_device, get_parallel, get_schedule from sglang.srt.server_args import ServerArgs from sglang.srt.utils.common import run_with_deadline @@ -472,8 +473,13 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager): backend_params.setdefault("thread_count", str(num_threads)) elif backend == "UCCL": backend_params.setdefault("num_cpus", str(num_threads)) + + def create_backend(): + torch.get_device_module(get_device().device).set_device(self.kv_args.gpu_id) + return self.agent.create_backend(backend, backend_params) + run_with_deadline( - lambda: self.agent.create_backend(backend, backend_params), + create_backend, timeout_s=envs.SGLANG_DISAGGREGATION_ENGINE_INIT_TIMEOUT.get(), what=f"NIXL create_backend({backend!r}, {backend_params})", ) diff --git a/test/registered/unit/disaggregation/test_nixl_backend_basic.py b/test/registered/unit/disaggregation/test_nixl_backend_basic.py index d8ebcde29..2bf28ec9b 100644 --- a/test/registered/unit/disaggregation/test_nixl_backend_basic.py +++ b/test/registered/unit/disaggregation/test_nixl_backend_basic.py @@ -24,6 +24,8 @@ from sglang.srt.disaggregation.nixl.conn import ( TransferKVChunk, TransferStatus, ) +from sglang.srt.disaggregation.utils import DisaggregationMode +from sglang.srt.runtime_context import get_context from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import CustomTestCase @@ -106,6 +108,106 @@ def _fake_staging_buffer_module(mock_gather=None): return module +class TestNixlBackendInitialization(CustomTestCase): + def _initialize_manager(self, agent, device_module, mode, gpu_id): + args = SimpleNamespace( + pp_rank=0, + engine_rank=0, + gpu_id=gpu_id, + kv_data_ptrs=[], + kv_item_lens=[], + ) + mgr = object.__new__(NixlKVManager) + mgr.kv_args = args + mgr.disaggregation_mode = mode + mgr.enable_deferred_decode_kv_release = False + + api = types.ModuleType("nixl._api") + api.nixl_agent = MagicMock(return_value=agent) + api.nixl_agent_config = MagicMock() + api.nixl_thread_sync_t = SimpleNamespace(NIXL_THREAD_SYNC_STRICT="strict") + nixl = types.ModuleType("nixl") + nixl._api = api + agent.get_plugin_list.return_value = ["UCX"] + + with ( + patch.dict(sys.modules, {"nixl": nixl, "nixl._api": api}), + patch.dict( + "os.environ", + { + "SGLANG_DISAGGREGATION_NIXL_BACKEND": "UCX", + "SGLANG_DISAGGREGATION_NIXL_BACKEND_PARAMS": "{}", + "SGLANG_DISAGGREGATION_ENGINE_INIT_TIMEOUT": "5", + "SGLANG_DISAGGREGATION_QUEUE_SIZE": "0", + "SGLANG_DISAGG_STAGING_BUFFER": "false", + }, + ), + get_context().override_server_args(device="cuda") as server_args, + patch.object(CommonKVManager, "__init__", return_value=None), + patch( + "sglang.srt.disaggregation.nixl.conn.get_parallel", + return_value=SimpleNamespace(tp_size=4), + ), + patch( + "sglang.srt.disaggregation.nixl.conn.torch.get_device_module", + return_value=device_module, + ) as get_device_module, + patch.object(NixlKVManager, "register_buffer_to_engine"), + patch.object(NixlKVManager, "_start_bootstrap_thread"), + patch.object(NixlKVManager, "_start_decode_listener_thread"), + patch.object(NixlKVManager, "_start_heartbeat_checker_thread"), + ): + self.assertIsNone(server_args.device) + NixlKVManager.__init__(mgr, args, mode, server_args) + get_device_module.assert_called_once_with("cuda") + + def test_backend_initialization_selects_device_in_deadline_thread(self): + caller_thread = threading.get_ident() + for mode, gpu_id in ( + (DisaggregationMode.PREFILL, 3), + (DisaggregationMode.DECODE, 1), + ): + with self.subTest(mode=mode, gpu_id=gpu_id): + calls = [] + device_module = MagicMock() + device_module.set_device.side_effect = lambda device: calls.append( + ("set_device", device, threading.get_ident()) + ) + agent = MagicMock() + agent.create_backend.side_effect = lambda backend, params: calls.append( + ("create_backend", backend, threading.get_ident()) + ) + + self._initialize_manager(agent, device_module, mode, gpu_id) + + self.assertEqual(len(calls), 2) + backend_thread = calls[1][2] + self.assertNotEqual(backend_thread, caller_thread) + self.assertEqual( + calls, + [ + ("set_device", gpu_id, backend_thread), + ("create_backend", "UCX", backend_thread), + ], + ) + agent.create_backend.assert_called_once_with( + "UCX", + {"num_threads": "8"} if mode == DisaggregationMode.PREFILL else {}, + ) + + def test_backend_initialization_propagates_error(self): + error = RuntimeError("backend initialization failed") + agent = MagicMock() + agent.create_backend.side_effect = error + + with self.assertRaises(RuntimeError) as raised: + self._initialize_manager( + agent, MagicMock(), DisaggregationMode.DECODE, gpu_id=3 + ) + + self.assertIs(raised.exception, error) + + class TestNixlTransferInfo(CustomTestCase): def test_from_zmq_parses_required_fields(self): kv_indices = np.array([3, 5, 8], dtype=np.int32)