weight cache: key daemon paths by GPU UUID (#36101)
Co-authored-by: siyu <liusy58@linux.alibaba.com> Co-authored-by: Alex Nails <alex.nails@radixark.ai>
This commit is contained in:
co-authored by
siyu
Alex Nails
parent
3865efc9f7
commit
6580d5cd9a
@@ -7,7 +7,9 @@ import unittest
|
||||
import requests
|
||||
import torch
|
||||
|
||||
from sglang.srt.platforms import current_platform
|
||||
from sglang.srt.utils import kill_process_tree
|
||||
from sglang.srt.weight_cache.protocol import get_ready_path, get_socket_path
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
from sglang.test.test_utils import (
|
||||
DEFAULT_TARGET_MODEL_EAGLE_DP_ATTN,
|
||||
@@ -44,6 +46,11 @@ PROMPTS = [
|
||||
]
|
||||
|
||||
|
||||
def _gpu_uuids(tp_size: int) -> list:
|
||||
# Single-node, default base_gpu_id/gpu_id_step: rank i runs on physical GPU i.
|
||||
return [current_platform.get_device_uuid(i) for i in range(tp_size)]
|
||||
|
||||
|
||||
@unittest.skipIf(
|
||||
torch.cuda.device_count() < 2,
|
||||
"TP=2 weight cache daemon test requires >=2 GPUs (skipped on the 1-gpu runner)",
|
||||
@@ -60,11 +67,11 @@ class TestWeightCacheDaemonTP2(CustomTestCase):
|
||||
cls.model = cls.model_override or DEFAULT_MODEL
|
||||
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||
cls.tp_size = 2
|
||||
cls.gpu_uuids = _gpu_uuids(cls.tp_size)
|
||||
|
||||
# Clean up stale ready/socket files from previous runs
|
||||
for rank in range(cls.tp_size):
|
||||
for suffix in (".ready", ".sock"):
|
||||
path = f"/tmp/sglang_weight_cache_rank{rank}{suffix}"
|
||||
for device_uuid in cls.gpu_uuids:
|
||||
for path in (get_ready_path(device_uuid), get_socket_path(device_uuid)):
|
||||
if os.path.exists(path):
|
||||
os.unlink(path)
|
||||
|
||||
@@ -86,13 +93,14 @@ class TestWeightCacheDaemonTP2(CustomTestCase):
|
||||
# Step 2: Wait for all daemon ready files
|
||||
timeout = DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH
|
||||
start = time.time()
|
||||
for rank in range(cls.tp_size):
|
||||
ready_path = f"/tmp/sglang_weight_cache_rank{rank}.ready"
|
||||
for device_uuid in cls.gpu_uuids:
|
||||
ready_path = get_ready_path(device_uuid)
|
||||
while not os.path.exists(ready_path):
|
||||
if time.time() - start > timeout:
|
||||
kill_process_tree(cls.daemon_process.pid)
|
||||
raise TimeoutError(
|
||||
f"Weight cache daemon rank {rank} not ready within {timeout}s"
|
||||
f"Weight cache daemon for GPU {device_uuid} not ready "
|
||||
f"within {timeout}s"
|
||||
)
|
||||
if cls.daemon_process.poll() is not None:
|
||||
raise RuntimeError(
|
||||
@@ -136,9 +144,8 @@ class TestWeightCacheDaemonTP2(CustomTestCase):
|
||||
os.unlink(path)
|
||||
except OSError:
|
||||
pass
|
||||
for rank in range(getattr(cls, "tp_size", 2)):
|
||||
for suffix in (".ready", ".sock"):
|
||||
path = f"/tmp/sglang_weight_cache_rank{rank}{suffix}"
|
||||
for device_uuid in getattr(cls, "gpu_uuids", ()):
|
||||
for path in (get_ready_path(device_uuid), get_socket_path(device_uuid)):
|
||||
if os.path.exists(path):
|
||||
try:
|
||||
os.unlink(path)
|
||||
@@ -220,11 +227,11 @@ class TestWeightCacheDaemonTP1Smoke(CustomTestCase):
|
||||
cls.model = DEFAULT_MODEL
|
||||
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||
cls.tp_size = 1
|
||||
cls.gpu_uuids = _gpu_uuids(cls.tp_size)
|
||||
|
||||
# Clean up stale ready/socket files from previous runs.
|
||||
for rank in range(cls.tp_size):
|
||||
for suffix in (".ready", ".sock"):
|
||||
path = f"/tmp/sglang_weight_cache_rank{rank}{suffix}"
|
||||
for device_uuid in cls.gpu_uuids:
|
||||
for path in (get_ready_path(device_uuid), get_socket_path(device_uuid)):
|
||||
if os.path.exists(path):
|
||||
os.unlink(path)
|
||||
|
||||
@@ -245,13 +252,14 @@ class TestWeightCacheDaemonTP1Smoke(CustomTestCase):
|
||||
# Step 2: Wait for the daemon ready file.
|
||||
timeout = DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH
|
||||
start = time.time()
|
||||
for rank in range(cls.tp_size):
|
||||
ready_path = f"/tmp/sglang_weight_cache_rank{rank}.ready"
|
||||
for device_uuid in cls.gpu_uuids:
|
||||
ready_path = get_ready_path(device_uuid)
|
||||
while not os.path.exists(ready_path):
|
||||
if time.time() - start > timeout:
|
||||
kill_process_tree(cls.daemon_process.pid)
|
||||
raise TimeoutError(
|
||||
f"Weight cache daemon rank {rank} not ready within {timeout}s"
|
||||
f"Weight cache daemon for GPU {device_uuid} not ready "
|
||||
f"within {timeout}s"
|
||||
)
|
||||
if cls.daemon_process.poll() is not None:
|
||||
raise RuntimeError(
|
||||
@@ -294,9 +302,8 @@ class TestWeightCacheDaemonTP1Smoke(CustomTestCase):
|
||||
os.unlink(path)
|
||||
except OSError:
|
||||
pass
|
||||
for rank in range(getattr(cls, "tp_size", 1)):
|
||||
for suffix in (".ready", ".sock"):
|
||||
path = f"/tmp/sglang_weight_cache_rank{rank}{suffix}"
|
||||
for device_uuid in getattr(cls, "gpu_uuids", ()):
|
||||
for path in (get_ready_path(device_uuid), get_socket_path(device_uuid)):
|
||||
if os.path.exists(path):
|
||||
try:
|
||||
os.unlink(path)
|
||||
|
||||
@@ -5,7 +5,7 @@ These cover the pure-Python logic that the GPU end-to-end test
|
||||
(test_weight_cache_daemon.py) cannot exercise cheaply:
|
||||
|
||||
- length-prefixed socket framing (send_msg/recv_msg) over socketpair()
|
||||
- CacheConfig fingerprint matching / (de)serialization
|
||||
- CacheConfig compatibility matching / (de)serialization
|
||||
- quant-config hashing and method-name extraction
|
||||
- daemon spawn configuration and socket/ready path derivation
|
||||
- the IPC quantization allowlist (the gate that keeps silently-wrong
|
||||
@@ -25,6 +25,7 @@ from types import SimpleNamespace
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.weight_cache.protocol import (
|
||||
IPC_QUANT_ALLOWLIST,
|
||||
CacheConfig,
|
||||
@@ -240,10 +241,25 @@ class TestGlobalRankAndPaths(CustomTestCase):
|
||||
self.assertEqual(compute_global_rank(tp_size=4, pp_rank=1, tp_rank=0), 4)
|
||||
self.assertEqual(compute_global_rank(tp_size=4, pp_rank=2, tp_rank=1), 9)
|
||||
|
||||
def test_socket_and_ready_paths_are_unique_per_rank(self):
|
||||
self.assertNotEqual(get_socket_path(0), get_socket_path(1))
|
||||
self.assertTrue(get_socket_path(3).endswith("rank3.sock"))
|
||||
self.assertTrue(get_ready_path(3).endswith("rank3.ready"))
|
||||
def test_socket_and_ready_paths_are_unique_per_device_uuid(self):
|
||||
self.assertNotEqual(get_socket_path("gpu-aaa"), get_socket_path("gpu-bbb"))
|
||||
self.assertTrue(get_socket_path("gpu-aaa").endswith("gpu-aaa.sock"))
|
||||
self.assertTrue(get_ready_path("gpu-aaa").endswith("gpu-aaa.ready"))
|
||||
|
||||
def test_custom_template_with_device_uuid_placeholder_is_honored(self):
|
||||
with envs.SGLANG_WEIGHT_CACHE_SOCKET_TEMPLATE.override(
|
||||
"/custom/dir/{device_uuid}.custom-sock"
|
||||
):
|
||||
self.assertEqual(
|
||||
get_socket_path("gpu-aaa"), "/custom/dir/gpu-aaa.custom-sock"
|
||||
)
|
||||
|
||||
def test_template_missing_device_uuid_placeholder_raises(self):
|
||||
with envs.SGLANG_WEIGHT_CACHE_SOCKET_TEMPLATE.override(
|
||||
"/tmp/sglang_weight_cache_rank{global_rank}.sock"
|
||||
):
|
||||
with self.assertRaises(ValueError):
|
||||
get_socket_path("gpu-aaa")
|
||||
|
||||
def test_compute_local_gpu_id_honors_base_and_step(self):
|
||||
# Single-node TP=4: identity mapping rank -> gpu.
|
||||
@@ -368,12 +384,12 @@ class TestIpcQuantAllowlist(CustomTestCase):
|
||||
|
||||
|
||||
class TestCleanupStaleDaemonFiles(CustomTestCase):
|
||||
# Use a rank far outside any realistic tp*pp layout so we never collide
|
||||
# with a daemon that might actually be running on the test host.
|
||||
RANK = 987654
|
||||
# A key no real daemon would ever compute, so this never collides with
|
||||
# one that might actually be running on the test host.
|
||||
KEY = "test-cleanup-stale-daemon-files"
|
||||
|
||||
def _paths(self):
|
||||
return get_ready_path(self.RANK), get_socket_path(self.RANK)
|
||||
return get_ready_path(self.KEY), get_socket_path(self.KEY)
|
||||
|
||||
def tearDown(self):
|
||||
for path in self._paths():
|
||||
@@ -382,7 +398,7 @@ class TestCleanupStaleDaemonFiles(CustomTestCase):
|
||||
|
||||
def test_no_files_is_noop(self):
|
||||
# Neither file present: must return quietly, not raise.
|
||||
cleanup_stale_daemon_files(self.RANK)
|
||||
cleanup_stale_daemon_files(self.KEY)
|
||||
|
||||
def test_stale_files_without_live_pid_are_removed(self):
|
||||
ready_path, socket_path = self._paths()
|
||||
@@ -392,7 +408,7 @@ class TestCleanupStaleDaemonFiles(CustomTestCase):
|
||||
f.write("stale contents, no pid line\n")
|
||||
open(socket_path, "w").close()
|
||||
|
||||
cleanup_stale_daemon_files(self.RANK)
|
||||
cleanup_stale_daemon_files(self.KEY)
|
||||
|
||||
self.assertFalse(os.path.exists(ready_path))
|
||||
self.assertFalse(os.path.exists(socket_path))
|
||||
@@ -405,7 +421,7 @@ class TestCleanupStaleDaemonFiles(CustomTestCase):
|
||||
open(socket_path, "w").close()
|
||||
|
||||
with self.assertRaises(RuntimeError):
|
||||
cleanup_stale_daemon_files(self.RANK)
|
||||
cleanup_stale_daemon_files(self.KEY)
|
||||
|
||||
self.assertTrue(os.path.exists(ready_path))
|
||||
self.assertTrue(os.path.exists(socket_path))
|
||||
@@ -423,11 +439,11 @@ class TestCleanupStaleDaemonFiles(CustomTestCase):
|
||||
f.write(f"pid={child.pid}\n")
|
||||
open(socket_path, "w").close()
|
||||
|
||||
cleanup_stale_daemon_files(self.RANK, force=True)
|
||||
cleanup_stale_daemon_files(self.KEY, force=True)
|
||||
|
||||
self.assertFalse(os.path.exists(ready_path))
|
||||
self.assertFalse(os.path.exists(socket_path))
|
||||
# The daemon holding the rank must have been killed.
|
||||
# The daemon holding the identity must have been killed.
|
||||
self.assertEqual(child.wait(timeout=5), -9)
|
||||
finally:
|
||||
if child.poll() is None:
|
||||
@@ -442,7 +458,7 @@ class TestDaemonModeRefusesDiskLoad(CustomTestCase):
|
||||
pointing the loader at a socket path that does not exist.
|
||||
"""
|
||||
|
||||
RANK = 987655
|
||||
KEY = "test-daemon-mode-refuses-disk-load"
|
||||
|
||||
def _model_config(self):
|
||||
from types import SimpleNamespace
|
||||
@@ -464,7 +480,7 @@ class TestDaemonModeRefusesDiskLoad(CustomTestCase):
|
||||
from sglang.srt.configs.load_config import LoadConfig, LoadFormat
|
||||
from sglang.srt.weight_cache.ipc_loader import IpcModelLoader
|
||||
|
||||
missing_socket = get_socket_path(self.RANK)
|
||||
missing_socket = get_socket_path(self.KEY)
|
||||
if os.path.exists(missing_socket):
|
||||
os.unlink(missing_socket)
|
||||
|
||||
@@ -481,6 +497,30 @@ class TestDaemonModeRefusesDiskLoad(CustomTestCase):
|
||||
# fall through to a disk load.
|
||||
self.assertIn("daemon", str(ctx.exception).lower())
|
||||
|
||||
def test_default_discovery_queries_device_uuid_for_the_caller_gpu(self):
|
||||
"""Without an explicit --weight-cache-socket (the production default),
|
||||
discovery must derive the socket from the caller's own gpu_id via
|
||||
get_device_uuid -- not silently substitute GPU 0."""
|
||||
from unittest import mock
|
||||
|
||||
from sglang.srt.configs.load_config import LoadConfig, LoadFormat
|
||||
from sglang.srt.weight_cache.ipc_loader import IpcModelLoader
|
||||
|
||||
loader = IpcModelLoader(
|
||||
load_config=LoadConfig(load_format=LoadFormat.IPC_CACHE),
|
||||
weight_cache_mode="client",
|
||||
fallback_load_format="auto",
|
||||
)
|
||||
with mock.patch(
|
||||
"sglang.srt.platforms.current_platform.get_device_uuid",
|
||||
return_value="gpu-under-test",
|
||||
) as get_uuid:
|
||||
result = loader._fetch_from_cache(
|
||||
self._model_config(), SimpleNamespace(gpu_id=5)
|
||||
)
|
||||
get_uuid.assert_called_once_with(5)
|
||||
self.assertIsNone(result) # no real daemon at that socket -> absent
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user