[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>
|
||||
<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>
|
||||
</tr>
|
||||
<tr>
|
||||
@@ -275,7 +275,7 @@ SGLang supports various environment variables that can be used to configure its
|
||||
<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", 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>
|
||||
<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.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
|
||||
# (e.g. macOS/MPS has no triton, and torch.mps lacks Stream / set_device /
|
||||
# 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,
|
||||
)
|
||||
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
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -175,12 +176,17 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
|
||||
envs.SGLANG_DIFFUSION_CACHE_ROOT, "torch_compile_cache"
|
||||
)
|
||||
tmp_root = tempfile.gettempdir()
|
||||
sglang_defaults = third_party_cache_defaults()
|
||||
for env_name, sub in (
|
||||
("TORCHINDUCTOR_CACHE_DIR", "inductor"),
|
||||
("TRITON_CACHE_DIR", "triton"),
|
||||
):
|
||||
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.
|
||||
continue
|
||||
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:
|
||||
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:
|
||||
"""Build a descriptive error for a scheduler process that died during init."""
|
||||
|
||||
@@ -5,7 +5,7 @@ import subprocess
|
||||
import warnings
|
||||
from contextlib import ExitStack, contextmanager
|
||||
from enum import IntEnum
|
||||
from typing import Any, Optional
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=1)
|
||||
@@ -25,6 +25,15 @@ def _default_hip() -> bool:
|
||||
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:
|
||||
_allow_set_name = True
|
||||
|
||||
@@ -737,7 +746,8 @@ class Envs:
|
||||
SGLANG_JIT_DEEPGEMM_FAST_WARMUP = EnvBool(False)
|
||||
SGLANG_JIT_DEEPGEMM_COMPILE_WORKERS = EnvInt(4)
|
||||
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_USE_DEEPGEMM_BMM = 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}")
|
||||
|
||||
|
||||
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():
|
||||
_print_deprecated_env("SGLANG_GC_LOG", "SGLANG_LOG_GC")
|
||||
_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()
|
||||
_FAST_WARMUP = envs.SGLANG_JIT_DEEPGEMM_FAST_WARMUP.get()
|
||||
|
||||
# Force redirect deep_gemm cache_dir
|
||||
os.environ["DG_JIT_CACHE_DIR"] = os.getenv(
|
||||
"SGLANG_DG_CACHE_DIR", os.path.join(os.path.expanduser("~"), ".cache", "deep_gemm")
|
||||
)
|
||||
# Force redirect deep_gemm cache_dir. Defaults under SGLANG_CACHE_DIR so it
|
||||
# sits with the other compiled-kernel caches; SGLANG_DG_CACHE_DIR still wins.
|
||||
os.environ["DG_JIT_CACHE_DIR"] = envs.SGLANG_DG_CACHE_DIR.get()
|
||||
|
||||
# Refer to https://github.com/deepseek-ai/DeepGEMM/commit/d75b218b7b8f4a5dd5406ac87905039ead3ae42f
|
||||
# NVRTC may have performance loss with some cases.
|
||||
|
||||
@@ -155,8 +155,20 @@ install_apt_packages() {
|
||||
}
|
||||
|
||||
clean_site_packages() {
|
||||
# Clear torch compilation cache
|
||||
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)'
|
||||
# Clear torch compilation cache from every location it can be in; sglang
|
||||
# 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)
|
||||
SITE_PACKAGES=$(python3 -c "import site; print(site.getsitepackages()[0])")
|
||||
|
||||
@@ -26,9 +26,15 @@ from math import ceil
|
||||
from pathlib import Path
|
||||
from typing import Dict, List
|
||||
|
||||
# Shared with warmup_server.py. Wipe alongside /root/.cache/deep_gemm if you
|
||||
# clear the DeepGEMM JIT cache — a stale marker → in-test JIT compile.
|
||||
MARKER_DIR = os.path.join(os.path.expanduser("~"), ".cache", "sglang", "warmup_markers")
|
||||
from sglang.srt.environ import envs
|
||||
|
||||
# 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.
|
||||
FALLBACK_TIMEOUT_SEC = 600
|
||||
@@ -67,11 +73,10 @@ CRASH_MARKERS = (
|
||||
"Received sigquit from a child",
|
||||
)
|
||||
|
||||
# Configure DeepGEMM cache before importing deep_gemm
|
||||
os.environ["DG_JIT_CACHE_DIR"] = os.getenv(
|
||||
"SGLANG_DG_CACHE_DIR",
|
||||
os.path.join(os.path.expanduser("~"), ".cache", "deep_gemm"),
|
||||
)
|
||||
# Configure DeepGEMM cache before importing deep_gemm. Read through envs so
|
||||
# this warms the directory the server will actually compile into; duplicating
|
||||
# the default here is what let the two drift apart.
|
||||
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")
|
||||
|
||||
BLOCK_SIZE = 128
|
||||
@@ -510,9 +515,7 @@ def main():
|
||||
)
|
||||
print(f"=== DeepGEMM Lightweight Warmup ({len(model_tp_pairs)} model(s)) ===")
|
||||
print(f" Fast warmup: {fast_warmup}")
|
||||
print(
|
||||
f" Cache dir: {os.environ.get('DG_JIT_CACHE_DIR', '~/.cache/deep_gemm')}\n"
|
||||
)
|
||||
print(f" Cache dir: {os.environ['DG_JIT_CACHE_DIR']}\n")
|
||||
|
||||
# Load configs and deduplicate by architecture
|
||||
seen_keys = {}
|
||||
|
||||
@@ -26,9 +26,8 @@ from pathlib import Path
|
||||
|
||||
# Reuse helpers from warmup_deep_gemm (same directory)
|
||||
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
|
||||
SERVER_STARTUP_TIMEOUT = 900 # 15 min max to wait for server ready
|
||||
DEFAULT_PORT = 39876
|
||||
|
||||
Reference in New Issue
Block a user