mirror of
https://github.com/tile-ai/tilelang.git
synced 2026-10-02 06:34:36 +08:00
[Enhancement] Speed up cold parallel/AOT compilation up to ~4x (#2809)
* [Enhancement] Speed up cold parallel/AOT compilation up to ~4x Three independent, low-risk changes (measured 45.6s -> 11.5s for 144 cold kernels on an H20, 180 cores): 1. cache: move kernel disk-save out of the global lock. KernelCache.cached() held the class lock around _save_kernel_to_disk (source + .so copy + cloudpickle), serializing every worker's save. The save stages+renames atomically and is idempotent, so it is already concurrency-safe without the lock (the disk load was already outside it). Lock now guards only _memory_cache. This is the bulk of the win (~3.1x). 2. jit: core-scale par_compile workers. Default was min(32, cpu+4); now min(len(funcs), available_cpus) when unset, with TILELANG_PAR_COMPILE_WORKERS override. Lowering is GIL-releasing C++ and nvcc is a subprocess, so threads parallelize. get_available_cpu_count is now cgroup-aware: it respects cpuset affinity AND caps at the cgroup CFS bandwidth quota (v2 cpu.max / v1 cpu.cfs_quota_us,period_us), so a CPU-throttled container does not oversubscribe nvcc jobs. 3. nvcc: opt-in parallel device compile (off by default). TL_NVCC_THREADS=N -> --threads (CUDA >= 11.2) + --split-compile (>= 12.1); scheduling-only, identical SASS. Helps large multi-kernel TUs; no effect on single-kernel TUs. get_available_cpu_count moved to tilelang.utils.device to break the autotuner<->jit import cycle; unit-tested in testing/python/utils. Adds benchmark/compile_speed/ to reproduce the numbers: a zoo of realistic inference kernels (GEMM, GQA attention, RMSNorm, SwiGLU, softmax) over Qwen2.5/Llama-3 dims (~126 distinct cold kernels at the default scale), with a before/after that reconstructs the pre-change baseline (lock-save + 32 workers) in one command (3.1x on the H20). * Simplify parallel compilation controls --------- Co-authored-by: SiriusNEO <chaofan@deepseek.com>
This commit is contained in:
@@ -0,0 +1,67 @@
|
||||
# Cold Parallel-Compilation Benchmark
|
||||
|
||||
Compiles a diverse zoo of realistic transformer inference kernels from a *cold*
|
||||
cache and reports the wall-clock before/after this PR. It exercises the parallel
|
||||
AOT path (`par_compile` → `KernelCache.cached` → nvcc) in the *many light
|
||||
kernels* regime (parallel AOT / autotune sweeps), where launch and disk-save
|
||||
serialization dominate.
|
||||
|
||||
## The zoo
|
||||
|
||||
`kernel_zoo.py` builds one translation unit per `(family, shape, tile)` over
|
||||
real model dims (Qwen2.5 / Llama-3). Families:
|
||||
|
||||
| family | kernel |
|
||||
|--------|--------|
|
||||
| `gemm` | dense q/k/v/o/gate/up/down projections |
|
||||
| `gqa` | GQA flash-attention forward (online softmax) |
|
||||
| `rmsnorm` | RMSNorm over the hidden dim |
|
||||
| `silu` | fused SwiGLU activation |
|
||||
| `softmax` | row-wise softmax (logits / router) |
|
||||
|
||||
Default `--scale 3` → **126 distinct kernels**; each is a genuine cold cache miss
|
||||
(throwaway `TILELANG_CACHE_DIR`, disjoint shapes) that actually runs nvcc.
|
||||
|
||||
## Environment
|
||||
|
||||
- GPU: `NVIDIA H20`, driver `535.161.08`, 180 CPU cores
|
||||
- nvcc: CUDA `12.9`
|
||||
|
||||
## How to Reproduce
|
||||
|
||||
```bash
|
||||
cd benchmark/compile_speed
|
||||
python benchmark_compile_speed.py # ~126 kernels
|
||||
python benchmark_compile_speed.py --scale 6 # larger zoo for many-core boxes
|
||||
```
|
||||
|
||||
`baseline` reconstructs the pre-PR behavior (disk-save under the global lock,
|
||||
capped at 32 workers); `current` is the shipped default (lock-free save, worker
|
||||
count scaled to cores).
|
||||
|
||||
## Results
|
||||
|
||||
126 cold kernels on an H20, `baseline` = global lock + `min(32, cores)` workers,
|
||||
`current` = lock-free save + `cores` workers. Speedup vs. core count (best of 2):
|
||||
|
||||
| cores | baseline (s) | current (s) | speedup |
|
||||
|-------|-------------|-------------|---------|
|
||||
| 8 | 50.7 | 39.2 | 1.29x |
|
||||
| 16 | 26.2 | 22.4 | 1.17x |
|
||||
| 32 | 19.9 | 17.8 | 1.12x |
|
||||
| 64 | 18.3 | 15.3 | 1.19x |
|
||||
| 180 | 18.4 | 14.7 | 1.25x |
|
||||
|
||||
Two effects compose. The **lock removal** helps at every core count (it is the
|
||||
whole win at `cores <= 32`, where both configs use the same worker count). The
|
||||
**worker scaling** only adds on top past 32 cores, where `current` outgrows the
|
||||
old `min(32, cores)` cap.
|
||||
|
||||
The largest speedups appear on a *fully cold* cache (first run of a fresh
|
||||
checkout / CI), where per-process PCH build is not yet amortized: 50.3s -> 16.2s
|
||||
= **3.1x** on the 180-core box. The table above is warm-toolchain (best of 2),
|
||||
so it isolates the compile-scheduling win from one-time startup cost.
|
||||
|
||||
Absolute numbers depend on core count, nvcc version, and disk speed; the
|
||||
reproducible result is the shape — lock removal is universal, worker scaling
|
||||
grows with cores and kernel count.
|
||||
@@ -0,0 +1,98 @@
|
||||
"""Cold parallel-compilation benchmark: a realistic inference kernel zoo.
|
||||
|
||||
Compiles a diverse zoo of ~real transformer inference kernels (GEMM projections,
|
||||
GQA flash-attention, RMSNorm, SwiGLU, softmax; see ``kernel_zoo.py``) from a
|
||||
*cold* cache and reports the wall-clock before/after this PR's two levers:
|
||||
|
||||
* **disk-save contention** — the per-kernel disk save no longer holds the global
|
||||
``KernelCache`` lock, so workers finishing near-simultaneously stop serializing
|
||||
their multi-file save behind one mutex.
|
||||
* **worker count** — ``par_compile`` now defaults to ``min(len(funcs),
|
||||
available_cpus)`` instead of the ``ThreadPoolExecutor`` default ``min(32, cpu+4)``.
|
||||
|
||||
The lock change is library code and cannot be toggled by a normal call, so the
|
||||
``baseline`` run reconstructs the old behavior (re-wrapping the save in a global
|
||||
lock, capped at 32 workers) for a self-contained before/after. Every kernel is a
|
||||
distinct cold cache miss (throwaway ``TILELANG_CACHE_DIR``, disjoint shapes), so
|
||||
each actually runs nvcc.
|
||||
|
||||
Usage::
|
||||
|
||||
python benchmark_compile_speed.py # default zoo (~126 kernels)
|
||||
python benchmark_compile_speed.py --scale 6 # larger zoo for many-core boxes
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import contextlib
|
||||
import os
|
||||
import shutil
|
||||
import tempfile
|
||||
import threading
|
||||
import time
|
||||
|
||||
import tilelang
|
||||
from tilelang.cache.kernel_cache import KernelCache
|
||||
from tilelang.utils.device import get_available_cpu_count
|
||||
|
||||
from kernel_zoo import build_zoo
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _serialized_save(enabled):
|
||||
"""Re-serialize the kernel disk-save behind one lock (the pre-PR behavior)."""
|
||||
if not enabled:
|
||||
yield
|
||||
return
|
||||
original = KernelCache._save_kernel_to_disk
|
||||
guard = threading.Lock()
|
||||
|
||||
def locked_save(self, *args, **kwargs):
|
||||
with guard:
|
||||
return original(self, *args, **kwargs)
|
||||
|
||||
KernelCache._save_kernel_to_disk = locked_save
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
KernelCache._save_kernel_to_disk = original
|
||||
|
||||
|
||||
def compile_zoo(scale, workers, salt, lock_save):
|
||||
"""Cold-compile the zoo at `workers`; return (num_kernels, wall_seconds)."""
|
||||
cache_dir = tempfile.mkdtemp(prefix="tl_compile_bench_")
|
||||
prev = os.environ.get("TILELANG_CACHE_DIR")
|
||||
os.environ["TILELANG_CACHE_DIR"] = cache_dir
|
||||
try:
|
||||
funcs = [pf for _, pf in build_zoo(scale=scale, salt=salt)]
|
||||
with _serialized_save(lock_save):
|
||||
start = time.time()
|
||||
tilelang.par_compile(funcs, target="cuda", num_workers=workers)
|
||||
return len(funcs), time.time() - start
|
||||
finally:
|
||||
if prev is None:
|
||||
os.environ.pop("TILELANG_CACHE_DIR", None)
|
||||
else:
|
||||
os.environ["TILELANG_CACHE_DIR"] = prev
|
||||
shutil.rmtree(cache_dir, ignore_errors=True)
|
||||
|
||||
|
||||
def main():
|
||||
cpu = get_available_cpu_count()
|
||||
parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
parser.add_argument("--scale", type=int, default=3, help="zoo replication factor (~44 kernels per unit)")
|
||||
args = parser.parse_args()
|
||||
|
||||
n = len(build_zoo(scale=args.scale))
|
||||
print(f"# {cpu} CPUs, {n} distinct cold kernels (scale={args.scale})")
|
||||
|
||||
compile_zoo(1, min(8, cpu), salt=9, lock_save=False) # warm up import/PCH
|
||||
|
||||
_, base = compile_zoo(args.scale, min(32, cpu), salt=1, lock_save=True)
|
||||
_, cur = compile_zoo(args.scale, cpu, salt=2, lock_save=False)
|
||||
print(f"baseline (global lock, {min(32, cpu)} workers) {base:7.2f}s")
|
||||
print(f"current (no lock, {cpu} workers) {cur:7.2f}s")
|
||||
print(f"speedup {base / cur:.2f}x")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,213 @@
|
||||
"""A zoo of realistic inference kernels for the compile-speed benchmark.
|
||||
|
||||
Each family mirrors a kernel that a real transformer inference stack compiles
|
||||
(GEMM projections, GQA flash-attention, RMSNorm, fused SiLU-gate, softmax), built
|
||||
from the corresponding ``examples/`` template so it actually compiles. Every
|
||||
family is swept over real model shapes (Qwen2.5 / Llama-3 dims), and each
|
||||
``(family, shape, tile)`` combination is a distinct translation unit — a genuine
|
||||
cold-cache miss that runs the full lowering + nvcc pipeline.
|
||||
|
||||
``build_zoo()`` returns a list of ``(name, prim_func)`` pairs; the benchmark
|
||||
compiles them with ``par_compile``. Keep the kernels light: the point is *many
|
||||
diverse* compiles (the parallel-AOT / autotune regime), not a few heavy ones.
|
||||
"""
|
||||
|
||||
import tilelang.language as T
|
||||
|
||||
# --- Realistic inference shapes ------------------------------------------------
|
||||
# (hidden, intermediate) pairs from common open models.
|
||||
MODEL_DIMS = [
|
||||
(2048, 5632), # Llama-3.2-1B-ish
|
||||
(3584, 18944), # Qwen2.5-7B
|
||||
(4096, 14336), # Llama-3-8B
|
||||
(5120, 13824), # Qwen2.5-14B-ish
|
||||
]
|
||||
# (num_q_heads, num_kv_heads, head_dim) GQA configs.
|
||||
ATTN_HEADS = [(32, 8, 128), (28, 4, 128), (40, 8, 128)]
|
||||
SEQ_TILES = [64, 128]
|
||||
GEMM_TILES = [(64, 64, 32), (128, 128, 32), (128, 64, 64)]
|
||||
|
||||
|
||||
def gemm(M, N, K, block_M, block_N, block_K, dtype="float16"):
|
||||
"""Dense projection GEMM (q/k/v/o/gate/up/down proj). examples/gemm."""
|
||||
|
||||
@T.prim_func
|
||||
def main(
|
||||
A: T.Tensor((M, K), dtype),
|
||||
B: T.Tensor((K, N), dtype),
|
||||
C: T.Tensor((M, N), dtype),
|
||||
):
|
||||
with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M), threads=128) as (bx, by):
|
||||
A_shared = T.alloc_shared((block_M, block_K), dtype)
|
||||
B_shared = T.alloc_shared((block_K, block_N), dtype)
|
||||
C_local = T.alloc_fragment((block_M, block_N), "float32")
|
||||
T.clear(C_local)
|
||||
for ko in T.Pipelined(T.ceildiv(K, block_K), num_stages=2):
|
||||
T.copy(A[by * block_M, ko * block_K], A_shared)
|
||||
T.copy(B[ko * block_K, bx * block_N], B_shared)
|
||||
T.gemm(A_shared, B_shared, C_local)
|
||||
T.copy(C_local, C[by * block_M, bx * block_N])
|
||||
|
||||
return main
|
||||
|
||||
|
||||
def gqa_attention(batch, heads, kv_heads, seq_len, dim, block_M, block_N):
|
||||
"""GQA flash-attention forward. examples/flash_attention/example_gqa_fwd."""
|
||||
scale = (1.0 / dim) ** 0.5 * 1.44269504
|
||||
q_shape = [batch, seq_len, heads, dim]
|
||||
kv_shape = [batch, seq_len, kv_heads, dim]
|
||||
group = heads // kv_heads
|
||||
dtype, accum = "float16", "float32"
|
||||
|
||||
@T.prim_func
|
||||
def main(
|
||||
Q: T.Tensor(q_shape, dtype),
|
||||
K: T.Tensor(kv_shape, dtype),
|
||||
V: T.Tensor(kv_shape, dtype),
|
||||
Output: T.Tensor(q_shape, dtype),
|
||||
):
|
||||
with T.Kernel(T.ceildiv(seq_len, block_M), heads, batch, threads=128) as (bx, by, bz):
|
||||
Q_shared = T.alloc_shared([block_M, dim], dtype)
|
||||
K_shared = T.alloc_shared([block_N, dim], dtype)
|
||||
V_shared = T.alloc_shared([block_N, dim], dtype)
|
||||
acc_s = T.alloc_fragment([block_M, block_N], accum)
|
||||
acc_s_cast = T.alloc_fragment([block_M, block_N], dtype)
|
||||
acc_o = T.alloc_fragment([block_M, dim], accum)
|
||||
scores_max = T.alloc_fragment([block_M], accum)
|
||||
scores_max_prev = T.alloc_fragment([block_M], accum)
|
||||
scores_scale = T.alloc_fragment([block_M], accum)
|
||||
scores_sum = T.alloc_fragment([block_M], accum)
|
||||
logsum = T.alloc_fragment([block_M], accum)
|
||||
kv_head = by // group
|
||||
|
||||
T.copy(Q[bz, bx * block_M : (bx + 1) * block_M, by, :], Q_shared)
|
||||
T.fill(acc_o, 0)
|
||||
T.fill(logsum, 0)
|
||||
T.fill(scores_max, -T.infinity(accum))
|
||||
for k in T.Pipelined(T.ceildiv(seq_len, block_N), num_stages=1):
|
||||
T.copy(K[bz, k * block_N : (k + 1) * block_N, kv_head, :], K_shared)
|
||||
T.clear(acc_s)
|
||||
T.gemm(Q_shared, K_shared, acc_s, transpose_B=True)
|
||||
T.copy(scores_max, scores_max_prev)
|
||||
T.reduce_max(acc_s, scores_max, dim=1, clear=False)
|
||||
for i in T.Parallel(block_M):
|
||||
scores_scale[i] = T.exp2((scores_max_prev[i] - scores_max[i]) * scale)
|
||||
for i, j in T.Parallel(block_M, block_N):
|
||||
acc_s[i, j] = T.exp2(acc_s[i, j] * scale - scores_max[i] * scale)
|
||||
T.reduce_sum(acc_s, scores_sum, dim=1)
|
||||
for i in T.Parallel(block_M):
|
||||
logsum[i] = logsum[i] * scores_scale[i] + scores_sum[i]
|
||||
T.copy(acc_s, acc_s_cast)
|
||||
for i, j in T.Parallel(block_M, dim):
|
||||
acc_o[i, j] *= scores_scale[i]
|
||||
T.copy(V[bz, k * block_N : (k + 1) * block_N, kv_head, :], V_shared)
|
||||
T.gemm(acc_s_cast, V_shared, acc_o)
|
||||
for i, j in T.Parallel(block_M, dim):
|
||||
acc_o[i, j] /= logsum[i]
|
||||
T.copy(acc_o, Output[bz, bx * block_M : (bx + 1) * block_M, by, :])
|
||||
|
||||
return main
|
||||
|
||||
|
||||
def rms_norm(M, N, block_M, dtype="float32"):
|
||||
"""RMSNorm over the hidden dim. examples/norm/rms_norm.py."""
|
||||
|
||||
@T.prim_func
|
||||
def main(A: T.Tensor((M, N), dtype), B: T.Tensor((M, N), dtype)):
|
||||
with T.Kernel(T.ceildiv(M, block_M), threads=128) as bx:
|
||||
A_local = T.alloc_fragment((block_M, N), dtype)
|
||||
A_pow = T.alloc_fragment((block_M, N), dtype)
|
||||
A_sum = T.alloc_fragment((block_M,), dtype)
|
||||
T.copy(A[bx * block_M, 0], A_local)
|
||||
for i, j in T.Parallel(block_M, N):
|
||||
A_pow[i, j] = A_local[i, j] * A_local[i, j]
|
||||
T.reduce_sum(A_pow, A_sum, dim=1)
|
||||
for i in T.Parallel(block_M):
|
||||
A_sum[i] = T.rsqrt(A_sum[i] / N + 1e-6)
|
||||
for i, j in T.Parallel(block_M, N):
|
||||
A_local[i, j] *= A_sum[i]
|
||||
T.copy(A_local, B[bx * block_M, 0])
|
||||
|
||||
return main
|
||||
|
||||
|
||||
def silu_gate(M, N, block_M, block_N, dtype="float16"):
|
||||
"""Fused SwiGLU activation: out = silu(gate) * up. FFN elementwise epilogue."""
|
||||
|
||||
@T.prim_func
|
||||
def main(
|
||||
Gate: T.Tensor((M, N), dtype),
|
||||
Up: T.Tensor((M, N), dtype),
|
||||
Out: T.Tensor((M, N), dtype),
|
||||
):
|
||||
with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M), threads=128) as (bx, by):
|
||||
g = T.alloc_fragment((block_M, block_N), dtype)
|
||||
u = T.alloc_fragment((block_M, block_N), dtype)
|
||||
T.copy(Gate[by * block_M, bx * block_N], g)
|
||||
T.copy(Up[by * block_M, bx * block_N], u)
|
||||
for i, j in T.Parallel(block_M, block_N):
|
||||
x = g[i, j].astype("float32")
|
||||
g[i, j] = (x * (1.0 / (1.0 + T.exp2(-x * 1.44269504)))).astype(dtype) * u[i, j]
|
||||
T.copy(g, Out[by * block_M, bx * block_N])
|
||||
|
||||
return main
|
||||
|
||||
|
||||
def softmax(M, N, block_M, dtype="float32"):
|
||||
"""Row-wise softmax (logits / router). Online max-sum reduction."""
|
||||
|
||||
@T.prim_func
|
||||
def main(A: T.Tensor((M, N), dtype), B: T.Tensor((M, N), dtype)):
|
||||
with T.Kernel(T.ceildiv(M, block_M), threads=128) as bx:
|
||||
row = T.alloc_fragment((block_M, N), dtype)
|
||||
row_max = T.alloc_fragment((block_M,), dtype)
|
||||
row_sum = T.alloc_fragment((block_M,), dtype)
|
||||
T.copy(A[bx * block_M, 0], row)
|
||||
T.reduce_max(row, row_max, dim=1, clear=True)
|
||||
for i, j in T.Parallel(block_M, N):
|
||||
row[i, j] = T.exp(row[i, j] - row_max[i])
|
||||
T.reduce_sum(row, row_sum, dim=1)
|
||||
for i, j in T.Parallel(block_M, N):
|
||||
row[i, j] = row[i, j] / row_sum[i]
|
||||
T.copy(row, B[bx * block_M, 0])
|
||||
|
||||
return main
|
||||
|
||||
|
||||
def build_zoo(scale=1, salt=0):
|
||||
"""Return ``[(name, prim_func), ...]`` — a diverse, realistic kernel set.
|
||||
|
||||
``scale`` replicates the sweep to grow the count for larger boxes without
|
||||
changing the kernel mix. ``salt`` shifts the row/token count so a second call
|
||||
compiles *fresh* kernels (distinct cache keys) rather than hitting the
|
||||
process-global in-memory cache from a prior run. Both perturb the *M* (row)
|
||||
dimension only — never the tile dims, which carry alignment constraints.
|
||||
"""
|
||||
zoo = []
|
||||
for rep in range(scale):
|
||||
# Perturb rows/tokens (M / seq_len), which is tile-alignment-agnostic, so
|
||||
# replicas and salted re-runs stay distinct cache keys without producing
|
||||
# illegal tile shapes.
|
||||
m = 256 + 128 * (rep + salt)
|
||||
for hidden, inter in MODEL_DIMS:
|
||||
for bm, bn, bk in GEMM_TILES:
|
||||
zoo.append((f"gemm_qkv_{hidden}_{bm}x{bn}x{bk}_{m}", gemm(m, hidden, hidden, bm, bn, bk)))
|
||||
zoo.append((f"gemm_up_{hidden}x{inter}_{bm}x{bn}x{bk}_{m}", gemm(m, inter, hidden, bm, bn, bk)))
|
||||
for h, kvh, hd in ATTN_HEADS:
|
||||
for bm in SEQ_TILES:
|
||||
zoo.append((f"gqa_{h}x{kvh}x{hd}_bm{bm}_{m}", gqa_attention(1, h, kvh, m, hd, bm, 64)))
|
||||
for hidden, inter in MODEL_DIMS:
|
||||
zoo.append((f"rmsnorm_{hidden}_{m}", rms_norm(m, hidden, 32)))
|
||||
zoo.append((f"silu_{inter}_{m}", silu_gate(m, inter, 32, 64)))
|
||||
zoo.append((f"softmax_{hidden}_{m}", softmax(m, hidden, 32)))
|
||||
return zoo
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
zoo = build_zoo()
|
||||
from collections import Counter
|
||||
|
||||
fams = Counter(name.split("_")[0] for name, _ in zoo)
|
||||
print(f"zoo size: {len(zoo)} kernels")
|
||||
for fam, n in sorted(fams.items()):
|
||||
print(f" {fam:10s} {n}")
|
||||
@@ -35,6 +35,7 @@ from pathlib import Path
|
||||
from tilelang.autotuner.param import CompileArgs, ProfileArgs, AutotuneResult
|
||||
from tilelang.autotuner.grouped_compile import compile_grouped_unit_tvm_ffi
|
||||
from tilelang.utils.language import get_prim_func_name
|
||||
from tilelang.utils.device import get_available_cpu_count
|
||||
from tilelang.autotuner.capture import get_autotune_inputs
|
||||
from tilelang.backend.target import determine_target
|
||||
from tilelang import __version__
|
||||
@@ -203,16 +204,6 @@ def _init_logger_handlers():
|
||||
_logger_handlers_initialized = True
|
||||
|
||||
|
||||
def get_available_cpu_count() -> int:
|
||||
"""Gets the number of CPU cores available to the current process."""
|
||||
try:
|
||||
cpu_count = len(os.sched_getaffinity(0))
|
||||
except AttributeError:
|
||||
cpu_count = os.cpu_count()
|
||||
|
||||
return cpu_count or 1
|
||||
|
||||
|
||||
def _normalize_value(value, sort_dict_items: bool = False):
|
||||
if isinstance(value, torch.Tensor):
|
||||
return ("tensor", str(value.dtype), tuple(value.shape), value.stride())
|
||||
|
||||
Vendored
+7
-6
@@ -401,12 +401,13 @@ class KernelCache:
|
||||
pass_configs=pass_configs,
|
||||
compile_flags=compile_flags,
|
||||
)
|
||||
with self._lock:
|
||||
if env.is_cache_enabled():
|
||||
cache_path = self._get_cache_path(key)
|
||||
self._save_kernel_to_disk(key, kernel, func, verbose)
|
||||
# Set cache path on adapter so it can save cubin after first execution
|
||||
self._set_adapter_cache_path(kernel, cache_path)
|
||||
# Save outside the lock: staging+rename is atomic and idempotent (like the
|
||||
# disk load above). Holding the lock here serialized every worker's save.
|
||||
if env.is_cache_enabled():
|
||||
cache_path = self._get_cache_path(key)
|
||||
self._save_kernel_to_disk(key, kernel, func, verbose)
|
||||
# Set cache path on adapter so it can save cubin after first execution
|
||||
self._set_adapter_cache_path(kernel, cache_path)
|
||||
|
||||
# Store in memory cache after compilation
|
||||
self._tag_kernel_cache_entry(kernel, key, self._get_cache_path(key))
|
||||
|
||||
@@ -25,6 +25,7 @@ from tvm.target import Target
|
||||
|
||||
from tilelang.jit.kernel import JITKernel
|
||||
from tilelang.cache import cached
|
||||
from tilelang.utils.device import get_available_cpu_count
|
||||
from os import path, makedirs
|
||||
from logging import getLogger
|
||||
from tilelang.jit.param import Kernel
|
||||
@@ -218,6 +219,15 @@ def par_compile(
|
||||
Set to "1", "true", "yes", or "on" to enable verbose compilation by default.
|
||||
"""
|
||||
|
||||
# funcs may be a one-shot iterable; materialize to size the pool and reuse below.
|
||||
funcs = list(funcs)
|
||||
|
||||
if num_workers is None and funcs:
|
||||
# Scale to available cores (affinity-aware), capped at #kernels; the stdlib
|
||||
# min(32, cpu+4) throttles large AOT batches. Lowering releases the GIL and
|
||||
# nvcc is a subprocess, so threads parallelize.
|
||||
num_workers = min(len(funcs), get_available_cpu_count())
|
||||
|
||||
with concurrent.futures.ThreadPoolExecutor(num_workers, "tl-par-comp") as executor:
|
||||
futures = []
|
||||
future_map = {}
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
import os
|
||||
|
||||
import torch
|
||||
|
||||
IS_CUDA = torch.cuda.is_available()
|
||||
@@ -19,3 +21,15 @@ def get_current_device():
|
||||
device = "mps:0"
|
||||
|
||||
return device
|
||||
|
||||
|
||||
def get_available_cpu_count() -> int:
|
||||
"""CPU cores available to this process (cpuset affinity), at least 1.
|
||||
|
||||
Falls back to ``os.cpu_count`` where affinity is unavailable (macOS/Windows).
|
||||
"""
|
||||
try:
|
||||
cpu_count = len(os.sched_getaffinity(0))
|
||||
except AttributeError:
|
||||
cpu_count = os.cpu_count()
|
||||
return max(1, cpu_count or 1)
|
||||
|
||||
Reference in New Issue
Block a user