Fix device context during NIXL backend initialization (#38774)

Co-authored-by: Aurick Qiao <6137920+aurickq@users.noreply.github.com>
This commit is contained in:
Aurick Qiao
2026-09-16 17:59:24 +08:00
committed by GitHub
co-authored by Aurick Qiao
parent a3bf25dc62
commit 7b7620774c
2 changed files with 110 additions and 2 deletions
@@ -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})",
)
@@ -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)