From 4229088a48612191756ad7aa0e6bb82275cc1b53 Mon Sep 17 00:00:00 2001 From: James Liu <51351043+chromecast56@users.noreply.github.com> Date: Wed, 2 Sep 2026 19:59:04 -0700 Subject: [PATCH] feat(kernels): generalize persistent CuTe JIT cache (#33911) --- .../docs/references/environment_variables.mdx | 5 + python/sglang/kernels/jit/cute_aot_cache.py | 429 ++++++++++++++++++ .../attention/flash_attn/cute/cache_utils.py | 289 ------------ .../attention/flash_attn/cute/interface.py | 28 +- .../kernels/ops/kimi_k3/kda_decode_mtp.py | 40 +- python/sglang/srt/environ.py | 4 + 6 files changed, 487 insertions(+), 308 deletions(-) create mode 100644 python/sglang/kernels/jit/cute_aot_cache.py delete mode 100644 python/sglang/kernels/ops/attention/flash_attn/cute/cache_utils.py diff --git a/docs/docs/references/environment_variables.mdx b/docs/docs/references/environment_variables.mdx index 82e0ea4df..b10ac265e 100644 --- a/docs/docs/references/environment_variables.mdx +++ b/docs/docs/references/environment_variables.mdx @@ -83,6 +83,11 @@ SGLang supports various environment variables that can be used to configure its 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 ~/.cache/sglang + + SGLANG_CUTE_AOT_CACHE_DIR + Trusted directory for persistent CuTe DSL AOT objects shared across process restarts. Artifacts are namespaced by source, runtime ABI, host platform, and target GPU architecture. Set to an empty string to keep compilation process-local. Only use a directory writable by trusted users because SGLang loads cached object files into the process. + {`{SGLANG_CACHE_DIR}`}/cute_aot + SGLANG_PREFETCH_BLOCK_SIZE_MB Block size (in MB) for sequential checkpoint prefetch reads that warm the OS page cache before workers load weights via mmap diff --git a/python/sglang/kernels/jit/cute_aot_cache.py b/python/sglang/kernels/jit/cute_aot_cache.py new file mode 100644 index 000000000..0824622c0 --- /dev/null +++ b/python/sglang/kernels/jit/cute_aot_cache.py @@ -0,0 +1,429 @@ +# Manage Ahead-of-Time (AOT) compiled kernels +"""In-memory and persistent caches for CuTe DSL JIT functions. + +Compiled objects persist under ``SGLANG_CUTE_AOT_CACHE_DIR`` (default +``{SGLANG_CACHE_DIR}/cute_aot``) and are shared across process restarts. Set +the variable to an empty string to keep compilation process-local. The +directory must be trusted: cached object files are loaded into the process. +""" + +import ctypes +import fcntl +import hashlib +import logging +import os +import pickle +import platform +import sys +import time +from collections.abc import Callable, Sequence +from functools import lru_cache +from pathlib import Path +from typing import Any, Hashable, TypeAlias + +logger = logging.getLogger(__name__) + +_runtime_library_handles: list[Any] = [] +_loaded_modules: list[Any] = [] + +CompileKeyType: TypeAlias = tuple[Hashable, ...] +CallableFunction: TypeAlias = Any + +_UNSET_CACHE_DIR = object() + + +def _normalize_disk_key(value: Any) -> Any: + value_type = type(value) + if value_type.__module__ == "torch" and value_type.__name__ == "device": + return ("torch.device", value.type) + if isinstance(value, tuple): + return tuple(_normalize_disk_key(item) for item in value) + return value + + +@lru_cache(maxsize=None) +def _compute_source_fingerprint( + source_paths: tuple[str, ...], enable_tvm_ffi: bool, target_arch: str +) -> str: + """ + Hash all CuTe Python sources plus runtime ABI stamps into a short fingerprint. + + The fingerprint changes with the supplied sources, Python/CuTe versions, + selected ABI, CUDA version, or target architecture. + + Computed once per process and cached. + """ + import cutlass + + h = hashlib.sha256() + + h.update(f"py{sys.version_info.major}.{sys.version_info.minor}".encode()) + # Exported objects contain host machine code, not just GPU code. + h.update(f"host={sys.platform}-{platform.machine()}".encode()) + h.update(f"cutlass={cutlass.__version__}".encode()) + h.update(f"cuda={getattr(cutlass, 'CUDA_VERSION', 'unknown')}".encode()) + h.update(f"tvm_ffi={enable_tvm_ffi}".encode()) + h.update(f"arch={target_arch}".encode()) + if enable_tvm_ffi: + import tvm_ffi + + h.update(f"tvm_ffi_version={tvm_ffi.__version__}".encode()) + + for index, raw_path in enumerate(source_paths): + source_path = Path(raw_path).resolve() + if source_path.is_dir(): + sources = sorted(source_path.rglob("*.py")) + root = source_path + elif source_path.is_file(): + sources = [source_path] + root = source_path.parent + else: + raise FileNotFoundError(source_path) + for src in sources: + if not src.is_file(): + continue + h.update(f"{index}:{src.relative_to(root).as_posix()}".encode()) + content = src.read_bytes() + h.update(len(content).to_bytes(8, "little")) + h.update(content) + + return h.hexdigest() + + +def _resolve_target_arch() -> str: + if target_arch := os.getenv("CUTE_DSL_ARCH"): + return target_arch + + import torch + + major, minor = torch.cuda.get_device_capability() + return f"sm_{major}{minor}" + + +# Pre-load cute DSL runtime libraries with RTLD_GLOBAL so that their symbols +# (e.g. _cudaLibraryLoadData) are visible to .so modules loaded later via dlopen. +# Upstream cute.runtime.load_module loads these without RTLD_GLOBAL, which causes +# "undefined symbol" errors when loading cached kernels from disk. +@lru_cache(maxsize=2) +def _preload_runtime_libraries(enable_tvm_ffi: bool) -> None: + import cutlass.cute as cute + + for raw_path in cute.runtime.find_runtime_libraries(enable_tvm_ffi=enable_tvm_ffi): + path = Path(raw_path) + if path.is_file(): + _runtime_library_handles.append( + ctypes.CDLL(str(path), mode=ctypes.RTLD_GLOBAL) + ) + + +def _load_object( + object_path: Path, function_prefix: str, enable_tvm_ffi: bool +) -> CallableFunction: + import cutlass.cute as cute + + _preload_runtime_libraries(enable_tvm_ffi) + module = cute.runtime.load_module(str(object_path), enable_tvm_ffi=enable_tvm_ffi) + try: + function = module[function_prefix] + except (KeyError, TypeError): + function = getattr(module, function_prefix) + _loaded_modules.append(module) + return function + + +class FileLock: + """Context manager for advisory file locks using fcntl.flock. + + Supports exclusive (write) and shared (read) locks. + Always blocks with polling until the lock is acquired or timeout is reached. + + Usage: + with FileLock(lock_path, exclusive=True, timeout=15, label="abc"): + # do work under lock + """ + + def __init__( + self, + lock_path: Path, + exclusive: bool, + timeout: float = 15, + label: str = "", + ): + """ + Args: + lock_path: Path to the lock file on disk. + exclusive: True for exclusive (write) lock, False for shared (read) lock. + timeout: Max seconds to wait for lock acquisition before raising RuntimeError. + label: Optional human-readable label for error messages. + """ + self.lock_path: Path = lock_path + self.exclusive: bool = exclusive + self.timeout: float = timeout + self.label: str = label + self._fd: int = -1 + + @property + def _lock_label(self) -> str: + kind = "exclusive" if self.exclusive else "shared" + return f"{kind} {self.label}" if self.label else kind + + def __enter__(self) -> "FileLock": + open_flags = ( + os.O_WRONLY | os.O_CREAT if self.exclusive else os.O_RDONLY | os.O_CREAT + ) + lock_type = fcntl.LOCK_EX if self.exclusive else fcntl.LOCK_SH + + self._fd = os.open(str(self.lock_path), open_flags) + + deadline = time.monotonic() + self.timeout + acquired = False + while time.monotonic() < deadline: + try: + fcntl.flock(self._fd, lock_type | fcntl.LOCK_NB) + acquired = True + break + except OSError: + time.sleep(0.1) + if not acquired: + os.close(self._fd) + self._fd = None + raise RuntimeError( + f"Timed out after {self.timeout}s waiting for " + f"{self._lock_label} lock: {self.lock_path}" + ) + + return self + + def __exit__(self, exc_type, exc_val, exc_tb) -> None: + if self._fd is not None: + fcntl.flock(self._fd, fcntl.LOCK_UN) + os.close(self._fd) + self._fd = None + + +class JITCache: + """ + In-memory cache for compiled functions. + """ + + def __init__(self): + self.cache: dict[CompileKeyType, CallableFunction] = {} + + def __setitem__(self, key: CompileKeyType, fn: CallableFunction) -> None: + self.cache[key] = fn + + def __getitem__(self, key: CompileKeyType) -> CallableFunction: + return self.cache[key] + + def __contains__(self, key: CompileKeyType) -> bool: + return key in self.cache + + def clear(self) -> None: + """ + Clear in-memory cache of compiled functions + """ + self.cache.clear() + + +class JITPersistentCache(JITCache): + """ + In-memory cache for compiled functions, which is also backed by persistent storage. + + ``cache_path`` may be a path or a zero-argument callable returning one; a + callable is resolved on first storage access, so constructing a cache at + import time never probes the GPU or hashes sources. + """ + + EXPORT_FUNCTION_PREFIX = "func" + LOCK_TIMEOUT_SECONDS = 15 + + def __init__( + self, + cache_path: Path | Callable[[], Path], + *, + enable_tvm_ffi: bool = True, + ): + super().__init__() + self._cache_path_source = cache_path + self._resolved_cache_path: Path | None = None + self.enable_tvm_ffi = enable_tvm_ffi + + @property + def cache_path(self) -> Path: + if self._resolved_cache_path is None: + source = self._cache_path_source + path = Path(source() if callable(source) else source) + path.mkdir(parents=True, exist_ok=True) + self._resolved_cache_path = path + return self._resolved_cache_path + + def __setitem__(self, key: CompileKeyType, fn: CallableFunction) -> None: + JITCache.__setitem__(self, key, fn) + self._try_export_to_storage(key, fn) + + def __getitem__(self, key: CompileKeyType) -> CallableFunction: + # Use __contains__ to try populating in-memory cache with persistent storage + self.__contains__(key) + return JITCache.__getitem__(self, key) + + def __contains__(self, key: CompileKeyType) -> bool: + # Checks in-memory cache first, then tries loading from storage. + # When returning True, guarantees the in-memory cache is populated. + if JITCache.__contains__(self, key): + return True + return self._try_load_from_storage(key) + + def _try_load_from_storage(self, key: CompileKeyType) -> bool: + """ + Try to load a function from persistent storage into in-memory cache. + Returns True if loaded successfully, False if not found on disk. + Holds a shared lock during loading to prevent concurrent writes. + """ + sha256_hex = self._key_to_hash(key) + obj_path = self.cache_path / f"{sha256_hex}.o" + invalid_inode = None + with FileLock( + self._lock_path(sha256_hex), + exclusive=False, + timeout=self.LOCK_TIMEOUT_SECONDS, + label=sha256_hex, + ): + if obj_path.exists(): + logger.debug("Loading compiled function from disk: %s", obj_path) + try: + fn = _load_object( + obj_path, self.EXPORT_FUNCTION_PREFIX, self.enable_tvm_ffi + ) + except Exception as error: + logger.warning("Invalid cache object %s: %s", obj_path, error) + try: + invalid_inode = obj_path.stat().st_ino + except OSError: + pass + else: + JITCache.__setitem__(self, key, fn) + return True + else: + logger.debug("Cache miss on disk for key hash %s", sha256_hex) + if invalid_inode is not None: + self._discard_invalid_object(sha256_hex, obj_path, invalid_inode) + return False + + def _discard_invalid_object( + self, sha256_hex: str, obj_path: Path, invalid_inode: int + ) -> None: + """Unlink a failed object under an exclusive lock. + + The shared load lock is released first (flock cannot upgrade), so a + writer may republish in the gap; the inode check keeps a fresh object + intact. Eviction is best-effort: on lock timeout the object is left + for the next process. + """ + try: + with FileLock( + self._lock_path(sha256_hex), + exclusive=True, + timeout=self.LOCK_TIMEOUT_SECONDS, + label=sha256_hex, + ): + try: + if obj_path.stat().st_ino == invalid_inode: + obj_path.unlink() + except OSError: + return + except RuntimeError as error: + logger.warning("Could not evict invalid object %s: %s", obj_path, error) + + def _try_export_to_storage(self, key: CompileKeyType, fn: CallableFunction) -> None: + """Export a compiled function to persistent storage under exclusive lock.""" + sha256_hex = self._key_to_hash(key) + with FileLock( + self._lock_path(sha256_hex), + exclusive=True, + timeout=self.LOCK_TIMEOUT_SECONDS, + label=sha256_hex, + ): + obj_path = self.cache_path / f"{sha256_hex}.o" + if obj_path.exists(): + # Another process already exported. + logger.debug("Skipping export, already on disk: %s", obj_path) + return + logger.debug("Exporting compiled function to disk: %s", obj_path) + temp_key = f".{sha256_hex}.tmp" + temp_obj_path = self.cache_path / f"{temp_key}.o" + temp_obj_path.unlink(missing_ok=True) + try: + if self.enable_tvm_ffi: + fn.export_to_c( + object_file_path=str(temp_obj_path), + function_name=self.EXPORT_FUNCTION_PREFIX, + ) + else: + fn.export_to_c( + str(self.cache_path), + temp_key, + function_prefix=self.EXPORT_FUNCTION_PREFIX, + ) + os.replace(temp_obj_path, obj_path) + finally: + temp_obj_path.unlink(missing_ok=True) + logger.debug( + "Successfully exported compiled function to disk: %s", obj_path + ) + + def _key_to_hash(self, key: CompileKeyType) -> str: + disk_key = (self.enable_tvm_ffi, _normalize_disk_key(key)) + return hashlib.sha256(pickle.dumps(disk_key)).hexdigest() + + def _lock_path(self, sha256_hex: str) -> Path: + return self.cache_path / f"{sha256_hex}.lock" + + def clear(self) -> None: + """ + Not only clear the in-memory cache. Also purge persistent compilation cache. + """ + logger.debug("Clearing persistent cache at %s", self.cache_path) + super().clear() + for child in self.cache_path.iterdir(): + child.unlink() + + +def get_jit_cache( + name: str | None = None, + *, + cache_dir: Any = _UNSET_CACHE_DIR, + source_paths: Sequence[str | os.PathLike[str]] = (), + enable_tvm_ffi: bool = True, +) -> JITCache: + """ + JIT cache factory. + `name` is an optional identifier to create subdirectories to manage cache. + + ``cache_dir`` defaults to ``SGLANG_CUTE_AOT_CACHE_DIR``; pass ``None`` (or + set the variable to an empty string) for a process-local cache. + + When persistent caching is enabled, artifacts are namespaced under a + source fingerprint directory so that code or dependency changes + automatically invalidate stale entries. + """ + if cache_dir is _UNSET_CACHE_DIR: + from sglang.srt.environ import envs + + cache_dir = envs.SGLANG_CUTE_AOT_CACHE_DIR.get() or None + if cache_dir is None: + logger.debug("Persistent cache disabled, using in-memory JIT cache") + return JITCache() + + def resolve_cache_path() -> Path: + paths = (str(Path(__file__).resolve()),) + tuple( + str(Path(path).resolve()) for path in source_paths + ) + path = Path(cache_dir).expanduser() / _compute_source_fingerprint( + paths, enable_tvm_ffi, _resolve_target_arch() + ) + if name: + path = path / name + logger.debug("Creating persistent JIT cache at %s", path) + return path + + return JITPersistentCache(resolve_cache_path, enable_tvm_ffi=enable_tvm_ffi) diff --git a/python/sglang/kernels/ops/attention/flash_attn/cute/cache_utils.py b/python/sglang/kernels/ops/attention/flash_attn/cute/cache_utils.py deleted file mode 100644 index 8c46e3cdc..000000000 --- a/python/sglang/kernels/ops/attention/flash_attn/cute/cache_utils.py +++ /dev/null @@ -1,289 +0,0 @@ -# Manage Ahead-of-Time (AOT) compiled kernels -import ctypes -import fcntl -import hashlib -import os -import pickle -import sys -import tempfile -import time -from functools import lru_cache -from getpass import getuser -from pathlib import Path -from typing import Hashable, TypeAlias - -import cutlass -import cutlass.cute as cute -import tvm_ffi -from cutlass.cutlass_dsl import JitCompiledFunction - -from sglang.kernels.ops.attention.flash_attn.cute.fa_logging import fa_log - -# Pre-load cute DSL runtime libraries with RTLD_GLOBAL so that their symbols -# (e.g. _cudaLibraryLoadData) are visible to .so modules loaded later via dlopen. -# Upstream cute.runtime.load_module loads these without RTLD_GLOBAL, which causes -# "undefined symbol" errors when loading cached kernels from disk. -for _lib_path in cute.runtime.find_runtime_libraries(enable_tvm_ffi=False): - if Path(_lib_path).exists(): - ctypes.CDLL(_lib_path, mode=ctypes.RTLD_GLOBAL) - -CompileKeyType: TypeAlias = tuple[Hashable, ...] -CallableFunction: TypeAlias = JitCompiledFunction | tvm_ffi.Function - -# Enable cache via `FLASH_ATTENTION_CUTE_DSL_CACHE_ENABLED=1` -CUTE_DSL_CACHE_ENABLED: bool = ( - os.getenv("FLASH_ATTENTION_CUTE_DSL_CACHE_ENABLED", "0") == "1" -) - - -# Customize cache dir via `FLASH_ATTENTION_CUTE_DSL_CACHE_DIR`, default is -# `/tmp/${USER}/flash_attention_cute_dsl_cache`` -CUTE_DSL_CACHE_DIR: str | None = os.getenv("FLASH_ATTENTION_CUTE_DSL_CACHE_DIR", None) - - -def get_cache_path() -> Path: - if CUTE_DSL_CACHE_DIR is not None: - cache_dir = Path(CUTE_DSL_CACHE_DIR) - else: - cache_dir = ( - Path(tempfile.gettempdir()) / getuser() / "flash_attention_cute_dsl_cache" - ) - cache_dir.mkdir(parents=True, exist_ok=True) - return cache_dir - - -@lru_cache(maxsize=1) -def _compute_source_fingerprint() -> str: - """ - Hash all CuTe Python sources plus runtime ABI stamps into a short fingerprint. - - The fingerprint changes whenever: - - Any .py file under flash_attn/cute is added, removed, renamed, or modified. - - The Python minor version changes (e.g. 3.13 -> 3.14). - - The cutlass or tvm_ffi package version changes. - - Computed once per process and cached. - """ - cute_root = Path(__file__).resolve().parent - h = hashlib.sha256() - - h.update(f"py{sys.version_info.major}.{sys.version_info.minor}".encode()) - h.update(f"cutlass={cutlass.__version__}".encode()) - h.update(f"tvm_ffi={tvm_ffi.__version__}".encode()) - - for src in sorted(cute_root.rglob("*.py")): - if not src.is_file(): - continue - h.update(src.relative_to(cute_root).as_posix().encode()) - content = src.read_bytes() - h.update(len(content).to_bytes(8, "little")) - h.update(content) - - return h.hexdigest() - - -class FileLock: - """Context manager for advisory file locks using fcntl.flock. - - Supports exclusive (write) and shared (read) locks. - Always blocks with polling until the lock is acquired or timeout is reached. - - Usage: - with FileLock(lock_path, exclusive=True, timeout=15, label="abc"): - # do work under lock - """ - - def __init__( - self, - lock_path: Path, - exclusive: bool, - timeout: float = 15, - label: str = "", - ): - """ - Args: - lock_path: Path to the lock file on disk. - exclusive: True for exclusive (write) lock, False for shared (read) lock. - timeout: Max seconds to wait for lock acquisition before raising RuntimeError. - label: Optional human-readable label for error messages. - """ - self.lock_path: Path = lock_path - self.exclusive: bool = exclusive - self.timeout: float = timeout - self.label: str = label - self._fd: int = -1 - - @property - def _lock_label(self) -> str: - kind = "exclusive" if self.exclusive else "shared" - return f"{kind} {self.label}" if self.label else kind - - def __enter__(self) -> "FileLock": - open_flags = ( - os.O_WRONLY | os.O_CREAT if self.exclusive else os.O_RDONLY | os.O_CREAT - ) - lock_type = fcntl.LOCK_EX if self.exclusive else fcntl.LOCK_SH - - self._fd = os.open(str(self.lock_path), open_flags) - - deadline = time.monotonic() + self.timeout - acquired = False - while time.monotonic() < deadline: - try: - fcntl.flock(self._fd, lock_type | fcntl.LOCK_NB) - acquired = True - break - except OSError: - time.sleep(0.1) - if not acquired: - os.close(self._fd) - self._fd = None - raise RuntimeError( - f"Timed out after {self.timeout}s waiting for " - f"{self._lock_label} lock: {self.lock_path}" - ) - - return self - - def __exit__(self, exc_type, exc_val, exc_tb) -> None: - if self._fd is not None: - fcntl.flock(self._fd, fcntl.LOCK_UN) - os.close(self._fd) - self._fd = None - - -class JITCache: - """ - In-memory cache for compiled functions. - """ - - def __init__(self): - self.cache: dict[CompileKeyType, CallableFunction] = {} - - def __setitem__(self, key: CompileKeyType, fn: JitCompiledFunction) -> None: - self.cache[key] = fn - - def __getitem__(self, key: CompileKeyType) -> CallableFunction: - return self.cache[key] - - def __contains__(self, key: CompileKeyType) -> bool: - return key in self.cache - - def clear(self) -> None: - """ - Clear in-memory cache of compiled functions - """ - self.cache.clear() - - -class JITPersistentCache(JITCache): - """ - In-memory cache for compiled functions, which is also backed by persistent storage. - Use cutedsl ahead-of-time (AOT) compilation, only supporting enable_tvm_ffi=True - """ - - EXPORT_FUNCTION_PREFIX = "func" - LOCK_TIMEOUT_SECONDS = 15 - - def __init__(self, cache_path: Path): - super().__init__() - cache_path.mkdir(parents=True, exist_ok=True) - self.cache_path: Path = cache_path - - def __setitem__(self, key: CompileKeyType, fn: JitCompiledFunction) -> None: - JITCache.__setitem__(self, key, fn) - self._try_export_to_storage(key, fn) - - def __getitem__(self, key: CompileKeyType) -> CallableFunction: - # Use __contains__ to try populating in-memory cache with persistent storage - self.__contains__(key) - return JITCache.__getitem__(self, key) - - def __contains__(self, key: CompileKeyType) -> bool: - # Checks in-memory cache first, then tries loading from storage. - # When returning True, guarantees the in-memory cache is populated. - if JITCache.__contains__(self, key): - return True - return self._try_load_from_storage(key) - - def _try_load_from_storage(self, key: CompileKeyType) -> bool: - """ - Try to load a function from persistent storage into in-memory cache. - Returns True if loaded successfully, False if not found on disk. - Holds a shared lock during loading to prevent concurrent writes. - """ - sha256_hex = self._key_to_hash(key) - obj_path = self.cache_path / f"{sha256_hex}.o" - with FileLock( - self._lock_path(sha256_hex), - exclusive=False, - timeout=self.LOCK_TIMEOUT_SECONDS, - label=sha256_hex, - ): - if obj_path.exists(): - fa_log(1, f"Loading compiled function from disk: {obj_path}") - m = cute.runtime.load_module(str(obj_path), enable_tvm_ffi=True) - fn = getattr(m, self.EXPORT_FUNCTION_PREFIX) - JITCache.__setitem__(self, key, fn) - return True - else: - fa_log(1, f"Cache miss on disk for key hash {sha256_hex}") - return False - - def _try_export_to_storage( - self, key: CompileKeyType, fn: JitCompiledFunction - ) -> None: - """Export a compiled function to persistent storage under exclusive lock.""" - sha256_hex = self._key_to_hash(key) - with FileLock( - self._lock_path(sha256_hex), - exclusive=True, - timeout=self.LOCK_TIMEOUT_SECONDS, - label=sha256_hex, - ): - obj_path = self.cache_path / f"{sha256_hex}.o" - if obj_path.exists(): - # Another process already exported. - fa_log(1, f"Skipping export, already on disk: {obj_path}") - return - fa_log(1, f"Exporting compiled function to disk: {obj_path}") - fn.export_to_c( - object_file_path=str(obj_path), - function_name=self.EXPORT_FUNCTION_PREFIX, - ) - fa_log(1, f"Successfully exported compiled function to disk: {obj_path}") - - def _key_to_hash(self, key: CompileKeyType) -> str: - return hashlib.sha256(pickle.dumps(key)).hexdigest() - - def _lock_path(self, sha256_hex: str) -> Path: - return self.cache_path / f"{sha256_hex}.lock" - - def clear(self) -> None: - """ - Not only clear the in-memory cache. Also purge persistent compilation cache. - """ - fa_log(1, f"Clearing persistent cache at {self.cache_path}") - super().clear() - for child in self.cache_path.iterdir(): - child.unlink() - - -def get_jit_cache(name: str | None = None) -> JITCache: - """ - JIT cache factory. - `name` is an optional identifier to create subdirectories to manage cache. - - When persistent caching is enabled, artifacts are namespaced under a - source fingerprint directory so that code or dependency changes - automatically invalidate stale entries. - """ - if CUTE_DSL_CACHE_ENABLED: - path = get_cache_path() / _compute_source_fingerprint() - if name: - path = path / name - fa_log(1, f"Creating persistent JIT cache at {path}") - return JITPersistentCache(path) - else: - fa_log(1, "Persistent cache disabled, using in-memory JIT cache") - return JITCache() diff --git a/python/sglang/kernels/ops/attention/flash_attn/cute/interface.py b/python/sglang/kernels/ops/attention/flash_attn/cute/interface.py index 437bca8ca..746b21c17 100644 --- a/python/sglang/kernels/ops/attention/flash_attn/cute/interface.py +++ b/python/sglang/kernels/ops/attention/flash_attn/cute/interface.py @@ -11,11 +11,11 @@ import torch from cutlass import Float32, Int32 from quack.compile_utils import make_fake_tensor as fake_tensor +from sglang.kernels.jit.cute_aot_cache import get_jit_cache from sglang.kernels.jit.utils import is_arch_support_pdl from sglang.kernels.ops.attention.flash_attn.cute.batch_invariance import ( is_batch_invariant, ) -from sglang.kernels.ops.attention.flash_attn.cute.cache_utils import get_jit_cache from sglang.kernels.ops.attention.flash_attn.cute.testing import is_fake_mode if os.environ.get("CUTE_DSL_PTXAS_PATH", None) is not None: @@ -27,6 +27,11 @@ if os.environ.get("CUTE_DSL_PTXAS_PATH", None) is not None: cute_dsl_ptxas.patch() +from sglang.kernels.ops.attention.fa4_sm120.dispatch import ( + get_forward_host, + try_cached_paged_decode, + try_cached_varlen, +) from sglang.kernels.ops.attention.flash_attn.cute import fa_logging, utils from sglang.kernels.ops.attention.flash_attn.cute.block_sparsity import ( BlockSparseTensorsTorch, @@ -58,11 +63,6 @@ from sglang.kernels.ops.attention.flash_attn.cute.flash_fwd_sm100 import ( DescaleTensors, FlashAttentionForwardSm100, ) -from sglang.kernels.ops.attention.fa4_sm120.dispatch import ( - get_forward_host, - try_cached_paged_decode, - try_cached_varlen, -) from sglang.kernels.ops.attention.flash_attn.cute.shearing_bias import ShearingBias # SM100 head_dim=256 2CTA kernel imports @@ -1851,9 +1851,17 @@ def _flash_attn_fwd( return out, lse -_flash_attn_fwd.compile_cache = get_jit_cache("fwd") -_flash_attn_fwd.compile_cache_shear_bias = get_jit_cache("fwd_shear_bias") -_flash_attn_fwd.compile_cache_prepare_shear_bias = get_jit_cache( +def _get_jit_cache(name: str): + return get_jit_cache( + name, + source_paths=(os.path.dirname(os.path.abspath(__file__)),), + enable_tvm_ffi=True, + ) + + +_flash_attn_fwd.compile_cache = _get_jit_cache("fwd") +_flash_attn_fwd.compile_cache_shear_bias = _get_jit_cache("fwd_shear_bias") +_flash_attn_fwd.compile_cache_prepare_shear_bias = _get_jit_cache( "fwd_prepare_shear_bias" ) @@ -2521,7 +2529,7 @@ def _flash_attn_fwd_combine( ) -_flash_attn_fwd_combine.compile_cache = get_jit_cache("fwd_combine") +_flash_attn_fwd_combine.compile_cache = _get_jit_cache("fwd_combine") def flash_attn_combine( diff --git a/python/sglang/kernels/ops/kimi_k3/kda_decode_mtp.py b/python/sglang/kernels/ops/kimi_k3/kda_decode_mtp.py index 8bd3d4326..b3485b54b 100644 --- a/python/sglang/kernels/ops/kimi_k3/kda_decode_mtp.py +++ b/python/sglang/kernels/ops/kimi_k3/kda_decode_mtp.py @@ -13,6 +13,8 @@ import cutlass.cute as cute from cutlass._mlir.dialects import nvvm from cutlass.cute.nvgpu import cpasync +from sglang.kernels.jit.cute_aot_cache import get_jit_cache + WARP_SIZE = 32 TILE_K = 128 KERNEL_WIDTH = 4 @@ -890,11 +892,21 @@ def _block_threads(*, H: int, N: int) -> int: return BLOCK_THREADS_WIDE if H * N <= num_sms else BLOCK_THREADS_NARROW -_DSPARK_COMPILED = {} +_DSPARK_COMPILED = get_jit_cache( + "kimi_k3_kda_mtp_verify", + source_paths=(__file__,), + enable_tvm_ffi=False, +) -def _tensor_layout_key(tensor): - return (tensor.device, tensor.dtype, tuple(tensor.shape), tuple(tensor.stride())) +def _tensor_layout_key(tensor, fits_32bit_stride): + return ( + tensor.device, + tensor.dtype, + tuple(tensor.shape), + tuple(tensor.stride()), + fits_32bit_stride, + ) def _fits_32bit_stride(tensor): @@ -909,15 +921,17 @@ def _fits_32bit_stride(tensor): return True -def _cute_tensor(tensor, *, dynamic=True): +def _cute_tensor(tensor, *, dynamic=True, fits_32bit_stride=None): from cutlass.cute.runtime import from_dlpack if tensor.requires_grad: tensor = tensor.detach() + if fits_32bit_stride is None: + fits_32bit_stride = _fits_32bit_stride(tensor) value = from_dlpack( tensor, assumed_align=16, - use_32bit_stride=_fits_32bit_stride(tensor), + use_32bit_stride=fits_32bit_stride, ) leading_dim = next( (dim for dim, stride in enumerate(tensor.stride()) if stride == 1), None @@ -1097,7 +1111,9 @@ def fused_kda_decode_mtp_dspark( onorm_gate if apply_onorm else x_v, onorm_weight if apply_onorm else dt_bias, ) + fits_32bit = tuple(_fits_32bit_stride(tensor) for tensor in args) key = ( + torch.cuda.get_device_capability(), H, N, num_spec, @@ -1107,12 +1123,15 @@ def fused_kda_decode_mtp_dspark( float(onorm_eps) if apply_onorm else 0.0, float(scale), float(lower_bound), - *(_tensor_layout_key(tensor) for tensor in args), + *(_tensor_layout_key(tensor, fits) for tensor, fits in zip(args, fits_32bit)), ) stream = cuda.CUstream(torch.cuda.current_stream().cuda_stream) - compiled = _DSPARK_COMPILED.get(key) + compiled = _DSPARK_COMPILED[key] if key in _DSPARK_COMPILED else None if compiled is None: - cute_args = tuple(_cute_tensor(tensor, dynamic=False) for tensor in args) + cute_args = tuple( + _cute_tensor(tensor, dynamic=False, fits_32bit_stride=fits) + for tensor, fits in zip(args, fits_32bit) + ) compiled = cute.compile( _run_kda_decode_mtp_dspark, *cute_args, @@ -1129,7 +1148,10 @@ def fused_kda_decode_mtp_dspark( ) _DSPARK_COMPILED[key] = compiled compiled( - *(_cute_tensor(tensor) for tensor in args), + *( + _cute_tensor(tensor, fits_32bit_stride=fits) + for tensor, fits in zip(args, fits_32bit) + ), stream, ) return out diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index d5bf1ccb1..646b74f1b 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -1086,6 +1086,10 @@ class Envs: # Cache directories # =================================================================== SGLANG_CACHE_DIR = EnvStr(os.path.expanduser("~/.cache/sglang")) + # Persistent CuTe DSL AOT objects. Resolved lazily so it tracks + # SGLANG_CACHE_DIR; set to an empty string to keep compilation + # process-local. Must be trusted: cached objects are loaded into the process. + SGLANG_CUTE_AOT_CACHE_DIR = EnvStr(lambda: _default_cache_subdir("cute_aot")) # JIT kernel build cache. None = unset, resolving to ~/.cache/sglang/jit; # point it at a persistent mount to share builds across CI jobs. SGLANG_JIT_CACHE_DIR = EnvStr(None)