feat(kernels): generalize persistent CuTe JIT cache (#33911)

This commit is contained in:
James Liu
2026-09-02 19:59:04 -07:00
committed by GitHub
parent 28262c20df
commit 4229088a48
6 changed files with 487 additions and 308 deletions
@@ -83,6 +83,11 @@ SGLang supports various environment variables that can be used to configure its
<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>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_CUTE_AOT_CACHE_DIR</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>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.</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>{`{SGLANG_CACHE_DIR}`}/cute_aot</code></td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_PREFETCH_BLOCK_SIZE_MB</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Block size (in MB) for sequential checkpoint prefetch reads that warm the OS page cache before workers load weights via mmap</td>
+429
View File
@@ -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)
@@ -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()
@@ -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(
@@ -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
+4
View File
@@ -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)