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
@@ -101,6 +101,7 @@ from sglang.srt.observability.startup_time import build_engine_startup_time
|
|||||||
from sglang.srt.observability.trace import process_tracing_init, trace_set_thread_info
|
from sglang.srt.observability.trace import process_tracing_init, trace_set_thread_info
|
||||||
from sglang.srt.parser.template_detection import resolve_auto_parsers
|
from sglang.srt.parser.template_detection import resolve_auto_parsers
|
||||||
from sglang.srt.parser.template_manager import TemplateManager
|
from sglang.srt.parser.template_manager import TemplateManager
|
||||||
|
from sglang.srt.platforms import current_platform
|
||||||
from sglang.srt.plugins import load_plugins
|
from sglang.srt.plugins import load_plugins
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
get_disagg,
|
get_disagg,
|
||||||
@@ -138,6 +139,12 @@ from sglang.srt.utils.network import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
|
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
|
||||||
from sglang.srt.utils.watchdog import SubprocessWatchdog
|
from sglang.srt.utils.watchdog import SubprocessWatchdog
|
||||||
|
from sglang.srt.weight_cache.daemon import spawn_weight_cache_daemon
|
||||||
|
from sglang.srt.weight_cache.protocol import (
|
||||||
|
cleanup_stale_daemon_files,
|
||||||
|
compute_local_gpu_id,
|
||||||
|
get_ready_path,
|
||||||
|
)
|
||||||
from sglang.version import __version__
|
from sglang.version import __version__
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -715,19 +722,18 @@ class Engine(EngineScoreMixin, EngineBase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# Validate and clean up stale .ready/.sock files from prior runs.
|
# Validate and clean up stale .ready/.sock files from prior runs.
|
||||||
# If a daemon is still alive at this rank, raise instead of clobbering.
|
# If a daemon is still alive at this GPU, raise instead of clobbering.
|
||||||
from sglang.srt.weight_cache.daemon import spawn_weight_cache_daemon
|
|
||||||
from sglang.srt.weight_cache.protocol import (
|
|
||||||
cleanup_stale_daemon_files,
|
|
||||||
compute_global_rank,
|
|
||||||
compute_local_gpu_id,
|
|
||||||
get_ready_path,
|
|
||||||
)
|
|
||||||
|
|
||||||
for pp_rank in pp_rank_range:
|
for pp_rank in pp_rank_range:
|
||||||
for tp_rank in tp_rank_range:
|
for tp_rank in tp_rank_range:
|
||||||
global_rank = compute_global_rank(tp_size, pp_rank, tp_rank)
|
gpu_id = compute_local_gpu_id(
|
||||||
cleanup_stale_daemon_files(global_rank)
|
pp_rank,
|
||||||
|
tp_rank,
|
||||||
|
pp_size_per_node,
|
||||||
|
tp_size_per_node,
|
||||||
|
base_gpu_id=server_args.base_gpu_id,
|
||||||
|
gpu_id_step=server_args.gpu_id_step,
|
||||||
|
)
|
||||||
|
cleanup_stale_daemon_files(current_platform.get_device_uuid(gpu_id))
|
||||||
|
|
||||||
for pp_rank in pp_rank_range:
|
for pp_rank in pp_rank_range:
|
||||||
for tp_rank in tp_rank_range:
|
for tp_rank in tp_rank_range:
|
||||||
@@ -759,8 +765,17 @@ class Engine(EngineScoreMixin, EngineBase):
|
|||||||
try:
|
try:
|
||||||
for pp_rank in pp_rank_range:
|
for pp_rank in pp_rank_range:
|
||||||
for tp_rank in tp_rank_range:
|
for tp_rank in tp_rank_range:
|
||||||
global_rank = compute_global_rank(tp_size, pp_rank, tp_rank)
|
gpu_id = compute_local_gpu_id(
|
||||||
ready_path = get_ready_path(global_rank)
|
pp_rank,
|
||||||
|
tp_rank,
|
||||||
|
pp_size_per_node,
|
||||||
|
tp_size_per_node,
|
||||||
|
base_gpu_id=server_args.base_gpu_id,
|
||||||
|
gpu_id_step=server_args.gpu_id_step,
|
||||||
|
)
|
||||||
|
ready_path = get_ready_path(
|
||||||
|
current_platform.get_device_uuid(gpu_id)
|
||||||
|
)
|
||||||
while not os.path.exists(ready_path):
|
while not os.path.exists(ready_path):
|
||||||
time.sleep(check_interval)
|
time.sleep(check_interval)
|
||||||
if time.time() - start_time > timeout:
|
if time.time() - start_time > timeout:
|
||||||
|
|||||||
@@ -1612,14 +1612,14 @@ class Envs:
|
|||||||
# Weight Cache Daemon
|
# Weight Cache Daemon
|
||||||
# ===================================================================
|
# ===================================================================
|
||||||
# Paths the daemon and the engine ranks it serves must agree on. Both are
|
# Paths the daemon and the engine ranks it serves must agree on. Both are
|
||||||
# format templates and must keep the {global_rank} placeholder: each rank
|
# format templates and must keep the {device_uuid} placeholder: each daemon
|
||||||
# talks to the daemon on its own GPU, so a rank-independent path would point
|
# is keyed by the physical GPU it runs on, so a GPU-independent path would
|
||||||
# every rank at one daemon and map another rank's shard.
|
# let one job's client discover another job's daemon.
|
||||||
SGLANG_WEIGHT_CACHE_SOCKET_TEMPLATE = EnvStr(
|
SGLANG_WEIGHT_CACHE_SOCKET_TEMPLATE = EnvStr(
|
||||||
"/tmp/sglang_weight_cache_rank{global_rank}.sock"
|
"/tmp/sglang_weight_cache_{device_uuid}.sock"
|
||||||
)
|
)
|
||||||
SGLANG_WEIGHT_CACHE_READY_TEMPLATE = EnvStr(
|
SGLANG_WEIGHT_CACHE_READY_TEMPLATE = EnvStr(
|
||||||
"/tmp/sglang_weight_cache_rank{global_rank}.ready"
|
"/tmp/sglang_weight_cache_{device_uuid}.ready"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1145,13 +1145,8 @@ class ModelRunner:
|
|||||||
weight_cache_socket=get_model().weight_cache_socket,
|
weight_cache_socket=get_model().weight_cache_socket,
|
||||||
)
|
)
|
||||||
|
|
||||||
# If the weight cache is enabled, override the load format to IPC_CACHE
|
|
||||||
# and derive the per-rank daemon socket. Idempotent across reloads.
|
|
||||||
maybe_enable_ipc_weight_cache(
|
maybe_enable_ipc_weight_cache(
|
||||||
load_config=self.load_config,
|
load_config=self.load_config,
|
||||||
tp_size=self.ps.tp_size,
|
|
||||||
pp_rank=self.ps.pp_rank,
|
|
||||||
tp_rank=self.ps.tp_rank,
|
|
||||||
)
|
)
|
||||||
if self.device == "cpu":
|
if self.device == "cpu":
|
||||||
self.model_config = adjust_config_with_unaligned_cpu_tp(
|
self.model_config = adjust_config_with_unaligned_cpu_tp(
|
||||||
|
|||||||
@@ -236,16 +236,13 @@ def build_load_config(
|
|||||||
def maybe_enable_ipc_weight_cache(
|
def maybe_enable_ipc_weight_cache(
|
||||||
*,
|
*,
|
||||||
load_config: LoadConfig,
|
load_config: LoadConfig,
|
||||||
tp_size: int,
|
|
||||||
pp_rank: int,
|
|
||||||
tp_rank: int,
|
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Switch ``load_config`` onto the IPC weight-cache path, in place.
|
"""Switch ``load_config`` onto the IPC weight-cache path, in place.
|
||||||
|
|
||||||
Overrides the load format to ``IPC_CACHE`` (remembering the original as the
|
Overrides the load format to ``IPC_CACHE`` (remembering the original as the
|
||||||
disk fallback) and derives the per-rank daemon socket if unset. Idempotent:
|
disk fallback). Idempotent: the format swap is guarded on ``!= IPC_CACHE``
|
||||||
the format swap is guarded on ``!= IPC_CACHE`` so a second call (e.g. a
|
so a second call (e.g. a weight reload) can't overwrite the captured
|
||||||
weight reload) can't overwrite the captured fallback format.
|
fallback format.
|
||||||
"""
|
"""
|
||||||
if get_model().weight_cache_mode == "off":
|
if get_model().weight_cache_mode == "off":
|
||||||
return
|
return
|
||||||
@@ -254,17 +251,6 @@ def maybe_enable_ipc_weight_cache(
|
|||||||
load_config.fallback_load_format = load_config.load_format
|
load_config.fallback_load_format = load_config.load_format
|
||||||
load_config.load_format = LoadFormat.IPC_CACHE
|
load_config.load_format = LoadFormat.IPC_CACHE
|
||||||
|
|
||||||
# Compute socket path using global rank (tp_size * pp_rank + tp_rank) so
|
|
||||||
# each daemon has a unique socket even across PP stages and nodes.
|
|
||||||
if load_config.weight_cache_socket is None:
|
|
||||||
from sglang.srt.weight_cache.protocol import (
|
|
||||||
compute_global_rank,
|
|
||||||
get_socket_path,
|
|
||||||
)
|
|
||||||
|
|
||||||
global_rank = compute_global_rank(tp_size, pp_rank, tp_rank)
|
|
||||||
load_config.weight_cache_socket = get_socket_path(global_rank=global_rank)
|
|
||||||
|
|
||||||
|
|
||||||
def load_model_with_memory_saver(
|
def load_model_with_memory_saver(
|
||||||
*,
|
*,
|
||||||
|
|||||||
@@ -4379,21 +4379,10 @@ def get_model_loader(
|
|||||||
|
|
||||||
if load_config.load_format == LoadFormat.IPC_CACHE:
|
if load_config.load_format == LoadFormat.IPC_CACHE:
|
||||||
from sglang.srt.weight_cache.ipc_loader import IpcModelLoader
|
from sglang.srt.weight_cache.ipc_loader import IpcModelLoader
|
||||||
from sglang.srt.weight_cache.protocol import (
|
|
||||||
compute_global_rank,
|
|
||||||
get_socket_path,
|
|
||||||
)
|
|
||||||
|
|
||||||
if load_config.weight_cache_socket:
|
|
||||||
socket_path = load_config.weight_cache_socket
|
|
||||||
else:
|
|
||||||
|
|
||||||
ps = get_parallel()
|
|
||||||
global_rank = compute_global_rank(ps.tp_size, ps.pp_rank, ps.tp_rank)
|
|
||||||
socket_path = get_socket_path(global_rank=global_rank)
|
|
||||||
return IpcModelLoader(
|
return IpcModelLoader(
|
||||||
load_config=load_config,
|
load_config=load_config,
|
||||||
socket_path=socket_path,
|
socket_path=load_config.weight_cache_socket,
|
||||||
weight_cache_mode=load_config.weight_cache_mode,
|
weight_cache_mode=load_config.weight_cache_mode,
|
||||||
fallback_load_format=load_config.fallback_load_format,
|
fallback_load_format=load_config.fallback_load_format,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -3619,7 +3619,8 @@ class ServerArgs:
|
|||||||
Optional[str],
|
Optional[str],
|
||||||
Arg(
|
Arg(
|
||||||
help="Unix socket path for weight cache daemon (client mode)."
|
help="Unix socket path for weight cache daemon (client mode)."
|
||||||
"If not set, uses /tmp/sglang_weight_cache_rank{global_rank}.sock",
|
"If not set, derives the path from SGLANG_WEIGHT_CACHE_SOCKET_TEMPLATE "
|
||||||
|
"using the caller's physical GPU UUID.",
|
||||||
),
|
),
|
||||||
NS("model"),
|
NS("model"),
|
||||||
] = None
|
] = None
|
||||||
|
|||||||
@@ -173,12 +173,9 @@ class WeightCacheDaemon:
|
|||||||
self.revision = cfg.revision
|
self.revision = cfg.revision
|
||||||
self.dist_init_method = dist_init_method
|
self.dist_init_method = dist_init_method
|
||||||
|
|
||||||
self.socket_path = get_socket_path(
|
device_uuid = current_platform.get_device_uuid(gpu_id)
|
||||||
compute_global_rank(self.tp_size, pp_rank, tp_rank)
|
self.socket_path = get_socket_path(device_uuid)
|
||||||
)
|
self.ready_path = get_ready_path(device_uuid)
|
||||||
self.ready_path = get_ready_path(
|
|
||||||
compute_global_rank(self.tp_size, pp_rank, tp_rank)
|
|
||||||
)
|
|
||||||
|
|
||||||
self.model = None
|
self.model = None
|
||||||
self.config: Optional[CacheConfig] = None
|
self.config: Optional[CacheConfig] = None
|
||||||
@@ -735,8 +732,17 @@ def launch_weight_cache_daemons(
|
|||||||
# Validate and clean up stale .ready/.sock files from prior runs.
|
# Validate and clean up stale .ready/.sock files from prior runs.
|
||||||
for pp_rank in pp_rank_range:
|
for pp_rank in pp_rank_range:
|
||||||
for tp_rank in tp_rank_range:
|
for tp_rank in tp_rank_range:
|
||||||
global_rank = compute_global_rank(cfg.tp_size, pp_rank, tp_rank)
|
gpu_id = compute_local_gpu_id(
|
||||||
cleanup_stale_daemon_files(global_rank, force=force)
|
pp_rank,
|
||||||
|
tp_rank,
|
||||||
|
pp_size_per_node,
|
||||||
|
tp_size_per_node,
|
||||||
|
base_gpu_id=cfg.base_gpu_id,
|
||||||
|
gpu_id_step=cfg.gpu_id_step,
|
||||||
|
)
|
||||||
|
cleanup_stale_daemon_files(
|
||||||
|
current_platform.get_device_uuid(gpu_id), force=force
|
||||||
|
)
|
||||||
|
|
||||||
procs = []
|
procs = []
|
||||||
for pp_rank in pp_rank_range:
|
for pp_rank in pp_rank_range:
|
||||||
@@ -768,8 +774,15 @@ def launch_weight_cache_daemons(
|
|||||||
start_time = time.time()
|
start_time = time.time()
|
||||||
for pp_rank in pp_rank_range:
|
for pp_rank in pp_rank_range:
|
||||||
for tp_rank in tp_rank_range:
|
for tp_rank in tp_rank_range:
|
||||||
global_rank = compute_global_rank(cfg.tp_size, pp_rank, tp_rank)
|
gpu_id = compute_local_gpu_id(
|
||||||
ready_path = get_ready_path(global_rank)
|
pp_rank,
|
||||||
|
tp_rank,
|
||||||
|
pp_size_per_node,
|
||||||
|
tp_size_per_node,
|
||||||
|
base_gpu_id=cfg.base_gpu_id,
|
||||||
|
gpu_id_step=cfg.gpu_id_step,
|
||||||
|
)
|
||||||
|
ready_path = get_ready_path(current_platform.get_device_uuid(gpu_id))
|
||||||
while not os.path.exists(ready_path):
|
while not os.path.exists(ready_path):
|
||||||
time.sleep(check_interval)
|
time.sleep(check_interval)
|
||||||
if time.time() - start_time > timeout:
|
if time.time() - start_time > timeout:
|
||||||
@@ -863,7 +876,7 @@ if __name__ == "__main__":
|
|||||||
else daemon_args.gpu_id
|
else daemon_args.gpu_id
|
||||||
)
|
)
|
||||||
cleanup_stale_daemon_files(
|
cleanup_stale_daemon_files(
|
||||||
compute_global_rank(server_args.tp_size, daemon_args.pp_rank, tp_rank),
|
current_platform.get_device_uuid(gpu_id),
|
||||||
force=daemon_args.force,
|
force=daemon_args.force,
|
||||||
)
|
)
|
||||||
run_weight_cache_daemon(
|
run_weight_cache_daemon(
|
||||||
|
|||||||
@@ -22,6 +22,7 @@ from sglang.srt.model_loader.loader import (
|
|||||||
BaseModelLoader,
|
BaseModelLoader,
|
||||||
_initialize_model,
|
_initialize_model,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.platforms import current_platform
|
||||||
from sglang.srt.runtime_context import get_exec, get_parallel
|
from sglang.srt.runtime_context import get_exec, get_parallel
|
||||||
|
|
||||||
from .protocol import (
|
from .protocol import (
|
||||||
@@ -29,6 +30,7 @@ from .protocol import (
|
|||||||
check_ipc_quant_support,
|
check_ipc_quant_support,
|
||||||
compute_env_stamp,
|
compute_env_stamp,
|
||||||
get_quant_method_name,
|
get_quant_method_name,
|
||||||
|
get_socket_path,
|
||||||
hash_quant_config,
|
hash_quant_config,
|
||||||
recv_msg,
|
recv_msg,
|
||||||
send_msg,
|
send_msg,
|
||||||
@@ -67,7 +69,7 @@ class IpcModelLoader(BaseModelLoader):
|
|||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
load_config: LoadConfig,
|
load_config: LoadConfig,
|
||||||
socket_path: str,
|
socket_path: Optional[str] = None,
|
||||||
fallback_loader_cls=None,
|
fallback_loader_cls=None,
|
||||||
weight_cache_mode: str = "client",
|
weight_cache_mode: str = "client",
|
||||||
fallback_load_format: str = "auto",
|
fallback_load_format: str = "auto",
|
||||||
@@ -103,7 +105,7 @@ class IpcModelLoader(BaseModelLoader):
|
|||||||
check_ipc_quant_support(quant_method, engine_quant_config, where="client")
|
check_ipc_quant_support(quant_method, engine_quant_config, where="client")
|
||||||
|
|
||||||
# Try to fetch state from daemon
|
# Try to fetch state from daemon
|
||||||
cache_data = self._fetch_from_cache(model_config)
|
cache_data = self._fetch_from_cache(model_config, device_config)
|
||||||
|
|
||||||
if cache_data is None:
|
if cache_data is None:
|
||||||
if self.weight_cache_mode == "daemon":
|
if self.weight_cache_mode == "daemon":
|
||||||
@@ -446,7 +448,7 @@ class IpcModelLoader(BaseModelLoader):
|
|||||||
|
|
||||||
return model
|
return model
|
||||||
|
|
||||||
def _fetch_from_cache(self, model_config) -> Optional[dict]:
|
def _fetch_from_cache(self, model_config, device_config) -> Optional[dict]:
|
||||||
"""Connect to daemon, validate config, fetch IPC handles.
|
"""Connect to daemon, validate config, fetch IPC handles.
|
||||||
|
|
||||||
Returns the daemon response dict on success, None if the daemon is
|
Returns the daemon response dict on success, None if the daemon is
|
||||||
@@ -455,6 +457,10 @@ class IpcModelLoader(BaseModelLoader):
|
|||||||
"""
|
"""
|
||||||
import socket as socket_mod
|
import socket as socket_mod
|
||||||
|
|
||||||
|
if self.socket_path is None:
|
||||||
|
device_uuid = current_platform.get_device_uuid(int(device_config.gpu_id))
|
||||||
|
self.socket_path = get_socket_path(device_uuid)
|
||||||
|
|
||||||
# Only connect to a real socket node owned by us: reject a symlink, a
|
# Only connect to a real socket node owned by us: reject a symlink, a
|
||||||
# plain file, or another user's socket planted at this /tmp path. An
|
# plain file, or another user's socket planted at this /tmp path. An
|
||||||
# absent socket means no daemon -> fall back to disk (return None).
|
# absent socket means no daemon -> fall back to disk (return None).
|
||||||
|
|||||||
@@ -269,12 +269,7 @@ def compute_env_stamp() -> Dict[str, str]:
|
|||||||
|
|
||||||
|
|
||||||
def compute_global_rank(tp_size: int, pp_rank: int, tp_rank: int) -> int:
|
def compute_global_rank(tp_size: int, pp_rank: int, tp_rank: int) -> int:
|
||||||
"""Single source of truth for the daemon rank formula.
|
"""Global rank for ``init_distributed_environment`` (tp_size * pp_rank + tp_rank)."""
|
||||||
|
|
||||||
global_rank = tp_size * pp_rank + tp_rank, so each daemon gets a unique
|
|
||||||
socket/ready path even across PP stages and nodes. Every call site (engine,
|
|
||||||
loader, model_runner, daemon) must go through this so the copies can't drift.
|
|
||||||
"""
|
|
||||||
return tp_size * pp_rank + tp_rank
|
return tp_size * pp_rank + tp_rank
|
||||||
|
|
||||||
|
|
||||||
@@ -302,38 +297,32 @@ def compute_local_gpu_id(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def _format_daemon_path(env_field, global_rank: int) -> str:
|
def _format_daemon_path(env_field, device_uuid: str) -> str:
|
||||||
"""Fill in a daemon path template, rejecting one that drops the rank.
|
"""Fill in a daemon path template, rejecting one that drops the GPU identity.
|
||||||
|
|
||||||
The template is user-overridable, and ``str.format`` silently ignores a
|
The template is user-overridable, and ``str.format`` silently ignores a
|
||||||
missing placeholder. Every rank would then derive the same path and map the
|
missing placeholder. Every physical GPU would then derive the same path,
|
||||||
shard belonging to whichever daemon got there first, so refuse up front
|
letting one job's client discover another job's daemon, so refuse up
|
||||||
rather than serve wrong weights.
|
front rather than serve wrong weights.
|
||||||
"""
|
"""
|
||||||
template = env_field.get()
|
template = env_field.get()
|
||||||
if "{global_rank}" not in template:
|
if "{device_uuid}" not in template:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"{env_field.name}={template!r} must contain '{{global_rank}}': each "
|
f"{env_field.name}={template!r} must contain '{{device_uuid}}': "
|
||||||
f"rank needs its own path, and a rank-independent one would point "
|
f"each physical GPU needs its own path, and a GPU-independent one "
|
||||||
f"every rank at a single daemon."
|
f"would point every caller at a single daemon."
|
||||||
)
|
)
|
||||||
return template.format(global_rank=global_rank)
|
return template.format(device_uuid=device_uuid)
|
||||||
|
|
||||||
|
|
||||||
def get_socket_path(global_rank: int) -> str:
|
def get_socket_path(device_uuid: str) -> str:
|
||||||
"""Get the Unix socket path for a weight cache daemon.
|
"""Get the Unix socket path for a weight cache daemon's physical GPU."""
|
||||||
|
return _format_daemon_path(envs.SGLANG_WEIGHT_CACHE_SOCKET_TEMPLATE, device_uuid)
|
||||||
global_rank = tp_size * pp_rank + tp_rank
|
|
||||||
"""
|
|
||||||
return _format_daemon_path(envs.SGLANG_WEIGHT_CACHE_SOCKET_TEMPLATE, global_rank)
|
|
||||||
|
|
||||||
|
|
||||||
def get_ready_path(global_rank: int) -> str:
|
def get_ready_path(device_uuid: str) -> str:
|
||||||
"""Get the ready-file path for a weight cache daemon.
|
"""Get the ready-file path for a weight cache daemon's physical GPU."""
|
||||||
|
return _format_daemon_path(envs.SGLANG_WEIGHT_CACHE_READY_TEMPLATE, device_uuid)
|
||||||
global_rank = tp_size * pp_rank + tp_rank
|
|
||||||
"""
|
|
||||||
return _format_daemon_path(envs.SGLANG_WEIGHT_CACHE_READY_TEMPLATE, global_rank)
|
|
||||||
|
|
||||||
|
|
||||||
def _read_ready_pid(ready_path: str) -> Optional[int]:
|
def _read_ready_pid(ready_path: str) -> Optional[int]:
|
||||||
@@ -359,8 +348,8 @@ def _is_pid_alive(pid: int) -> bool:
|
|||||||
return True
|
return True
|
||||||
|
|
||||||
|
|
||||||
def cleanup_stale_daemon_files(global_rank: int, *, force: bool = False) -> None:
|
def cleanup_stale_daemon_files(device_uuid: str, *, force: bool = False) -> None:
|
||||||
"""Validate and clean up .ready/.sock files for a daemon rank.
|
"""Validate and clean up .ready/.sock files for a daemon's physical GPU.
|
||||||
|
|
||||||
If the .ready file exists and the recorded PID is still alive, the daemon
|
If the .ready file exists and the recorded PID is still alive, the daemon
|
||||||
is still running — raise RuntimeError so the caller doesn't clobber it,
|
is still running — raise RuntimeError so the caller doesn't clobber it,
|
||||||
@@ -369,8 +358,8 @@ def cleanup_stale_daemon_files(global_rank: int, *, force: bool = False) -> None
|
|||||||
If the PID is dead (or unreadable), the files are stale leftovers from a
|
If the PID is dead (or unreadable), the files are stale leftovers from a
|
||||||
crashed/killed daemon and are safe to remove.
|
crashed/killed daemon and are safe to remove.
|
||||||
"""
|
"""
|
||||||
ready_path = get_ready_path(global_rank)
|
ready_path = get_ready_path(device_uuid)
|
||||||
socket_path = get_socket_path(global_rank)
|
socket_path = get_socket_path(device_uuid)
|
||||||
|
|
||||||
if not os.path.exists(ready_path) and not os.path.exists(socket_path):
|
if not os.path.exists(ready_path) and not os.path.exists(socket_path):
|
||||||
return
|
return
|
||||||
@@ -380,14 +369,14 @@ def cleanup_stale_daemon_files(global_rank: int, *, force: bool = False) -> None
|
|||||||
if pid is not None and _is_pid_alive(pid):
|
if pid is not None and _is_pid_alive(pid):
|
||||||
if not force:
|
if not force:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
f"Weight cache daemon for rank {global_rank} is already running "
|
f"Weight cache daemon for GPU {device_uuid} is already running "
|
||||||
f"(pid={pid}, ready={ready_path}). Stop the existing daemon before "
|
f"(pid={pid}, ready={ready_path}). Stop the existing daemon before "
|
||||||
f"launching a new one, or pass force=True (--force) to kill it and "
|
f"launching a new one, or pass force=True (--force) to kill it and "
|
||||||
f"take over."
|
f"take over."
|
||||||
)
|
)
|
||||||
logger.warning(
|
logger.warning(
|
||||||
f"[weight_cache] force takeover: killing existing daemon pid={pid} "
|
f"[weight_cache] force takeover: killing existing daemon pid={pid} "
|
||||||
f"for rank {global_rank} and reclaiming its socket/ready files."
|
f"for GPU {device_uuid} and reclaiming its socket/ready files."
|
||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
os.kill(pid, signal.SIGKILL)
|
os.kill(pid, signal.SIGKILL)
|
||||||
|
|||||||
@@ -7,7 +7,9 @@ import unittest
|
|||||||
import requests
|
import requests
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.platforms import current_platform
|
||||||
from sglang.srt.utils import kill_process_tree
|
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.ci.ci_register import register_cuda_ci
|
||||||
from sglang.test.test_utils import (
|
from sglang.test.test_utils import (
|
||||||
DEFAULT_TARGET_MODEL_EAGLE_DP_ATTN,
|
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(
|
@unittest.skipIf(
|
||||||
torch.cuda.device_count() < 2,
|
torch.cuda.device_count() < 2,
|
||||||
"TP=2 weight cache daemon test requires >=2 GPUs (skipped on the 1-gpu runner)",
|
"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.model = cls.model_override or DEFAULT_MODEL
|
||||||
cls.base_url = DEFAULT_URL_FOR_TEST
|
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||||
cls.tp_size = 2
|
cls.tp_size = 2
|
||||||
|
cls.gpu_uuids = _gpu_uuids(cls.tp_size)
|
||||||
|
|
||||||
# Clean up stale ready/socket files from previous runs
|
# Clean up stale ready/socket files from previous runs
|
||||||
for rank in range(cls.tp_size):
|
for device_uuid in cls.gpu_uuids:
|
||||||
for suffix in (".ready", ".sock"):
|
for path in (get_ready_path(device_uuid), get_socket_path(device_uuid)):
|
||||||
path = f"/tmp/sglang_weight_cache_rank{rank}{suffix}"
|
|
||||||
if os.path.exists(path):
|
if os.path.exists(path):
|
||||||
os.unlink(path)
|
os.unlink(path)
|
||||||
|
|
||||||
@@ -86,13 +93,14 @@ class TestWeightCacheDaemonTP2(CustomTestCase):
|
|||||||
# Step 2: Wait for all daemon ready files
|
# Step 2: Wait for all daemon ready files
|
||||||
timeout = DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH
|
timeout = DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH
|
||||||
start = time.time()
|
start = time.time()
|
||||||
for rank in range(cls.tp_size):
|
for device_uuid in cls.gpu_uuids:
|
||||||
ready_path = f"/tmp/sglang_weight_cache_rank{rank}.ready"
|
ready_path = get_ready_path(device_uuid)
|
||||||
while not os.path.exists(ready_path):
|
while not os.path.exists(ready_path):
|
||||||
if time.time() - start > timeout:
|
if time.time() - start > timeout:
|
||||||
kill_process_tree(cls.daemon_process.pid)
|
kill_process_tree(cls.daemon_process.pid)
|
||||||
raise TimeoutError(
|
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:
|
if cls.daemon_process.poll() is not None:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
@@ -136,9 +144,8 @@ class TestWeightCacheDaemonTP2(CustomTestCase):
|
|||||||
os.unlink(path)
|
os.unlink(path)
|
||||||
except OSError:
|
except OSError:
|
||||||
pass
|
pass
|
||||||
for rank in range(getattr(cls, "tp_size", 2)):
|
for device_uuid in getattr(cls, "gpu_uuids", ()):
|
||||||
for suffix in (".ready", ".sock"):
|
for path in (get_ready_path(device_uuid), get_socket_path(device_uuid)):
|
||||||
path = f"/tmp/sglang_weight_cache_rank{rank}{suffix}"
|
|
||||||
if os.path.exists(path):
|
if os.path.exists(path):
|
||||||
try:
|
try:
|
||||||
os.unlink(path)
|
os.unlink(path)
|
||||||
@@ -220,11 +227,11 @@ class TestWeightCacheDaemonTP1Smoke(CustomTestCase):
|
|||||||
cls.model = DEFAULT_MODEL
|
cls.model = DEFAULT_MODEL
|
||||||
cls.base_url = DEFAULT_URL_FOR_TEST
|
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||||
cls.tp_size = 1
|
cls.tp_size = 1
|
||||||
|
cls.gpu_uuids = _gpu_uuids(cls.tp_size)
|
||||||
|
|
||||||
# Clean up stale ready/socket files from previous runs.
|
# Clean up stale ready/socket files from previous runs.
|
||||||
for rank in range(cls.tp_size):
|
for device_uuid in cls.gpu_uuids:
|
||||||
for suffix in (".ready", ".sock"):
|
for path in (get_ready_path(device_uuid), get_socket_path(device_uuid)):
|
||||||
path = f"/tmp/sglang_weight_cache_rank{rank}{suffix}"
|
|
||||||
if os.path.exists(path):
|
if os.path.exists(path):
|
||||||
os.unlink(path)
|
os.unlink(path)
|
||||||
|
|
||||||
@@ -245,13 +252,14 @@ class TestWeightCacheDaemonTP1Smoke(CustomTestCase):
|
|||||||
# Step 2: Wait for the daemon ready file.
|
# Step 2: Wait for the daemon ready file.
|
||||||
timeout = DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH
|
timeout = DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH
|
||||||
start = time.time()
|
start = time.time()
|
||||||
for rank in range(cls.tp_size):
|
for device_uuid in cls.gpu_uuids:
|
||||||
ready_path = f"/tmp/sglang_weight_cache_rank{rank}.ready"
|
ready_path = get_ready_path(device_uuid)
|
||||||
while not os.path.exists(ready_path):
|
while not os.path.exists(ready_path):
|
||||||
if time.time() - start > timeout:
|
if time.time() - start > timeout:
|
||||||
kill_process_tree(cls.daemon_process.pid)
|
kill_process_tree(cls.daemon_process.pid)
|
||||||
raise TimeoutError(
|
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:
|
if cls.daemon_process.poll() is not None:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
@@ -294,9 +302,8 @@ class TestWeightCacheDaemonTP1Smoke(CustomTestCase):
|
|||||||
os.unlink(path)
|
os.unlink(path)
|
||||||
except OSError:
|
except OSError:
|
||||||
pass
|
pass
|
||||||
for rank in range(getattr(cls, "tp_size", 1)):
|
for device_uuid in getattr(cls, "gpu_uuids", ()):
|
||||||
for suffix in (".ready", ".sock"):
|
for path in (get_ready_path(device_uuid), get_socket_path(device_uuid)):
|
||||||
path = f"/tmp/sglang_weight_cache_rank{rank}{suffix}"
|
|
||||||
if os.path.exists(path):
|
if os.path.exists(path):
|
||||||
try:
|
try:
|
||||||
os.unlink(path)
|
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:
|
(test_weight_cache_daemon.py) cannot exercise cheaply:
|
||||||
|
|
||||||
- length-prefixed socket framing (send_msg/recv_msg) over socketpair()
|
- 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
|
- quant-config hashing and method-name extraction
|
||||||
- daemon spawn configuration and socket/ready path derivation
|
- daemon spawn configuration and socket/ready path derivation
|
||||||
- the IPC quantization allowlist (the gate that keeps silently-wrong
|
- the IPC quantization allowlist (the gate that keeps silently-wrong
|
||||||
@@ -25,6 +25,7 @@ from types import SimpleNamespace
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.weight_cache.protocol import (
|
from sglang.srt.weight_cache.protocol import (
|
||||||
IPC_QUANT_ALLOWLIST,
|
IPC_QUANT_ALLOWLIST,
|
||||||
CacheConfig,
|
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=1, tp_rank=0), 4)
|
||||||
self.assertEqual(compute_global_rank(tp_size=4, pp_rank=2, tp_rank=1), 9)
|
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):
|
def test_socket_and_ready_paths_are_unique_per_device_uuid(self):
|
||||||
self.assertNotEqual(get_socket_path(0), get_socket_path(1))
|
self.assertNotEqual(get_socket_path("gpu-aaa"), get_socket_path("gpu-bbb"))
|
||||||
self.assertTrue(get_socket_path(3).endswith("rank3.sock"))
|
self.assertTrue(get_socket_path("gpu-aaa").endswith("gpu-aaa.sock"))
|
||||||
self.assertTrue(get_ready_path(3).endswith("rank3.ready"))
|
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):
|
def test_compute_local_gpu_id_honors_base_and_step(self):
|
||||||
# Single-node TP=4: identity mapping rank -> gpu.
|
# Single-node TP=4: identity mapping rank -> gpu.
|
||||||
@@ -368,12 +384,12 @@ class TestIpcQuantAllowlist(CustomTestCase):
|
|||||||
|
|
||||||
|
|
||||||
class TestCleanupStaleDaemonFiles(CustomTestCase):
|
class TestCleanupStaleDaemonFiles(CustomTestCase):
|
||||||
# Use a rank far outside any realistic tp*pp layout so we never collide
|
# A key no real daemon would ever compute, so this never collides with
|
||||||
# with a daemon that might actually be running on the test host.
|
# one that might actually be running on the test host.
|
||||||
RANK = 987654
|
KEY = "test-cleanup-stale-daemon-files"
|
||||||
|
|
||||||
def _paths(self):
|
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):
|
def tearDown(self):
|
||||||
for path in self._paths():
|
for path in self._paths():
|
||||||
@@ -382,7 +398,7 @@ class TestCleanupStaleDaemonFiles(CustomTestCase):
|
|||||||
|
|
||||||
def test_no_files_is_noop(self):
|
def test_no_files_is_noop(self):
|
||||||
# Neither file present: must return quietly, not raise.
|
# 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):
|
def test_stale_files_without_live_pid_are_removed(self):
|
||||||
ready_path, socket_path = self._paths()
|
ready_path, socket_path = self._paths()
|
||||||
@@ -392,7 +408,7 @@ class TestCleanupStaleDaemonFiles(CustomTestCase):
|
|||||||
f.write("stale contents, no pid line\n")
|
f.write("stale contents, no pid line\n")
|
||||||
open(socket_path, "w").close()
|
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(ready_path))
|
||||||
self.assertFalse(os.path.exists(socket_path))
|
self.assertFalse(os.path.exists(socket_path))
|
||||||
@@ -405,7 +421,7 @@ class TestCleanupStaleDaemonFiles(CustomTestCase):
|
|||||||
open(socket_path, "w").close()
|
open(socket_path, "w").close()
|
||||||
|
|
||||||
with self.assertRaises(RuntimeError):
|
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(ready_path))
|
||||||
self.assertTrue(os.path.exists(socket_path))
|
self.assertTrue(os.path.exists(socket_path))
|
||||||
@@ -423,11 +439,11 @@ class TestCleanupStaleDaemonFiles(CustomTestCase):
|
|||||||
f.write(f"pid={child.pid}\n")
|
f.write(f"pid={child.pid}\n")
|
||||||
open(socket_path, "w").close()
|
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(ready_path))
|
||||||
self.assertFalse(os.path.exists(socket_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)
|
self.assertEqual(child.wait(timeout=5), -9)
|
||||||
finally:
|
finally:
|
||||||
if child.poll() is None:
|
if child.poll() is None:
|
||||||
@@ -442,7 +458,7 @@ class TestDaemonModeRefusesDiskLoad(CustomTestCase):
|
|||||||
pointing the loader at a socket path that does not exist.
|
pointing the loader at a socket path that does not exist.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
RANK = 987655
|
KEY = "test-daemon-mode-refuses-disk-load"
|
||||||
|
|
||||||
def _model_config(self):
|
def _model_config(self):
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
@@ -464,7 +480,7 @@ class TestDaemonModeRefusesDiskLoad(CustomTestCase):
|
|||||||
from sglang.srt.configs.load_config import LoadConfig, LoadFormat
|
from sglang.srt.configs.load_config import LoadConfig, LoadFormat
|
||||||
from sglang.srt.weight_cache.ipc_loader import IpcModelLoader
|
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):
|
if os.path.exists(missing_socket):
|
||||||
os.unlink(missing_socket)
|
os.unlink(missing_socket)
|
||||||
|
|
||||||
@@ -481,6 +497,30 @@ class TestDaemonModeRefusesDiskLoad(CustomTestCase):
|
|||||||
# fall through to a disk load.
|
# fall through to a disk load.
|
||||||
self.assertIn("daemon", str(ctx.exception).lower())
|
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__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user