[CUDA][Cache] Publish immutable binary cache directories (#3177)

* [CUDA][Cache] Publish immutable binary cache directories

* [Example] Drop stale torch JIT build locks in minference example

The vertical_slash_index extension builds into a cache directory shared
across CI runs. torch's FileBaton lock file is not tied to the owning
process, so a run cancelled mid-compile leaves it behind and every later
load() polls on it forever (hung CI at test_vs_sparse_attn). Guard the
build with an OS-level flock, which dies with its holder, and remove any
leftover FileBaton lock before invoking torch.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>

---------

Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
Lei Wang
2026-09-08 02:46:47 +08:00
committed by GitHub
co-authored by Claude Fable 5
parent 62bba8d20d
commit 7525499330
6 changed files with 365 additions and 76 deletions
@@ -16,6 +16,8 @@ from tilelang.profiler import do_bench
def _load_vertical_slash_index_ops():
import fcntl
from torch.utils.cpp_extension import load
current_dir = os.path.dirname(os.path.abspath(__file__))
@@ -55,7 +57,18 @@ def _load_vertical_slash_index_ops():
os.replace(tmp_path, stable_path)
stable_sources.append(stable_path)
return load(name=name, sources=stable_sources, build_directory=build_dir, verbose=False)
# torch's JIT build guards build_dir with a plain lock *file* (FileBaton)
# that is not tied to the owning process: a build killed mid-compile leaves
# the file behind and every later load() polls on it forever. Serialize
# builds with an OS-level flock instead (released automatically when the
# holder dies); any FileBaton lock still present while we hold the flock is
# necessarily stale, so drop it before handing over to torch.
baton_path = os.path.join(build_dir, "lock")
with open(os.path.join(extension_root, ".build.flock"), "w") as flock_file:
fcntl.flock(flock_file, fcntl.LOCK_EX)
if os.path.exists(baton_path):
os.remove(baton_path)
return load(name=name, sources=stable_sources, build_directory=build_dir, verbose=False)
@tilelang.jit(out_idx=[3])
+276 -12
View File
@@ -1,9 +1,19 @@
from __future__ import annotations
import builtins
import errno
import io
import json
import os
import threading
from concurrent.futures import ThreadPoolExecutor
from hashlib import sha256
from pathlib import Path
import pytest
import tilelang
from tilelang import tvm
import tilelang.cache.cuda_binary_cache as cuda_binary_cache_mod
import tilelang.cache.kernel_cache as kernel_cache_mod
from tilelang.backend import create_backend_context
from tilelang.cache.cuda_binary_cache import CUDABinaryCache
@@ -69,8 +79,11 @@ def test_cuda_binary_cache_hit_skips_nvcc_compile(monkeypatch, tmp_path):
# first compiles, second hits; third compiles (new options), fourth hits
assert len(compile_calls) == 2
assert compile_calls[0][3] != compile_calls[1][3]
cache_files = list((tmp_path / "cache").glob("*/cuda-binaries/*.cubin"))
cache_files = list((tmp_path / "cache").glob("*/cuda-binaries/*/kernel.cubin"))
assert len(cache_files) == 2
for cache_file in cache_files:
assert cache_file.read_bytes() == b"fake-cubin"
assert (cache_file.parent / "metadata.json").is_file()
def test_cuda_binary_cache_corrupted_entry_recompiles(monkeypatch, tmp_path):
@@ -93,30 +106,281 @@ def test_cuda_binary_cache_corrupted_entry_recompiles(monkeypatch, tmp_path):
cuda_backend.tilelang_callback_cuda_compile(source, target)
assert len(compile_calls) == 1
[cache_file] = (tmp_path / "cache").glob("*/cuda-binaries/*.cubin")
assert cache_file.with_name(cache_file.name + ".sha256").exists()
[cache_file] = (tmp_path / "cache").glob("*/cuda-binaries/*/kernel.cubin")
# Same-size corruption, as left behind by a crashed writer/filesystem client.
cache_file.write_bytes(b"\x00" * len(b"fake-cubin"))
corrupted = b"\x00" * cache_file.stat().st_size
cache_file.write_bytes(corrupted)
metadata_path = cache_file.parent / "metadata.json"
original_metadata = metadata_path.read_bytes()
recompiled = cuda_backend.tilelang_callback_cuda_compile(source, target)
assert bytes(recompiled) == b"fake-cubin"
assert len(compile_calls) == 2
# The corrupted entry was rewritten, so the next call hits the cache again.
cuda_backend.tilelang_callback_cuda_compile(source, target)
assert len(compile_calls) == 2
# Shared entries stay immutable even after a miss. Recompilation is usable
# for this call, but repairing the on-disk entry needs offline cleanup.
assert bytes(cuda_backend.tilelang_callback_cuda_compile(source, target)) == b"fake-cubin"
assert len(compile_calls) == 3
assert cache_file.read_bytes() == corrupted
assert metadata_path.read_bytes() == original_metadata
def test_cuda_binary_cache_accepts_legacy_entry_without_sidecar(monkeypatch, tmp_path):
def test_cuda_binary_cache_directory_coexists_with_legacy_sidecar_entry(monkeypatch, tmp_path):
_set_cache_dirs(monkeypatch, tmp_path)
key = "legacy-key"
path = CUDABinaryCache.get_path(key, "cubin")
os.makedirs(os.path.dirname(path), exist_ok=True)
with open(path, "wb") as f:
cache_root = CUDABinaryCache._get_cache_root()
legacy_path = os.path.join(cache_root, f"{key}.cubin")
os.makedirs(cache_root, exist_ok=True)
with open(legacy_path, "wb") as f:
f.write(b"legacy-cubin")
with open(legacy_path + ".sha256", "w") as f:
f.write(sha256(b"legacy-cubin").hexdigest())
assert CUDABinaryCache.load(key, "cubin") == b"legacy-cubin"
assert CUDABinaryCache.load(key, "cubin") is None
assert os.path.exists(legacy_path)
CUDABinaryCache.save(key, "cubin", b"new-cubin")
assert CUDABinaryCache.load(key, "cubin") == b"new-cubin"
assert Path(legacy_path).read_bytes() == b"legacy-cubin"
assert Path(legacy_path + ".sha256").read_text() == sha256(b"legacy-cubin").hexdigest()
@pytest.mark.parametrize("failure", ["missing", "empty", "malformed", "invalid-encoding", "wrong-type", "unreadable"])
def test_cuda_binary_cache_rejects_bad_metadata(monkeypatch, tmp_path, failure):
_set_cache_dirs(monkeypatch, tmp_path)
key = "bad-metadata-key"
CUDABinaryCache.save(key, "cubin", b"valid-cubin")
path = Path(CUDABinaryCache.get_path(key, "cubin"))
metadata_path = path.parent / "metadata.json"
metadata_path.unlink()
if failure == "unreadable":
metadata_path.mkdir()
elif failure != "missing":
contents = {"empty": b"", "malformed": b"{", "invalid-encoding": b"\xff", "wrong-type": b"[]"}
metadata_path.write_bytes(contents[failure])
assert CUDABinaryCache.load(key, "cubin") is None
assert path.read_bytes() == b"valid-cubin"
# A later writer must leave even an invalid published directory alone.
CUDABinaryCache.save(key, "cubin", b"new-cubin")
assert path.read_bytes() == b"valid-cubin"
assert CUDABinaryCache.load(key, "cubin") is None
@pytest.mark.parametrize(
"field,value",
[
("size", None),
("size", True),
("size", 0),
("size", "11"),
("sha256", None),
("sha256", ""),
("sha256", "x" * 64),
("sha256", " " * 64),
("format", "tilelang.cuda-binary-cache.v2"),
],
)
def test_cuda_binary_cache_rejects_invalid_metadata_fields(monkeypatch, tmp_path, field, value):
_set_cache_dirs(monkeypatch, tmp_path)
key = "invalid-metadata-key"
CUDABinaryCache.save(key, "cubin", b"valid-cubin")
path = Path(CUDABinaryCache.get_path(key, "cubin"))
metadata_path = path.parent / "metadata.json"
metadata = json.loads(metadata_path.read_text())
if value is None:
del metadata[field]
else:
metadata[field] = value
metadata_path.write_text(json.dumps(metadata))
assert CUDABinaryCache.load(key, "cubin") is None
assert path.exists()
assert json.loads(metadata_path.read_text()) == metadata
def test_cuda_binary_cache_rejects_empty_entry(monkeypatch, tmp_path):
_set_cache_dirs(monkeypatch, tmp_path)
key = "empty-key"
CUDABinaryCache.save(key, "fatbin", b"valid-fatbin")
path = Path(CUDABinaryCache.get_path(key, "fatbin"))
path.write_bytes(b"")
metadata_path = path.parent / "metadata.json"
metadata = json.loads(metadata_path.read_text())
metadata.update(size=0, sha256=sha256(b"").hexdigest())
metadata_path.write_text(json.dumps(metadata))
assert CUDABinaryCache.load(key, "fatbin") is None
assert path.exists()
@pytest.mark.parametrize("missing_metadata", [False, True])
def test_cuda_binary_cache_empty_read_is_miss(monkeypatch, tmp_path, missing_metadata):
_set_cache_dirs(monkeypatch, tmp_path)
key = "empty-read-key"
CUDABinaryCache.save(key, "fatbin", b"valid-fatbin")
path = Path(CUDABinaryCache.get_path(key, "fatbin"))
metadata_path = path.parent / "metadata.json"
def empty_read(file, *args, **kwargs):
if Path(file) == path:
# Model a 3FS fd returning EOF despite a nonempty cached binary.
return io.BytesIO(b"")
if missing_metadata and Path(file) == metadata_path:
raise FileNotFoundError(errno.ENOENT, "metadata disappeared")
return builtins.open(file, *args, **kwargs)
monkeypatch.setattr(cuda_binary_cache_mod, "open", empty_read, raising=False)
assert CUDABinaryCache.load(key, "fatbin") is None
assert path.read_bytes() == b"valid-fatbin"
assert metadata_path.exists()
def test_cuda_binary_cache_hash_mismatch_does_not_delete_entry(monkeypatch, tmp_path):
_set_cache_dirs(monkeypatch, tmp_path)
key = "mismatched-key"
CUDABinaryCache.save(key, "cubin", b"valid-cubin")
path = Path(CUDABinaryCache.get_path(key, "cubin"))
corrupted = b"other-cubin"
assert len(corrupted) == path.stat().st_size
path.write_bytes(corrupted)
assert CUDABinaryCache.load(key, "cubin") is None
assert path.read_bytes() == corrupted
assert (path.parent / "metadata.json").exists()
def test_cuda_binary_cache_rejects_truncated_binary(monkeypatch, tmp_path):
_set_cache_dirs(monkeypatch, tmp_path)
key = "truncated-binary-key"
CUDABinaryCache.save(key, "cubin", b"valid-cubin")
path = Path(CUDABinaryCache.get_path(key, "cubin"))
path.write_bytes(path.read_bytes()[:-1])
assert CUDABinaryCache.load(key, "cubin") is None
assert path.exists()
@pytest.mark.parametrize("compile_format", ["cubin", "fatbin"])
def test_cuda_binary_cache_first_valid_writer_wins(monkeypatch, tmp_path, compile_format):
_set_cache_dirs(monkeypatch, tmp_path)
key = "first-writer-key"
path = Path(CUDABinaryCache.get_path(key, compile_format))
CUDABinaryCache.save(key, compile_format, b"first-binary")
first_inode = path.stat().st_ino
with path.open("rb") as reader:
CUDABinaryCache.save(key, compile_format, b"second-binary")
assert reader.read() == b"first-binary"
assert CUDABinaryCache.load(key, compile_format) == b"first-binary"
assert path.stat().st_ino == first_inode
assert sorted(p.name for p in path.parent.iterdir()) == [f"kernel.{compile_format}", "metadata.json"]
def test_cuda_binary_cache_directory_publication_is_atomic(monkeypatch, tmp_path):
_set_cache_dirs(monkeypatch, tmp_path)
key = "atomic-key"
path = Path(CUDABinaryCache.get_path(key, "fatbin"))
synced_inodes = set()
real_fsync = os.fsync
real_rename = os.rename
publications = []
def track_fsync(fd):
real_fsync(fd)
stat = os.fstat(fd)
synced_inodes.add((stat.st_dev, stat.st_ino))
def check_publication(src, dst):
staging = Path(src)
assert staging.parent == Path(CUDABinaryCache._get_staging_root())
assert staging.stat().st_dev == path.parent.parent.stat().st_dev
assert not path.parent.exists()
assert CUDABinaryCache.load(key, "fatbin") is None
assert sorted(p.name for p in staging.iterdir()) == ["kernel.fatbin", "metadata.json"]
assert (staging / "kernel.fatbin").read_bytes() == b"valid-fatbin"
for entry in [staging, *staging.iterdir()]:
stat = entry.stat()
assert (stat.st_dev, stat.st_ino) in synced_inodes
real_rename(src, dst)
assert CUDABinaryCache.load(key, "fatbin") == b"valid-fatbin"
publications.append(dst)
monkeypatch.setattr(os, "fsync", track_fsync)
monkeypatch.setattr(os, "rename", check_publication)
CUDABinaryCache.save(key, "fatbin", b"valid-fatbin")
assert publications == [str(path.parent)]
assert not list(Path(CUDABinaryCache._get_staging_root()).iterdir())
def test_cuda_binary_cache_concurrent_publishers_do_not_replace_winner(monkeypatch, tmp_path):
_set_cache_dirs(monkeypatch, tmp_path)
writer_count = 16
barrier = threading.Barrier(writer_count)
payloads = [f"cubin-{index}".encode() for index in range(writer_count)]
real_rename = os.rename
winners = []
def concurrent_rename(src, dst):
# Force every writer to stage its files before any can publish.
barrier.wait(timeout=10)
real_rename(src, dst)
winners.append((Path(dst) / "kernel.cubin").read_bytes())
monkeypatch.setattr(os, "rename", concurrent_rename)
with ThreadPoolExecutor(max_workers=writer_count) as executor:
list(executor.map(lambda data: CUDABinaryCache.save("concurrent-key", "cubin", data), payloads))
assert len(winners) == 1
assert winners[0] in payloads
assert CUDABinaryCache.load("concurrent-key", "cubin") == winners[0]
cache_entries = os.listdir(CUDABinaryCache._get_cache_root())
assert cache_entries == ["concurrent-key"]
assert not list(Path(CUDABinaryCache._get_staging_root()).iterdir())
@pytest.mark.parametrize("operation", ["fsync", "rename"])
def test_cuda_binary_cache_failed_publish_cleans_only_own_staging(monkeypatch, tmp_path, operation):
_set_cache_dirs(monkeypatch, tmp_path)
staging_root = Path(CUDABinaryCache._get_staging_root())
other_writer = staging_root / "other-writer"
other_writer.mkdir(parents=True)
(other_writer / "kernel.cubin").write_bytes(b"other-cubin")
def fail(*args, **kwargs):
raise OSError(errno.EIO, "injected I/O failure")
monkeypatch.setattr(os, operation, fail)
with pytest.raises(OSError, match="injected I/O failure"):
CUDABinaryCache.save("failed-publish-key", "cubin", b"valid-cubin")
assert CUDABinaryCache.load("failed-publish-key", "cubin") is None
assert not Path(CUDABinaryCache.get_path("failed-publish-key", "cubin")).parent.exists()
assert list(staging_root.iterdir()) == [other_writer]
assert (other_writer / "kernel.cubin").read_bytes() == b"other-cubin"
def test_cuda_binary_cache_rejects_empty_save(monkeypatch, tmp_path):
_set_cache_dirs(monkeypatch, tmp_path)
with pytest.raises(ValueError, match="empty CUDA binary"):
CUDABinaryCache.save("empty-save-key", "fatbin", b"")
def test_disk_cache_load_failure_is_cache_miss(monkeypatch, tmp_path):
@@ -128,20 +128,6 @@ def test_load_removes_entry_with_same_size_corruption(cache_dirs, tmp_path, monk
assert not cache_path.exists()
def test_hash_check_disabled_still_catches_truncation(cache_dirs, tmp_path, monkeypatch):
monkeypatch.setattr(env, "TILELANG_CACHE_VERIFY_HASH", "0")
cache = KernelCache()
key = "size-only-check"
cache._save_kernel_to_disk(key, _make_fake_kernel(tmp_path))
cache_path = Path(cache._get_cache_path(key))
lib_file = cache_path / cache.kernel_lib_path
lib_file.write_bytes(b"short")
assert _load_expecting_no_build(cache, key, monkeypatch) is None
assert not cache_path.exists()
def test_source_files_are_size_checked_but_not_hashed(cache_dirs, tmp_path, monkeypatch):
# Hashing sources on load would regress lazy source loading; only their
# size is checked, so a same-size rewrite of a source file must not
+74 -41
View File
@@ -3,9 +3,11 @@
from __future__ import annotations
import contextlib
import errno
import functools
import json
import os
import shutil
import sys
import uuid
from hashlib import sha256
@@ -15,9 +17,14 @@ from tilelang import __version__
from tilelang.env import env
_CACHE_FORMAT = "tilelang.cuda-binary-cache.v1"
class CUDABinaryCache:
"""Cache cubin/fatbin bytes independently from host executable artifacts."""
# Each key is an immutable directory containing metadata.json and a raw
# kernel binary. Legacy `<key>.<format>` files and sidecars are ignored.
cache_root_dir = "cuda-binaries"
@staticmethod
@@ -44,6 +51,10 @@ class CUDABinaryCache:
def _get_cache_root(cls) -> str:
return os.path.join(cls._get_namespace_root(), cls.cache_root_dir)
@classmethod
def _get_staging_root(cls) -> str:
return os.path.join(cls._get_namespace_root(), ".staging", cls.cache_root_dir)
@staticmethod
@functools.cache
def _get_tilelang_lib_stamp() -> str | None:
@@ -120,12 +131,7 @@ class CUDABinaryCache:
@classmethod
def get_path(cls, key: str, compile_format: str) -> str:
filename = f"{key}.{compile_format}"
return os.path.join(cls._get_cache_root(), filename)
@classmethod
def _sidecar_path(cls, path: str) -> str:
return path + ".sha256"
return os.path.join(cls._get_cache_root(), key, f"kernel.{compile_format}")
@classmethod
def load(cls, key: str, compile_format: str) -> bytes | None:
@@ -135,55 +141,82 @@ class CUDABinaryCache:
try:
with open(path, "rb") as f:
data = f.read()
except FileNotFoundError:
with open(os.path.join(os.path.dirname(path), "metadata.json"), encoding="utf-8") as f:
metadata = json.load(f)
except (OSError, ValueError):
return None
if not env.should_verify_cache_hash():
return data
try:
with open(cls._sidecar_path(path)) as f:
expected_hash = f.read().strip()
except OSError:
# Entries written before content hashes were recorded.
return data
if sha256(data).hexdigest() == expected_hash:
return data
# Corrupted entry (e.g. truncated by a crashed writer): feeding it to
# cuModuleLoadData would fail with CUDA_ERROR_INVALID_IMAGE on every
# future run. Drop it so the caller recompiles and rewrites it.
for stale in (path, cls._sidecar_path(path)):
with contextlib.suppress(OSError):
os.remove(stale)
return None
if not isinstance(metadata, dict) or metadata.get("format") != _CACHE_FORMAT:
return None
# Empty/short reads and missing metadata are always misses. Never delete
# shared cache entries: on 3FS, unlinking a binary can invalidate another
# reader's open fd.
payload_size = metadata.get("size")
if type(payload_size) is not int or payload_size <= 0 or len(data) != payload_size:
return None
if sha256(data).hexdigest() != metadata.get("sha256"):
return None
return data
@classmethod
def save(cls, key: str, compile_format: str, data: bytes) -> None:
if not data:
raise ValueError("Cannot cache an empty CUDA binary")
if not env.is_cache_enabled():
return
cache_root = cls._get_cache_root()
os.makedirs(cache_root, exist_ok=True)
path = cls.get_path(key, compile_format)
# Sidecar first: a crash between the two renames then leaves a hash
# without a payload (a plain cache miss) instead of an unverifiable
# payload.
cls._write_atomic(cls._sidecar_path(path), sha256(data).hexdigest().encode())
cls._write_atomic(path, data)
cache_path = os.path.dirname(path)
# Published directories are immutable, even if a load found corruption.
# The caller can use its fresh compilation; repairing shared entries
# requires offline cleanup to avoid invalidating concurrent readers.
if os.path.lexists(cache_path):
return
@classmethod
def _write_atomic(cls, path: str, data: bytes) -> None:
directory, filename = os.path.split(path)
# Atomic replacement requires the temporary file and destination to be
# on the same filesystem, so keep the temporary file next to the cache
# entry.
temp_path = os.path.join(directory, f".{filename}.{os.getpid()}_{uuid.uuid4().hex}.tmp")
staging_root = cls._get_staging_root()
os.makedirs(staging_root, exist_ok=True)
staging_path = os.path.join(staging_root, f"{key}.{os.getpid()}.{uuid.uuid4().hex}")
os.mkdir(staging_path)
try:
with open(temp_path, "wb") as f:
data = bytes(data)
metadata = {"format": _CACHE_FORMAT, "size": len(data), "sha256": sha256(data).hexdigest()}
with open(os.path.join(staging_path, os.path.basename(path)), "wb") as f:
f.write(data)
# Without this barrier a crash can persist the rename below
# before the file data, publishing a truncated binary.
f.flush()
os.fsync(f.fileno())
os.replace(temp_path, path)
with open(os.path.join(staging_path, "metadata.json"), "w", encoding="utf-8") as f:
json.dump(metadata, f, indent=2, sort_keys=True)
f.write("\n")
f.flush()
os.fsync(f.fileno())
cls._fsync_dir(staging_path)
try:
# Both roots live in the same namespace/filesystem. Rename
# publishes both files together and cannot overwrite another
# writer's nonempty directory: the first publication wins.
os.rename(staging_path, cache_path)
except OSError as exc:
if exc.errno not in (errno.EEXIST, errno.ENOTEMPTY):
raise
else:
cls._fsync_dir(cache_root)
cls._fsync_dir(staging_root)
finally:
# Only remove this writer's private staging directory.
shutil.rmtree(staging_path, ignore_errors=True)
@staticmethod
def _fsync_dir(path: str) -> None:
"""Best-effort durability barrier for the published directory entry."""
try:
fd = os.open(path, os.O_RDONLY)
except OSError:
return
try:
with contextlib.suppress(OSError):
os.remove(temp_path)
os.fsync(fd)
finally:
os.close(fd)
+1 -2
View File
@@ -823,13 +823,12 @@ class KernelCache:
required_names = {os.path.basename(path) for path in self._get_required_files(cache_path)}
if not required_names.issubset(files):
return False
verify_hash = env.should_verify_cache_hash()
for name, meta in entries:
path = os.path.join(cache_path, name)
try:
if os.path.getsize(path) != meta["size"]:
return False
if verify_hash and name in required_names and KernelCache._hash_file(path) != meta["sha256"]:
if name in required_names and KernelCache._hash_file(path) != meta["sha256"]:
return False
except (OSError, KeyError, TypeError):
return False
-6
View File
@@ -375,9 +375,6 @@ class Environment:
TILELANG_KERNEL_CACHE_USE_LIB_STAMP = EnvVar(
"TILELANG_KERNEL_CACHE_USE_LIB_STAMP", "0"
) # include native TileLang library content hash in kernel cache keys
TILELANG_CACHE_VERIFY_HASH = EnvVar(
"TILELANG_CACHE_VERIFY_HASH", "1"
) # verify content hashes of cached binary artifacts at load time (set to 0 to only check file sizes)
TILELANG_CLEANUP_TEMP_FILES = EnvVar(
"TILELANG_CLEANUP_TEMP_FILES", "1"
) # cleanup temporary compiler files/dirs after compilation (set to 0 to keep for debugging)
@@ -450,9 +447,6 @@ class Environment:
def should_use_kernel_cache_lib_stamp(self) -> bool:
return str(self.TILELANG_KERNEL_CACHE_USE_LIB_STAMP).lower() in ("1", "true", "yes", "on")
def should_verify_cache_hash(self) -> bool:
return str(self.TILELANG_CACHE_VERIFY_HASH).lower() in ("1", "true", "yes", "on")
def is_autotune_cache_disabled(self) -> bool:
return self.TILELANG_AUTO_TUNING_DISABLE_CACHE.lower() in ("1", "true", "yes", "on")