[Weight Cache] Make daemon socket/ready paths configurable via env (#36299)
This commit is contained in:
@@ -1540,6 +1540,20 @@ class Envs:
|
|||||||
# Most batched requests one /generate HTTP call may expand into.
|
# Most batched requests one /generate HTTP call may expand into.
|
||||||
SGLANG_MAX_BATCH_REQS_PER_HTTP_REQ = EnvInt(4096)
|
SGLANG_MAX_BATCH_REQS_PER_HTTP_REQ = EnvInt(4096)
|
||||||
|
|
||||||
|
# ===================================================================
|
||||||
|
# Weight Cache Daemon
|
||||||
|
# ===================================================================
|
||||||
|
# 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
|
||||||
|
# talks to the daemon on its own GPU, so a rank-independent path would point
|
||||||
|
# every rank at one daemon and map another rank's shard.
|
||||||
|
SGLANG_WEIGHT_CACHE_SOCKET_TEMPLATE = EnvStr(
|
||||||
|
"/tmp/sglang_weight_cache_rank{global_rank}.sock"
|
||||||
|
)
|
||||||
|
SGLANG_WEIGHT_CACHE_READY_TEMPLATE = EnvStr(
|
||||||
|
"/tmp/sglang_weight_cache_rank{global_rank}.ready"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
envs = Envs()
|
envs = Envs()
|
||||||
EnvField._allow_set_name = False
|
EnvField._allow_set_name = False
|
||||||
|
|||||||
@@ -15,17 +15,11 @@ from typing import Any, Dict, Optional
|
|||||||
|
|
||||||
import msgspec
|
import msgspec
|
||||||
|
|
||||||
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.utils.common import safe_pickle_loads
|
from sglang.srt.utils.common import safe_pickle_loads
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
# Socket path template for weight cache daemons (keyed by global rank
|
|
||||||
# = tp_size * pp_rank + tp_rank, so multi-node / multi-PP don't collide)
|
|
||||||
WEIGHT_CACHE_SOCKET_TEMPLATE = "/tmp/sglang_weight_cache_rank{global_rank}.sock"
|
|
||||||
|
|
||||||
# Ready file template — daemon writes this after loading completes
|
|
||||||
WEIGHT_CACHE_READY_TEMPLATE = "/tmp/sglang_weight_cache_rank{global_rank}.ready"
|
|
||||||
|
|
||||||
|
|
||||||
class CacheConfig(msgspec.Struct):
|
class CacheConfig(msgspec.Struct):
|
||||||
"""Fingerprint of the cached weights. Used to validate compatibility
|
"""Fingerprint of the cached weights. Used to validate compatibility
|
||||||
@@ -308,12 +302,30 @@ def compute_local_gpu_id(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _format_daemon_path(env_field, global_rank: int) -> str:
|
||||||
|
"""Fill in a daemon path template, rejecting one that drops the rank.
|
||||||
|
|
||||||
|
The template is user-overridable, and ``str.format`` silently ignores a
|
||||||
|
missing placeholder. Every rank would then derive the same path and map the
|
||||||
|
shard belonging to whichever daemon got there first, so refuse up front
|
||||||
|
rather than serve wrong weights.
|
||||||
|
"""
|
||||||
|
template = env_field.get()
|
||||||
|
if "{global_rank}" not in template:
|
||||||
|
raise ValueError(
|
||||||
|
f"{env_field.name}={template!r} must contain '{{global_rank}}': each "
|
||||||
|
f"rank needs its own path, and a rank-independent one would point "
|
||||||
|
f"every rank at a single daemon."
|
||||||
|
)
|
||||||
|
return template.format(global_rank=global_rank)
|
||||||
|
|
||||||
|
|
||||||
def get_socket_path(global_rank: int) -> str:
|
def get_socket_path(global_rank: int) -> str:
|
||||||
"""Get the Unix socket path for a weight cache daemon.
|
"""Get the Unix socket path for a weight cache daemon.
|
||||||
|
|
||||||
global_rank = tp_size * pp_rank + tp_rank
|
global_rank = tp_size * pp_rank + tp_rank
|
||||||
"""
|
"""
|
||||||
return WEIGHT_CACHE_SOCKET_TEMPLATE.format(global_rank=global_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(global_rank: int) -> str:
|
||||||
@@ -321,7 +333,7 @@ def get_ready_path(global_rank: int) -> str:
|
|||||||
|
|
||||||
global_rank = tp_size * pp_rank + tp_rank
|
global_rank = tp_size * pp_rank + tp_rank
|
||||||
"""
|
"""
|
||||||
return WEIGHT_CACHE_READY_TEMPLATE.format(global_rank=global_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]:
|
||||||
|
|||||||
Reference in New Issue
Block a user