[core] Consolidate compiled-kernel caches under SGLANG_CACHE_DIR (#32434)
This commit is contained in:
@@ -80,7 +80,7 @@ SGLang supports various environment variables that can be used to configure its
|
|||||||
</tr>
|
</tr>
|
||||||
<tr>
|
<tr>
|
||||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_CACHE_DIR</code></td>
|
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_CACHE_DIR</code></td>
|
||||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Cache directory for model weights and other data</td>
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Cache directory for model weights and other data. Also the default root for compiled-kernel caches: Triton, Inductor, FlashInfer, the CUDA driver and DeepGEMM are pointed under it unless their own env vars (`TRITON_CACHE_DIR`, `TORCHINDUCTOR_CACHE_DIR`, `FLASHINFER_WORKSPACE_BASE`, `CUDA_CACHE_PATH`, `SGLANG_DG_CACHE_DIR`) are set explicitly</td>
|
||||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>~/.cache/sglang</code></td>
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>~/.cache/sglang</code></td>
|
||||||
</tr>
|
</tr>
|
||||||
<tr>
|
<tr>
|
||||||
@@ -275,7 +275,7 @@ SGLang supports various environment variables that can be used to configure its
|
|||||||
<tr>
|
<tr>
|
||||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>`SGLANG_DG_CACHE_DIR`</td>
|
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>`SGLANG_DG_CACHE_DIR`</td>
|
||||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Directory for caching compiled DeepGEMM kernels</td>
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Directory for caching compiled DeepGEMM kernels</td>
|
||||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>`~/.cache/deep_gemm`</td>
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>`{SGLANG_CACHE_DIR}/deep_gemm`</td>
|
||||||
</tr>
|
</tr>
|
||||||
<tr>
|
<tr>
|
||||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_DG_USE_NVRTC</code></td>
|
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_DG_USE_NVRTC</code></td>
|
||||||
|
|||||||
@@ -1,5 +1,13 @@
|
|||||||
# SGLang public APIs
|
# SGLang public APIs
|
||||||
|
|
||||||
|
# sglang.srt.environ must run before the rest of this file's imports
|
||||||
|
# (hf_transformers_patches, lang.api, ...), which pull in torch and
|
||||||
|
# FlashInfer: those claim these cache dirs early, and the first value set is
|
||||||
|
# the one that sticks. Safe here -- environ has no heavy dependency (no torch).
|
||||||
|
from sglang.srt.environ import redirect_third_party_caches
|
||||||
|
|
||||||
|
redirect_third_party_caches()
|
||||||
|
|
||||||
# Install stubs early for platforms where certain dependencies are unavailable
|
# Install stubs early for platforms where certain dependencies are unavailable
|
||||||
# (e.g. macOS/MPS has no triton, and torch.mps lacks Stream / set_device /
|
# (e.g. macOS/MPS has no triton, and torch.mps lacks Stream / set_device /
|
||||||
# get_device_properties). This must run before any downstream imports.
|
# get_device_properties). This must run before any downstream imports.
|
||||||
|
|||||||
@@ -82,6 +82,7 @@ from sglang.multimodal_gen.runtime.utils.trace_wrapper import (
|
|||||||
trace_slice,
|
trace_slice,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.utils import kill_itself_when_parent_died
|
from sglang.multimodal_gen.utils import kill_itself_when_parent_died
|
||||||
|
from sglang.srt.environ import third_party_cache_defaults
|
||||||
from sglang.srt.utils.network import NetworkAddress
|
from sglang.srt.utils.network import NetworkAddress
|
||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
@@ -175,12 +176,17 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
|
|||||||
envs.SGLANG_DIFFUSION_CACHE_ROOT, "torch_compile_cache"
|
envs.SGLANG_DIFFUSION_CACHE_ROOT, "torch_compile_cache"
|
||||||
)
|
)
|
||||||
tmp_root = tempfile.gettempdir()
|
tmp_root = tempfile.gettempdir()
|
||||||
|
sglang_defaults = third_party_cache_defaults()
|
||||||
for env_name, sub in (
|
for env_name, sub in (
|
||||||
("TORCHINDUCTOR_CACHE_DIR", "inductor"),
|
("TORCHINDUCTOR_CACHE_DIR", "inductor"),
|
||||||
("TRITON_CACHE_DIR", "triton"),
|
("TRITON_CACHE_DIR", "triton"),
|
||||||
):
|
):
|
||||||
current = os.environ.get(env_name)
|
current = os.environ.get(env_name)
|
||||||
if current and not current.startswith(tmp_root):
|
if (
|
||||||
|
current
|
||||||
|
and current != sglang_defaults.get(env_name)
|
||||||
|
and not current.startswith(tmp_root)
|
||||||
|
):
|
||||||
# Respect an explicit, non-ephemeral user-provided cache dir.
|
# Respect an explicit, non-ephemeral user-provided cache dir.
|
||||||
continue
|
continue
|
||||||
cache_path = os.path.join(compile_cache_root, sub)
|
cache_path = os.path.join(compile_cache_root, sub)
|
||||||
|
|||||||
@@ -1673,6 +1673,34 @@ def _set_envs_and_config(server_args: ServerArgs):
|
|||||||
if gc_threshold := server_args.gc_threshold:
|
if gc_threshold := server_args.gc_threshold:
|
||||||
gc.set_threshold(*gc_threshold)
|
gc.set_threshold(*gc_threshold)
|
||||||
|
|
||||||
|
_log_legacy_kernel_cache_dirs()
|
||||||
|
|
||||||
|
|
||||||
|
def _log_legacy_kernel_cache_dirs():
|
||||||
|
"""Note the pre-SGLANG_CACHE_DIR cache dirs without touching them: other
|
||||||
|
frameworks on the box may still be using them."""
|
||||||
|
# TODO(shuwang21): drop once SGLANG_CACHE_DIR has been the default for a
|
||||||
|
# few releases.
|
||||||
|
legacy_dirs = [
|
||||||
|
d
|
||||||
|
for d in (
|
||||||
|
os.path.expanduser("~/.triton"),
|
||||||
|
os.path.expanduser("~/.cache/flashinfer"),
|
||||||
|
os.path.expanduser("~/.cache/deep_gemm"),
|
||||||
|
)
|
||||||
|
if os.path.isdir(d)
|
||||||
|
]
|
||||||
|
if not legacy_dirs:
|
||||||
|
return
|
||||||
|
logger.info(
|
||||||
|
"Compiled-kernel caches now live under SGLANG_CACHE_DIR (%s). These "
|
||||||
|
"older directories are no longer used by sglang, but may still be "
|
||||||
|
"used by other frameworks on this machine, so they were left alone: "
|
||||||
|
"%s. Remove them yourself if nothing else needs them.",
|
||||||
|
envs.SGLANG_CACHE_DIR.get(),
|
||||||
|
", ".join(legacy_dirs),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _scheduler_died_error(rank: int, proc) -> RuntimeError:
|
def _scheduler_died_error(rank: int, proc) -> RuntimeError:
|
||||||
"""Build a descriptive error for a scheduler process that died during init."""
|
"""Build a descriptive error for a scheduler process that died during init."""
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ import subprocess
|
|||||||
import warnings
|
import warnings
|
||||||
from contextlib import ExitStack, contextmanager
|
from contextlib import ExitStack, contextmanager
|
||||||
from enum import IntEnum
|
from enum import IntEnum
|
||||||
from typing import Any, Optional
|
from typing import Any, Dict, Optional
|
||||||
|
|
||||||
|
|
||||||
@functools.lru_cache(maxsize=1)
|
@functools.lru_cache(maxsize=1)
|
||||||
@@ -25,6 +25,15 @@ def _default_hip() -> bool:
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def _default_cache_subdir(name: str) -> str:
|
||||||
|
"""A directory under SGLANG_CACHE_DIR, for env defaults that track it.
|
||||||
|
|
||||||
|
Pass as a callable default: SGLANG_CACHE_DIR is declared further down the
|
||||||
|
Envs body, and resolving late also lets tests override it.
|
||||||
|
"""
|
||||||
|
return os.path.join(os.path.expanduser(envs.SGLANG_CACHE_DIR.get()), name)
|
||||||
|
|
||||||
|
|
||||||
class EnvField:
|
class EnvField:
|
||||||
_allow_set_name = True
|
_allow_set_name = True
|
||||||
|
|
||||||
@@ -737,7 +746,8 @@ class Envs:
|
|||||||
SGLANG_JIT_DEEPGEMM_FAST_WARMUP = EnvBool(False)
|
SGLANG_JIT_DEEPGEMM_FAST_WARMUP = EnvBool(False)
|
||||||
SGLANG_JIT_DEEPGEMM_COMPILE_WORKERS = EnvInt(4)
|
SGLANG_JIT_DEEPGEMM_COMPILE_WORKERS = EnvInt(4)
|
||||||
SGLANG_IN_DEEPGEMM_PRECOMPILE_STAGE = EnvBool(False)
|
SGLANG_IN_DEEPGEMM_PRECOMPILE_STAGE = EnvBool(False)
|
||||||
SGLANG_DG_CACHE_DIR = EnvStr(os.path.expanduser("~/.cache/deep_gemm"))
|
# Resolved lazily so it tracks SGLANG_CACHE_DIR, which is defined below.
|
||||||
|
SGLANG_DG_CACHE_DIR = EnvStr(lambda: _default_cache_subdir("deep_gemm"))
|
||||||
SGLANG_DG_USE_NVRTC = EnvBool(False)
|
SGLANG_DG_USE_NVRTC = EnvBool(False)
|
||||||
SGLANG_USE_DEEPGEMM_BMM = EnvBool(False)
|
SGLANG_USE_DEEPGEMM_BMM = EnvBool(False)
|
||||||
SGLANG_DEEPGEMM_SANITY_CHECK = EnvBool(False)
|
SGLANG_DEEPGEMM_SANITY_CHECK = EnvBool(False)
|
||||||
@@ -1367,6 +1377,33 @@ def _warn_deprecated_env_to_cli_flag(env_name: str, suggestion: str):
|
|||||||
warnings.warn(f"Environment variable {env_name} is deprecated. {suggestion}")
|
warnings.warn(f"Environment variable {env_name} is deprecated. {suggestion}")
|
||||||
|
|
||||||
|
|
||||||
|
def third_party_cache_defaults() -> Dict[str, str]:
|
||||||
|
base = os.path.expanduser(envs.SGLANG_CACHE_DIR.get())
|
||||||
|
return {
|
||||||
|
"TRITON_CACHE_DIR": os.path.join(base, "triton"),
|
||||||
|
"TORCHINDUCTOR_CACHE_DIR": os.path.join(base, "inductor"),
|
||||||
|
"CUDA_CACHE_PATH": os.path.join(base, "nv"),
|
||||||
|
# FlashInfer appends ".cache/flashinfer" to this base itself, so this
|
||||||
|
# is the base dir rather than the final cache dir.
|
||||||
|
"FLASHINFER_WORKSPACE_BASE": base,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def redirect_third_party_caches():
|
||||||
|
"""Point third-party JIT caches at SGLANG_CACHE_DIR, so a run's compiled
|
||||||
|
kernels can be cleaned, warmed or volume-mounted as one directory.
|
||||||
|
|
||||||
|
Must be called early. The redirect silently does nothing if either of
|
||||||
|
these has already happened:
|
||||||
|
|
||||||
|
- FlashInfer was imported. It resolves its workspace at import time.
|
||||||
|
- Inductor made its first ``cache_dir()`` call. That call setdefaults
|
||||||
|
TORCHINDUCTOR_CACHE_DIR itself.
|
||||||
|
"""
|
||||||
|
for key, value in third_party_cache_defaults().items():
|
||||||
|
os.environ.setdefault(key, value)
|
||||||
|
|
||||||
|
|
||||||
def _convert_SGL_to_SGLANG():
|
def _convert_SGL_to_SGLANG():
|
||||||
_print_deprecated_env("SGLANG_GC_LOG", "SGLANG_LOG_GC")
|
_print_deprecated_env("SGLANG_GC_LOG", "SGLANG_LOG_GC")
|
||||||
_print_deprecated_env(
|
_print_deprecated_env(
|
||||||
|
|||||||
@@ -35,10 +35,9 @@ _IS_FIRST_RANK_ON_NODE = envs.SGLANG_IS_FIRST_RANK_ON_NODE.get()
|
|||||||
_IN_PRECOMPILE_STAGE = envs.SGLANG_IN_DEEPGEMM_PRECOMPILE_STAGE.get()
|
_IN_PRECOMPILE_STAGE = envs.SGLANG_IN_DEEPGEMM_PRECOMPILE_STAGE.get()
|
||||||
_FAST_WARMUP = envs.SGLANG_JIT_DEEPGEMM_FAST_WARMUP.get()
|
_FAST_WARMUP = envs.SGLANG_JIT_DEEPGEMM_FAST_WARMUP.get()
|
||||||
|
|
||||||
# Force redirect deep_gemm cache_dir
|
# Force redirect deep_gemm cache_dir. Defaults under SGLANG_CACHE_DIR so it
|
||||||
os.environ["DG_JIT_CACHE_DIR"] = os.getenv(
|
# sits with the other compiled-kernel caches; SGLANG_DG_CACHE_DIR still wins.
|
||||||
"SGLANG_DG_CACHE_DIR", os.path.join(os.path.expanduser("~"), ".cache", "deep_gemm")
|
os.environ["DG_JIT_CACHE_DIR"] = envs.SGLANG_DG_CACHE_DIR.get()
|
||||||
)
|
|
||||||
|
|
||||||
# Refer to https://github.com/deepseek-ai/DeepGEMM/commit/d75b218b7b8f4a5dd5406ac87905039ead3ae42f
|
# Refer to https://github.com/deepseek-ai/DeepGEMM/commit/d75b218b7b8f4a5dd5406ac87905039ead3ae42f
|
||||||
# NVRTC may have performance loss with some cases.
|
# NVRTC may have performance loss with some cases.
|
||||||
|
|||||||
@@ -155,8 +155,20 @@ install_apt_packages() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
clean_site_packages() {
|
clean_site_packages() {
|
||||||
# Clear torch compilation cache
|
# Clear torch compilation cache from every location it can be in; sglang
|
||||||
python3 -c 'import os, shutil, tempfile, getpass; cache_dir = os.environ.get("TORCHINDUCTOR_CACHE_DIR") or os.path.join(tempfile.gettempdir(), "torchinductor_" + getpass.getuser()); shutil.rmtree(cache_dir, ignore_errors=True)'
|
# is not installed yet, so it cannot be asked which one is in use.
|
||||||
|
python3 -c '
|
||||||
|
import getpass, os, shutil, tempfile
|
||||||
|
|
||||||
|
sglang_cache_dir = os.environ.get("SGLANG_CACHE_DIR") or "~/.cache/sglang"
|
||||||
|
for cache_dir in (
|
||||||
|
os.environ.get("TORCHINDUCTOR_CACHE_DIR"),
|
||||||
|
os.path.join(tempfile.gettempdir(), "torchinductor_" + getpass.getuser()),
|
||||||
|
os.path.join(os.path.expanduser(sglang_cache_dir), "inductor"),
|
||||||
|
):
|
||||||
|
if cache_dir:
|
||||||
|
shutil.rmtree(cache_dir, ignore_errors=True)
|
||||||
|
'
|
||||||
|
|
||||||
# Remove broken dist-info directories (missing METADATA per PEP 376)
|
# Remove broken dist-info directories (missing METADATA per PEP 376)
|
||||||
SITE_PACKAGES=$(python3 -c "import site; print(site.getsitepackages()[0])")
|
SITE_PACKAGES=$(python3 -c "import site; print(site.getsitepackages()[0])")
|
||||||
|
|||||||
@@ -26,9 +26,15 @@ from math import ceil
|
|||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Dict, List
|
from typing import Dict, List
|
||||||
|
|
||||||
# Shared with warmup_server.py. Wipe alongside /root/.cache/deep_gemm if you
|
from sglang.srt.environ import envs
|
||||||
# clear the DeepGEMM JIT cache — a stale marker → in-test JIT compile.
|
|
||||||
MARKER_DIR = os.path.join(os.path.expanduser("~"), ".cache", "sglang", "warmup_markers")
|
# Shared with warmup_server.py. Uses the same root as DG_JIT_CACHE_DIR below,
|
||||||
|
# so overriding SGLANG_CACHE_DIR moves the markers and the cache together.
|
||||||
|
# If only one moved, a marker could report a model as warmed while its cache
|
||||||
|
# is empty, and the test would pay for the JIT compilation it should skip.
|
||||||
|
MARKER_DIR = os.path.join(
|
||||||
|
os.path.expanduser(envs.SGLANG_CACHE_DIR.get()), "warmup_markers"
|
||||||
|
)
|
||||||
|
|
||||||
# Outer cap for stuck fallback subprocesses; CRASH_MARKERS abort sooner.
|
# Outer cap for stuck fallback subprocesses; CRASH_MARKERS abort sooner.
|
||||||
FALLBACK_TIMEOUT_SEC = 600
|
FALLBACK_TIMEOUT_SEC = 600
|
||||||
@@ -67,11 +73,10 @@ CRASH_MARKERS = (
|
|||||||
"Received sigquit from a child",
|
"Received sigquit from a child",
|
||||||
)
|
)
|
||||||
|
|
||||||
# Configure DeepGEMM cache before importing deep_gemm
|
# Configure DeepGEMM cache before importing deep_gemm. Read through envs so
|
||||||
os.environ["DG_JIT_CACHE_DIR"] = os.getenv(
|
# this warms the directory the server will actually compile into; duplicating
|
||||||
"SGLANG_DG_CACHE_DIR",
|
# the default here is what let the two drift apart.
|
||||||
os.path.join(os.path.expanduser("~"), ".cache", "deep_gemm"),
|
os.environ["DG_JIT_CACHE_DIR"] = envs.SGLANG_DG_CACHE_DIR.get()
|
||||||
)
|
|
||||||
os.environ["DG_JIT_USE_NVRTC"] = os.getenv("SGL_DG_USE_NVRTC", "0")
|
os.environ["DG_JIT_USE_NVRTC"] = os.getenv("SGL_DG_USE_NVRTC", "0")
|
||||||
|
|
||||||
BLOCK_SIZE = 128
|
BLOCK_SIZE = 128
|
||||||
@@ -510,9 +515,7 @@ def main():
|
|||||||
)
|
)
|
||||||
print(f"=== DeepGEMM Lightweight Warmup ({len(model_tp_pairs)} model(s)) ===")
|
print(f"=== DeepGEMM Lightweight Warmup ({len(model_tp_pairs)} model(s)) ===")
|
||||||
print(f" Fast warmup: {fast_warmup}")
|
print(f" Fast warmup: {fast_warmup}")
|
||||||
print(
|
print(f" Cache dir: {os.environ['DG_JIT_CACHE_DIR']}\n")
|
||||||
f" Cache dir: {os.environ.get('DG_JIT_CACHE_DIR', '~/.cache/deep_gemm')}\n"
|
|
||||||
)
|
|
||||||
|
|
||||||
# Load configs and deduplicate by architecture
|
# Load configs and deduplicate by architecture
|
||||||
seen_keys = {}
|
seen_keys = {}
|
||||||
|
|||||||
@@ -26,9 +26,8 @@ from pathlib import Path
|
|||||||
|
|
||||||
# Reuse helpers from warmup_deep_gemm (same directory)
|
# Reuse helpers from warmup_deep_gemm (same directory)
|
||||||
sys.path.insert(0, os.path.dirname(__file__))
|
sys.path.insert(0, os.path.dirname(__file__))
|
||||||
from warmup_deep_gemm import get_architecture_key, get_config_json
|
from warmup_deep_gemm import MARKER_DIR, get_architecture_key, get_config_json
|
||||||
|
|
||||||
MARKER_DIR = os.path.join(os.path.expanduser("~"), ".cache", "sglang", "warmup_markers")
|
|
||||||
HEALTH_POLL_INTERVAL = 10 # seconds between health checks
|
HEALTH_POLL_INTERVAL = 10 # seconds between health checks
|
||||||
SERVER_STARTUP_TIMEOUT = 900 # 15 min max to wait for server ready
|
SERVER_STARTUP_TIMEOUT = 900 # 15 min max to wait for server ready
|
||||||
DEFAULT_PORT = 39876
|
DEFAULT_PORT = 39876
|
||||||
|
|||||||
Reference in New Issue
Block a user