[Ascend] Rebase onto upstream main (through #3186) and adopt dialect-owned T.Kernel

Merge upstream main (85fd8fc2) into the Ascend comparison snapshot. This brings
in #3177, #3181, #3182, #3185 (AutoSchedule -> AutoWarpSpecialization) and
#3186 (target-neutral T.Kernel with dialect-owned launch annotations).

Conflict resolution:

- AutoSchedule -> AutoWarpSpecialization (#3185) supersedes the snapshot's own
  disambiguation rename to `tl.cuda_auto_schedule`; upstream's name is already
  distinct from Ascend's boolean `tl.enable_auto_schedule`, so the CUDA rename is
  taken wholesale and `TL_ENABLE_AUTO_SCHEDULE` (Ascend-only) is kept.
  `testing/python/backend/test_tilelang_backend_auto_schedule.py` follows.
- `tilelang/language/kernel.py` is now upstream's target-neutral implementation,
  unchanged. The Ascend additions it used to carry (`check_ascend_availability`,
  `_current_target_is_ascend`, the `is_cpu`/`is_ascend` branches in `Kernel`, the
  annotation-sniffing grid/thread split in `KernelLaunchFrame.__enter__`, and the
  SimtVF hooks in the thread accessors) are relocated, not adapted: see below.
- `tilelang/language/__init__.py` keeps the fork's Ascend facade and gains
  upstream's typed `Kernel` re-export.

Ascend now owns its launch surface, so the shared launch path carries no backend
branches and no runtime hardware probe:

- `tilelang/ascend/language/kernel.py` defines the dialect's `Kernel(*blocks,
  prelude=)`: the NPU 1-D core grid, with no `threads=`/`cluster_dims=` to be
  rejected at runtime because Python rejects the keywords. It also owns the
  SimtVF thread-scope state and the `get_thread_binding(s)`/`get_thread_extent(s)`
  accessors that resolve against an active `T.SimtVF` scope.
- `MaterializeKernelLaunch` splits "how a launch dimension is lowered" from
  "whether SIMT threads exist", which the old single `lower_thread_binding` flag
  conflated: Ascend needs the grid materialized as a real block-level
  thread_extent while having no threadIdx at kernel scope. `lower_grid_binding`
  defaults to None and follows `lower_thread_binding`, so every existing caller
  and test keeps its meaning. `launch_dim_tags` lets a backend name extra launch
  dimensions; Ascend passes `cthread` for `T.MixedKernel`'s sub-block-id.
- `MixedKernelLaunch` populates the frame's new grid_vars/thread_vars fields and
  emits the thread placeholders like `KernelLaunch` does, so a body referencing a
  thread index gets the actionable diagnostic instead of an out-of-range lookup.
- `tilelang.is_cpu_kernel_frame` / `tilelang.is_npu_kernel_frame` are dropped from
  `src/transform/common/attr.h`: the explicit grid_vars/thread_vars fields replaced
  every reader the snapshot had added for them.

Combining them also removes the cross-backend landmine where, on any host with
torch_npu available, the shared `T.Kernel(N, threads=...)` raised even for CUDA
targets.
This commit is contained in:
LeiWang1999
2026-09-10 17:45:32 +08:00
93 changed files with 2620 additions and 940 deletions
+1 -2
View File
@@ -53,8 +53,7 @@ tests:
```python
@tilelang.testing.requires_cuda
@pytest.mark.parametrize("op", ["sum", "max"])
def test_cuda_packed_codegen(op):
...
def test_cuda_packed_codegen(op): ...
```
## Validate the Current Head
+7 -7
View File
@@ -275,26 +275,26 @@ jobs:
sysctl kern.corefile kern.coredump kern.sugid_coredump
- name: Run performance regression test
id: perf
run: |
source test_regression/bin/activate
OLD_PYTHON=./old/bin/python NEW_PYTHON=./new/bin/python \
PERF_REGRESSION_MD=regression_result.md PERF_REGRESSION_PNG=regression_result.png \
python ./maint/scripts/test_perf_regression.py
- name: Read markdown table
id: read_md
run: |
echo "content<<EOF" >> $GITHUB_OUTPUT
cat regression_result.md >> $GITHUB_OUTPUT
echo "EOF" >> $GITHUB_OUTPUT
# An `if:` without a status-check function is evaluated as `success() && ...`,
# which is false once the step above exits 1 on a regression. `!cancelled()`
# supplies the status check explicitly, so the report is still posted before
# the job goes red.
- name: Upload result image as artifact
if: ${{ !cancelled() && (steps.perf.outcome == 'success' || steps.perf.outcome == 'failure') }}
uses: actions/upload-artifact@v7
with:
name: perf-regression-${{ github.run_id }}
path: regression_result.png
- name: Post test results as PR comment
if: ${{ !cancelled() && (steps.perf.outcome == 'success' || steps.perf.outcome == 'failure') }}
uses: actions/github-script@v9
with:
github-token: ${{ secrets.GITHUB_TOKEN }}
+2 -2
View File
@@ -30,12 +30,12 @@ repos:
args: [--ignore-case]
files: ^docs/spelling_wordlist\.txt$
- repo: https://github.com/pre-commit/mirrors-clang-format
rev: v22.1.8 # sync with requirements-lint.txt
rev: v23.1.0 # sync with requirements-lint.txt
hooks:
- id: clang-format
types_or: [c++, c]
- repo: https://github.com/astral-sh/ruff-pre-commit
rev: v0.16.1 # sync with requirements-lint.txt
rev: v0.16.6 # sync with requirements-lint.txt
hooks:
- id: ruff-check
args: [--fix, --exit-non-zero-on-fix]
+1 -3
View File
@@ -133,9 +133,7 @@ TileLang exposes an explicit pass configuration key, `tilelang.PassConfigKey.TL_
from tilelang import transform
from tilelang.engine.phase import LowerAndLegalize
with transform.PassContext(
config={transform.PassConfigKey.TL_FORCE_LET_INLINE: True}
):
with transform.PassContext(config={transform.PassConfigKey.TL_FORCE_LET_INLINE: True}):
lowered_mod = LowerAndLegalize(input_mod, target)
```
@@ -46,9 +46,9 @@ sections explain the mechanisms behind this path.
```python
with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M), threads=128) as (bx, by):
T.gemm(
A[by * block_M:(by + 1) * block_M, 0:K],
B[0:K, bx * block_N:(bx + 1) * block_N],
C[by * block_M:(by + 1) * block_M, bx * block_N:(bx + 1) * block_N],
A[by * block_M : (by + 1) * block_M, 0:K],
B[0:K, bx * block_N : (bx + 1) * block_N],
C[by * block_M : (by + 1) * block_M, bx * block_N : (bx + 1) * block_N],
clear_accum=True,
)
```
+31 -25
View File
@@ -63,8 +63,7 @@ def main(
A: T.Tensor((m,), dtype),
B: T.Tensor((m + n,), dtype),
C: T.Tensor((n * k,), dtype),
):
...
): ...
```
This enables enforcing cross-tensor relationships like `len(B) == m + n` and `len(C) == n * k` at runtime.
@@ -88,6 +87,8 @@ Passing `None` raises: `main.A_handle is expected to have non-NULL pointer`.
2) Still must be non-NULL (constant-true branch)
```python
some_cond: bool = True
@T.prim_func
def main(A: T.Tensor((M, K), dtype)):
if some_cond:
@@ -97,6 +98,8 @@ def main(A: T.Tensor((M, K), dtype)):
3) Nullable (constant-false branch, statically unreachable)
```python
some_cond: bool = False
@T.prim_func
def main(A: T.Tensor((M, K), dtype)):
if some_cond:
@@ -201,6 +204,7 @@ def matmul_relu_kernel(
T.gemm(A_shared, B_shared, C_local)
T.copy(C_local, C[by * block_M, bx * block_N])
# For debugging, print the host source
print(matmul_relu_kernel.get_host_source())
```
@@ -258,8 +262,8 @@ print(fn.get_host_source())
```python
import torch
A = torch.empty((M, K), device='cuda', dtype=torch.float16)
B = torch.empty((K, N), device='cuda', dtype=torch.float16)
A = torch.empty((M, K), device="cuda", dtype=torch.float16)
B = torch.empty((K, N), device="cuda", dtype=torch.float16)
# Missing C
fn(A, B)
```
@@ -271,8 +275,8 @@ Fix: pass all arguments per the signature.
```python
import torch
B = torch.empty((K, N), device='cuda', dtype=torch.float16)
C = torch.empty((M, N), device='cuda', dtype=torch.float16)
B = torch.empty((K, N), device="cuda", dtype=torch.float16)
C = torch.empty((M, N), device="cuda", dtype=torch.float16)
fn(1, B, C)
```
Expected: `<kernel>: Expect arg[0] to be pointer`.
@@ -283,9 +287,9 @@ Fix: pass a DLPack-compatible tensor (e.g., torch.Tensor).
```python
import torch
A = torch.empty((M, K, 1), device='cuda', dtype=torch.float16) # rank=3
B = torch.empty((K, N), device='cuda', dtype=torch.float16)
C = torch.empty((M, N), device='cuda', dtype=torch.float16)
A = torch.empty((M, K, 1), device="cuda", dtype=torch.float16) # rank=3
B = torch.empty((K, N), device="cuda", dtype=torch.float16)
C = torch.empty((M, N), device="cuda", dtype=torch.float16)
fn(A, B, C)
```
Expected: `<kernel>.A_handle.ndim is expected to equal 2, but got mismatched ndim`.
@@ -296,9 +300,9 @@ Fix: ensure runtime rank equals compiled rank.
```python
import torch
A = torch.empty((M, K), device='cuda', dtype=torch.float32) # should be float16
B = torch.empty((K, N), device='cuda', dtype=torch.float16)
C = torch.empty((M, N), device='cuda', dtype=torch.float16)
A = torch.empty((M, K), device="cuda", dtype=torch.float32) # should be float16
B = torch.empty((K, N), device="cuda", dtype=torch.float16)
C = torch.empty((M, N), device="cuda", dtype=torch.float16)
fn(A, B, C)
```
Expected: `<kernel>.A_handle.dtype is expected to be float16, but got incompatible dtype`.
@@ -309,9 +313,9 @@ Fix: `A = A.to(torch.float16)` or create with the correct dtype.
```python
import torch
A = torch.empty((M, K + 1), device='cuda', dtype=torch.float16) # K mismatched
B = torch.empty((K, N), device='cuda', dtype=torch.float16)
C = torch.empty((M, N), device='cuda', dtype=torch.float16)
A = torch.empty((M, K + 1), device="cuda", dtype=torch.float16) # K mismatched
B = torch.empty((K, N), device="cuda", dtype=torch.float16)
C = torch.empty((M, N), device="cuda", dtype=torch.float16)
fn(A, B, C)
```
Expected: `Argument <kernel>.A_handle.shape[i] has an unsatisfied constraint: ... == <expected>`.
@@ -322,10 +326,10 @@ Fix: satisfy linear constraints and constants across tensors.
```python
import torch
A = torch.empty((M, K), device='cuda', dtype=torch.float16)
A = torch.empty((M, K), device="cuda", dtype=torch.float16)
A_nc = A.t() # transpose -> non-contiguous
B = torch.empty((K, N), device='cuda', dtype=torch.float16)
C = torch.empty((M, N), device='cuda', dtype=torch.float16)
B = torch.empty((K, N), device="cuda", dtype=torch.float16)
C = torch.empty((M, N), device="cuda", dtype=torch.float16)
fn(A_nc, B, C)
```
Expected: `Argument <kernel>.A_handle.strides[1] has an unsatisfied constraint: ... == 1`.
@@ -336,9 +340,9 @@ Fix: pass `A_nc.contiguous()` or align the layout expectation in the kernel.
```python
import torch
A = torch.empty((M, K), device='cpu', dtype=torch.float16)
B = torch.empty((K, N), device='cpu', dtype=torch.float16)
C = torch.empty((M, N), device='cpu', dtype=torch.float16)
A = torch.empty((M, K), device="cpu", dtype=torch.float16)
B = torch.empty((K, N), device="cpu", dtype=torch.float16)
C = torch.empty((M, N), device="cpu", dtype=torch.float16)
fn(A, B, C) # CUDA-targeted kernel
```
Expected: `<kernel>.A_handle.device_type mismatch [expected: 2 (cuda)] ...`.
@@ -349,9 +353,9 @@ Fix: move tensors to the CUDA device.
```python
import torch
A = torch.empty((M, K), device='cuda:0', dtype=torch.float16)
B = torch.empty((K, N), device='cuda:1', dtype=torch.float16)
C = torch.empty((M, N), device='cuda:0', dtype=torch.float16)
A = torch.empty((M, K), device="cuda:0", dtype=torch.float16)
B = torch.empty((K, N), device="cuda:1", dtype=torch.float16)
C = torch.empty((M, N), device="cuda:0", dtype=torch.float16)
fn(A, B, C)
```
Expected: `Argument <kernel>.B_handle.device_id has an unsatisfied constraint: ... == ...`.
@@ -369,12 +373,14 @@ Fix: ensure valid underlying storage; in PyTorch scenarios, avoid constructing t
```python
import tilelang.language as T
@T.prim_func
def scalar_check(x: T.int32, flag: T.bool()):
T.evaluate(0)
scalar_check(1.0, True) # x is float -> Expect arg[0] to be int
scalar_check(1, 2.5) # flag is float -> Expect arg[1] to be boolean
scalar_check(1, 2.5) # flag is float -> Expect arg[1] to be boolean
```
Fix: pass correct scalar types, e.g., `scalar_check(1, True)`.
+5 -3
View File
@@ -114,9 +114,11 @@ One common strategy to address bank conflicts is shared memory swizzling. This t
Similarly, TileLang also supports shared memory swizzling. Users only need to add a single line of Python code:
```python
T.annotate_layout({
S_shared: TileLang.layout.make_swizzled_layout(S_shared),
})
T.annotate_layout(
{
S_shared: TileLang.layout.make_swizzled_layout(S_shared),
}
)
```
Here, `T.annotate_layout` allows users to specify any desired layout for a buffer. For convenience, TileLang provides the `make_swizzled_layout` primitive to automatically generate a swizzled layout.
+8 -6
View File
@@ -56,6 +56,8 @@ The vector add operation can also be extended to two-dimensional cases, where bo
```python
import tilelang.language as T
def elementwise_add(
M,
N,
@@ -67,15 +69,15 @@ def elementwise_add(
):
@T.prim_func
def main(
A: T.Tensor((M, N), in_dtype),
B: T.Tensor((M, N), in_dtype),
C: T.Tensor((M, N), out_dtype),
A: T.Tensor((M, N), in_dtype),
B: T.Tensor((M, N), in_dtype),
C: T.Tensor((M, N), out_dtype),
):
with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M), threads=threads) as (bx, by):
start_x = bx * block_N
start_y = by * block_M
for (local_y, local_x) in T.Parallel(block_M, block_N):
for local_y, local_x in T.Parallel(block_M, block_N):
y = start_y + local_y
x = start_x + local_x
@@ -237,8 +239,8 @@ def elementwise_add(N, NUM_ELE_PER_THREAD=8, threads=256, dtype=T.bfloat16):
# vector add.
for tid, i in T.Parallel(threads, NUM_ELE_PER_THREAD):
C_register[tid * NUM_ELE_PER_THREAD + i] = (
A_register[tid * NUM_ELE_PER_THREAD + i] +
B_register[tid * NUM_ELE_PER_THREAD + i])
A_register[tid * NUM_ELE_PER_THREAD + i] + B_register[tid * NUM_ELE_PER_THREAD + i]
)
# STG. 128
T.copy(
+47 -38
View File
@@ -23,8 +23,11 @@ A simple Triton kernel for GEMV might look like this:
```python
@triton.jit
def _gemv_naive(
x_ptr, A_ptr, y_ptr,
N, K,
x_ptr,
A_ptr,
y_ptr,
N,
K,
BLOCK_SIZE_K: tl.constexpr,
):
n = tl.program_id(0)
@@ -54,9 +57,9 @@ def naive_gemv(
@T.prim_func
def main(
A: T.Buffer((K,), dtype),
B: T.Buffer((N, K), dtype),
C: T.Buffer((N,), dtype),
A: T.Buffer((K,), dtype),
B: T.Buffer((N, K), dtype),
C: T.Buffer((N,), dtype),
):
with T.Kernel(T.ceildiv(N, BLOCK_N)) as bn:
tn = T.get_thread_binding(0) # tn = threadIdx.x
@@ -69,8 +72,7 @@ def naive_gemv(
A_shared[tk] = A[bk * BLOCK_K + tk]
B_shared[tn, tk] = B[bn * BLOCK_N + tn, bk * BLOCK_K + tk]
for tk in T.serial(BLOCK_K):
C_reg[0] += A_shared[tk].astype(accum_dtype) * B_shared[tn,
tk].astype(accum_dtype)
C_reg[0] += A_shared[tk].astype(accum_dtype) * B_shared[tn, tk].astype(accum_dtype)
C[bn * BLOCK_N + tn] = C_reg[0]
return main
@@ -137,9 +139,9 @@ def naive_splitk_gemv(
@T.prim_func
def main(
A: T.Buffer((K,), dtype),
B: T.Buffer((N, K), dtype),
C: T.Buffer((N,), dtype),
A: T.Buffer((K,), dtype),
B: T.Buffer((N, K), dtype),
C: T.Buffer((N,), dtype),
):
with T.Kernel(T.ceildiv(N, BLOCK_N), threads=(BLOCK_N, BLOCK_K)) as bn:
tn = T.get_thread_binding(0)
@@ -180,9 +182,9 @@ def splitk_gemv(
@T.prim_func
def main(
A: T.Buffer((K,), dtype),
B: T.Buffer((N, K), dtype),
C: T.Buffer((N,), dtype),
A: T.Buffer((K,), dtype),
B: T.Buffer((N, K), dtype),
C: T.Buffer((N,), dtype),
):
with T.Kernel(T.ceildiv(N, BLOCK_N), threads=(BLOCK_N, reduce_threads)) as bn:
tn = T.get_thread_binding(0)
@@ -225,9 +227,9 @@ def splitk_gemv_vectorized(
@T.prim_func
def main(
A: T.Buffer((K,), dtype),
B: T.Buffer((N, K), dtype),
C: T.Buffer((N,), dtype),
A: T.Buffer((K,), dtype),
B: T.Buffer((N, K), dtype),
C: T.Buffer((N,), dtype),
):
with T.Kernel(T.ceildiv(N, BLOCK_N), threads=(BLOCK_N, reduce_threads)) as bn:
tn = T.get_thread_binding(0)
@@ -272,9 +274,9 @@ def splitk_gemv_vectorized_tvm(
@T.prim_func
def main(
A: T.Buffer((K,), dtype),
B: T.Buffer((N, K), dtype),
C: T.Buffer((N,), dtype),
A: T.Buffer((K,), dtype),
B: T.Buffer((N, K), dtype),
C: T.Buffer((N,), dtype),
):
with T.Kernel(T.ceildiv(N, BLOCK_N), threads=(BLOCK_N, reduce_threads)) as bn:
tn = T.get_thread_binding(0)
@@ -292,9 +294,9 @@ def splitk_gemv_vectorized_tvm(
C_accum[0] += A_local[k].astype(accum_dtype) * B_local[k].astype(accum_dtype)
C_reduced = T.alloc_local((1,), accum_dtype)
with T.attr(
T.comm_reducer(lambda x, y: x + y, [T.cast(0, accum_dtype)]),
"reduce_scope",
T.reinterpret(T.uint64(0), dtype="handle"),
T.comm_reducer(lambda x, y: x + y, [T.cast(0, accum_dtype)]),
"reduce_scope",
T.reinterpret(T.uint64(0), dtype="handle"),
):
T.evaluate(
T.tvm_thread_allreduce(
@@ -304,7 +306,8 @@ def splitk_gemv_vectorized_tvm(
C_reduced[0],
tk,
dtype="handle",
))
)
)
C[bn * BLOCK_N + tn] = C_reduced[0]
@@ -323,14 +326,19 @@ def get_best_config(N, K):
def get_configs():
BLOCK_N = [2, 4, 8, 32, 64, 128]
reduce_threads = [4, 8, 32]
_configs = list(itertools.product(
BLOCK_N,
reduce_threads,
))
configs = [{
"BLOCK_N": c[0],
"reduce_threads": c[1],
} for c in _configs]
_configs = list(
itertools.product(
BLOCK_N,
reduce_threads,
)
)
configs = [
{
"BLOCK_N": c[0],
"reduce_threads": c[1],
}
for c in _configs
]
return configs
@autotune(
@@ -357,9 +365,9 @@ def get_best_config(N, K):
@T.prim_func
def main(
A: T.Buffer((K,), dtype),
B: T.Buffer((N, K), dtype),
C: T.Buffer((N,), dtype),
A: T.Buffer((K,), dtype),
B: T.Buffer((N, K), dtype),
C: T.Buffer((N,), dtype),
):
with T.Kernel(T.ceildiv(N, BLOCK_N), threads=(BLOCK_N, reduce_threads)) as bn:
tn = T.get_thread_binding(0)
@@ -377,9 +385,9 @@ def get_best_config(N, K):
C_accum[0] += A_local[k].astype(accum_dtype) * B_local[k].astype(accum_dtype)
C_reduced = T.alloc_local((1,), accum_dtype)
with T.attr(
T.comm_reducer(lambda x, y: x + y, [T.cast(0, accum_dtype)]),
"reduce_scope",
T.reinterpret(T.uint64(0), dtype="handle"),
T.comm_reducer(lambda x, y: x + y, [T.cast(0, accum_dtype)]),
"reduce_scope",
T.reinterpret(T.uint64(0), dtype="handle"),
):
T.evaluate(
T.tvm_thread_allreduce(
@@ -389,7 +397,8 @@ def get_best_config(N, K):
C_reduced[0],
tk,
dtype="handle",
))
)
)
C[bn * BLOCK_N + tn] = C_reduced[0]
+4 -2
View File
@@ -64,6 +64,7 @@ import tilelang
import tilelang.language as T
from tilelang.cuda.intrinsics import make_mma_swizzle_layout
def matmul(M, N, K, block_M, block_N, block_K, dtype="float16", accum_dtype="float"):
@T.prim_func
def main(
@@ -75,7 +76,7 @@ def matmul(M, N, K, block_M, block_N, block_K, dtype="float16", accum_dtype="flo
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), accum_dtype)
C_local = T.alloc_fragment((block_M, block_N), accum_dtype)
# Optional layout hints (commented out by default)
# T.annotate_layout({
@@ -105,6 +106,7 @@ def matmul(M, N, K, block_M, block_N, block_K, dtype="float16", accum_dtype="flo
return main
# 1. Create the TileLang function
func = matmul(1024, 1024, 1024, 128, 128, 32)
@@ -158,7 +160,7 @@ with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M), threads=128) as (bx,
```python
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), accum_dtype)
C_local = T.alloc_fragment((block_M, block_N), accum_dtype)
```
- `T.alloc_shared` allocates shared memory across the entire thread block.
+16 -9
View File
@@ -42,8 +42,9 @@ A compressor is provided with the sparse GEMM example in `examples/gemm_sp/spars
```python
from examples.gemm_sp.sparse_utils import compress
A_sparse, E = compress(A) # default: int16 metadata for fp16/bf16
A_sparse, E = compress(A.t().contiguous()) # compress the transposed layout
A_sparse, E = compress(A) # default: int16 metadata for fp16/bf16
A_sparse, E = compress(A.t().contiguous()) # compress the transposed layout
```
Here, `A_sparse` contains all the non-zero elements of `A`, while `E` stores the corresponding metadata (indexing information) required to reconstruct the original sparse pattern. The metadata uses a natural row-major layout that `T.gemm_sp` consumes directly — no additional layout annotation is needed.
@@ -58,11 +59,19 @@ The default metadata dtype for fp16/bf16 is `int16` with an E-factor of 16 (one
import tilelang.language as T
from examples.gemm_sp.sparse_utils import get_e_factor
def matmul_sp(
M, N, K,
block_M, block_N, block_K,
in_dtype, accum_dtype, e_dtype,
num_stages, threads,
M,
N,
K,
block_M,
block_N,
block_K,
in_dtype,
accum_dtype,
e_dtype,
num_stages,
threads,
policy=T.GemmWarpPolicy.Square,
):
e_factor = get_e_factor(in_dtype, e_dtype)
@@ -85,9 +94,7 @@ def matmul_sp(
T.copy(A_sparse[by * block_M, k * block_K // 2], A_shared)
T.copy(E[by * block_M, k * block_K // e_factor], E_shared)
T.copy(B[k * block_K, bx * block_N], B_shared)
T.gemm_sp(A_shared, E_shared, B_shared, C_local,
transpose_A=False, transpose_E=False, transpose_B=False,
policy=policy)
T.gemm_sp(A_shared, E_shared, B_shared, C_local, transpose_A=False, transpose_E=False, transpose_B=False, policy=policy)
T.copy(C_local, C_shared)
T.copy(C_shared, C[by * block_M, bx * block_N])
+6 -4
View File
@@ -26,6 +26,8 @@ To add options, pass a target config dictionary. For example:
```python
target = {"kind": "cuda", "arch": "sm_90"}
kernel = tilelang.compile(func, target=target, execution_backend="cython")
# or
@tilelang.jit(target=target)
def compiled_kernel(*args):
@@ -38,10 +40,10 @@ Most TileLang APIs that accept a target, such as `tilelang.compile`, `tilelang.j
same input forms:
```python
target = "auto" # detect CUDA, HIP, or Metal
target = "cuda" # bare TVM target kind
target = {"kind": "cuda", "arch": "sm_90"} # target config dict
target = tvm.target.Target({"kind": "cuda"}) # already-built TVM Target
target = "auto" # detect CUDA, HIP, or Metal
target = "cuda" # bare TVM target kind
target = {"kind": "cuda", "arch": "sm_90"} # target config dict
target = tvm.target.Target({"kind": "cuda"}) # already-built TVM Target
```
Use the bare string form for simple cases. Use a config dictionary when you need target attributes such as CUDA
+40 -26
View File
@@ -21,6 +21,7 @@ values from your config space.
import tilelang
import tilelang.language as T
def matmul_configs(M, N, K):
# Example space — tailor to your target
tiles = [64, 128]
@@ -35,17 +36,24 @@ def matmul_configs(M, N, K):
for TH in threads
]
@tilelang.autotune(configs=matmul_configs, warmup=25, rep=100, timeout=60)
@tilelang.jit(out_idx=[-1])
def matmul(M: int, N: int, K: int,
block_M: int = 128, block_N: int = 128, block_K: int = 32,
threads: int = 128, num_stages: int = 3,
dtype: str = 'float16', accum_dtype: str = 'float32'):
def matmul(
M: int,
N: int,
K: int,
block_M: int = 128,
block_N: int = 128,
block_K: int = 32,
threads: int = 128,
num_stages: int = 3,
dtype: str = "float16",
accum_dtype: str = "float32",
):
@T.prim_func
def kernel(A: T.Tensor((M, K), dtype),
B: T.Tensor((K, N), dtype),
C: T.Tensor((M, N), dtype)):
def kernel(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=threads) as (bx, by):
A_s = T.alloc_shared((block_M, block_K), dtype)
B_s = T.alloc_shared((block_K, block_N), dtype)
@@ -61,18 +69,21 @@ def matmul(M: int, N: int, K: int,
return kernel
# Usage
# Provide inputs via context (recommended for reproducibility across configs)
import torch
M = N = K = 1024
A = torch.randn(M, K, device='cuda', dtype=torch.float16)
B = torch.randn(K, N, device='cuda', dtype=torch.float16)
C = torch.empty(M, N, device='cuda', dtype=torch.float16)
A = torch.randn(M, K, device="cuda", dtype=torch.float16)
B = torch.randn(K, N, device="cuda", dtype=torch.float16)
C = torch.empty(M, N, device="cuda", dtype=torch.float16)
from tilelang.autotuner import set_autotune_inputs
with set_autotune_inputs(A, B, C):
tuned_kernel = matmul(M, N, K) # compiles, tunes, returns best kernel
tuned_kernel(A, B, C) # run best kernel
tuned_kernel = matmul(M, N, K) # compiles, tunes, returns best kernel
tuned_kernel(A, B, C) # run best kernel
```
Notes
@@ -94,22 +105,24 @@ kernel_factory = matmul # the function above (already @tilelang.jit)
tuner = AutoTuner.from_kernel(kernel_factory(M, N, K), configs=matmul_configs(M, N, K))
tuner.set_profile_args(
warmup=25, rep=100, timeout=60,
warmup=25,
rep=100,
timeout=60,
supply_type=tilelang.TensorSupplyType.Auto, # or provide supply_prog/ref_prog
ref_prog=lambda A, B, C: torch.allclose(C, (A @ B).to(C.dtype), rtol=1e-2, atol=1e-2),
)
tuner.set_compile_args(
target='auto', # or 'cuda'/'hip'/'metal'
execution_backend='auto', # resolves per-target
out_idx=[-1], # which outputs to return if multiple
pass_configs={ # optional TVM passes/flags
target="auto", # or 'cuda'/'hip'/'metal'
execution_backend="auto", # resolves per-target
out_idx=[-1], # which outputs to return if multiple
pass_configs={ # optional TVM passes/flags
# tilelang.PassConfigKey.EXAMPLE_KEY: value,
},
)
artifact = tuner.run() # compiles + runs + validates all configs
best_kernel = artifact.kernel # JITKernel
artifact = tuner.run() # compiles + runs + validates all configs
best_kernel = artifact.kernel # JITKernel
best_latency = artifact.latency
best_config = artifact.config
@@ -145,6 +158,7 @@ def supply_prog(signature):
# Return a list of torch tensors matching the kernel’s arguments
return [A, B, C]
tuner.set_profile_args(supply_prog=supply_prog)
```
@@ -234,12 +248,11 @@ configs and drive your own benchmarking:
@tilelang.jit
def factory(M, N, K, block_M=128, block_N=128, block_K=32):
@T.prim_func
def k(A: T.Tensor((M, K), 'float16'),
B: T.Tensor((K, N), 'float16'),
C: T.Tensor((M, N), 'float16')):
...
def k(A: T.Tensor((M, K), "float16"), B: T.Tensor((K, N), "float16"), C: T.Tensor((M, N), "float16")): ...
return k
impl = factory # JITImpl
cfgs = [
dict(block_M=64, block_N=128, block_K=32),
@@ -259,11 +272,13 @@ artifact = tuner.run() # AutotuneResult
# Save to disk
from pathlib import Path
save_dir = Path('out/best/matmul_1024')
save_dir = Path("out/best/matmul_1024")
artifact.save_to_disk(save_dir, verbose=True)
# Reload later
from tilelang.autotuner.param import AutotuneResult, CompileArgs
restored = AutotuneResult.load_from_disk(save_dir, CompileArgs())
best = restored.kernel
best(A, B, C)
@@ -287,8 +302,7 @@ def matmul_configs(M, N, K):
for BK in [32, 64]:
for S in [2, 3]:
for TH in [128, 256]:
yield dict(block_M=BM, block_N=BN, block_K=BK,
num_stages=S, threads=TH)
yield dict(block_M=BM, block_N=BN, block_K=BK, num_stages=S, threads=TH)
```
## Device and Backend Selection
+16 -22
View File
@@ -20,8 +20,8 @@ memory via the `shared::cluster` address space.
```python
with T.ClusterKernel(grid_x, grid_y, threads=128, cluster_dims=(4, 1, 1)) as (bx, by):
rank = T.block_rank_in_cluster() # 0..3 within this cluster
T.cluster_sync() # barrier across all CTAs in cluster
rank = T.block_rank_in_cluster() # 0..3 within this cluster
T.cluster_sync() # barrier across all CTAs in cluster
```
---
@@ -61,6 +61,7 @@ outside the mask perform a regular TMA load for their own tile.
import tilelang
import tilelang.language as T
def make_tma_multicast_kernel(M, N, block_M, block_N, cluster_mask):
@T.prim_func
def kernel(
@@ -68,19 +69,13 @@ def make_tma_multicast_kernel(M, N, block_M, block_N, cluster_mask):
B: T.Tensor((M, N), "float16"),
):
# 4 CTAs per cluster; ranks 0 and 1 share the same tile via multicast.
with T.ClusterKernel(
T.ceildiv(N, block_N),
T.ceildiv(M, block_M),
threads=128,
cluster_dims=(4, 1, 1)
) as (bx, by):
with T.ClusterKernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M), threads=128, cluster_dims=(4, 1, 1)) as (bx, by):
A_shared = T.alloc_shared((block_M, block_N), "float16")
# cluster_mask=0b0011: ranks 0 and 1 participate.
# Rank 0 issues tma_load_multicast; rank 1 receives passively.
# Ranks 2 and 3 each issue a regular tma_load.
T.copy_cluster(A[by * block_M, bx * block_N], A_shared,
cluster_mask=cluster_mask)
T.copy_cluster(A[by * block_M, bx * block_N], A_shared, cluster_mask=cluster_mask)
T.copy(A_shared, B[by * block_M, bx * block_N])
@@ -159,6 +154,7 @@ Steps:
import tilelang
import tilelang.language as T
@tilelang.jit(execution_backend="cython")
def make_cluster_copy_kernel(N: int):
@T.prim_func
@@ -167,8 +163,8 @@ def make_cluster_copy_kernel(N: int):
B: T.Tensor((N,), "float32"),
):
with T.ClusterKernel(2, threads=128, cluster_dims=(2, 1, 1)) as pid:
s_src = T.alloc_shared((N,), "float32")
s_dst = T.alloc_shared((N,), "float32")
s_src = T.alloc_shared((N,), "float32")
s_dst = T.alloc_shared((N,), "float32")
s_barrier = T.alloc_cluster_barrier([1])
T.fill(s_src, 0.0)
@@ -182,8 +178,7 @@ def make_cluster_copy_kernel(N: int):
s_src[i] = A[i]
# Async-push s_src → s_dst in CTA 1, signal CTA 1's barrier.
T.copy_cluster(s_src, s_dst, dst_block=1,
remote_barrier=s_barrier[0])
T.copy_cluster(s_src, s_dst, dst_block=1, remote_barrier=s_barrier[0])
if pid == 1:
# Wait until CTA 0 finishes writing.
@@ -218,7 +213,7 @@ completes only after all rows are transferred.
# 2-D non-contiguous copy: N_tile < N_full → compiler emits M TMA calls
s_src = T.alloc_shared((M, N_full), "float32")
s_dst = T.alloc_shared((M, N_full), "float32")
s_barrier = T.alloc_cluster_barrier([1]) # arrive_count updated to M at compile time
s_barrier = T.alloc_cluster_barrier([1]) # arrive_count updated to M at compile time
T.copy_cluster(
s_src[0:M, 0:N_tile],
@@ -301,11 +296,11 @@ SM-to-SM copy (saving global-memory round trips).
@T.prim_func
def split_k_gemm(A, B, C):
with T.ClusterKernel(grid_x, grid_y, threads=256, cluster_dims=(4, 1, 1)) as (bx, by):
rank = T.block_rank_in_cluster()
A_s = T.alloc_shared((BM, BK), "float16")
B_s = T.alloc_shared((BK, BN), "float16")
C_f = T.alloc_fragment((BM, BN), "float32")
C_s = T.alloc_shared((BM, BN), "float32")
rank = T.block_rank_in_cluster()
A_s = T.alloc_shared((BM, BK), "float16")
B_s = T.alloc_shared((BK, BN), "float16")
C_f = T.alloc_fragment((BM, BN), "float32")
C_s = T.alloc_shared((BM, BN), "float32")
barrier = T.alloc_cluster_barrier([3])
T.clear(C_f)
@@ -333,8 +328,7 @@ def split_k_gemm(A, B, C):
if rank != 0:
# Push this rank's slot to the *same* slot index in rank 0's
# C_parts — different offsets, so no destination race.
T.copy_cluster(C_parts[rank], C_parts[rank],
dst_block=0, remote_barrier=barrier[0])
T.copy_cluster(C_parts[rank], C_parts[rank], dst_block=0, remote_barrier=barrier[0])
if rank == 0:
T.mbarrier_wait_parity(barrier[0], 0) # wakes after all 3 arrivals
+7 -7
View File
@@ -21,13 +21,13 @@ treated as compile‑time constants and will be folded.
```python
for i in T.serial(N):
if i < N: # TIR condition
if i < N: # TIR condition
C[i] = A[i] + B[i]
else:
pass
# Ternary
x = (A[i] if i < N else 0)
x = A[i] if i < N else 0
```
Short‑circuit boolean ops are supported. For multi‑dimensional bounds, use
@@ -52,10 +52,10 @@ Boundary handling note
```python
for i in T.serial(N):
... # 0..N-1
... # 0..N-1
for i in T.serial(0, N, 2):
... # 0, 2, 4, ...
... # 0, 2, 4, ...
```
### Unroll
@@ -96,7 +96,7 @@ pipelines.
for ko in T.Pipelined(T.ceildiv(K, BK), num_stages=3):
T.copy(A[by * BM, ko * BK], A_s) # stage: copy A tile
T.copy(B[ko * BK, bx * BN], B_s) # stage: copy B tile
T.gemm(A_s, B_s, C_f) # stage: compute
T.gemm(A_s, B_s, C_f) # stage: compute
```
For manual `stage` / `order` annotations and the rule that scalar `Bind`
@@ -138,7 +138,7 @@ the explicit guard can be omitted when you don’t need a custom edge path.
for i, j in T.Parallel(M, N):
gi = by * BM + i
gj = bx * BN + j
if T.all_of(gi < M, gj < N): # optional in many cases
if T.all_of(gi < M, gj < N): # optional in many cases
C[gi, gj] = A[gi, gj] + B[gi, gj]
```
@@ -149,5 +149,5 @@ from a single thread to avoid duplicate outputs.
```python
if i == 0:
T.print(C, msg='C tile:')
T.print(C, msg="C tile:")
```
+8 -8
View File
@@ -117,22 +117,22 @@ If you need debugging or explicit checks:
```python
@T.prim_func
def gemm(
A: T.Tensor((M, K), 'float16'),
B: T.Tensor((K, N), 'float16'),
C: T.Tensor((M, N), 'float16'),
A: T.Tensor((M, K), "float16"),
B: T.Tensor((K, N), "float16"),
C: T.Tensor((M, N), "float16"),
):
with T.Kernel(T.ceildiv(N, BN), T.ceildiv(M, BM), threads=128) as (bx, by):
A_s = T.alloc_shared((BM, BK), 'float16')
B_s = T.alloc_shared((BK, BN), 'float16')
C_f = T.alloc_fragment((BM, BN), 'float32')
A_s = T.alloc_shared((BM, BK), "float16")
B_s = T.alloc_shared((BK, BN), "float16")
C_f = T.alloc_fragment((BM, BN), "float32")
T.clear(C_f)
for ko in T.Pipelined(T.ceildiv(K, BK), num_stages=3):
T.copy(A[by * BM, ko * BK], A_s) # Global → Shared
T.copy(B[ko * BK, bx * BN], B_s)
T.gemm(A_s, B_s, C_f) # compute into fragment
T.gemm(A_s, B_s, C_f) # compute into fragment
T.copy(C_f, C[by * BM, bx * BN]) # store back
T.copy(C_f, C[by * BM, bx * BN]) # store back
```
## Instruction Reference (Concise)
+63 -31
View File
@@ -26,11 +26,10 @@ Note on dtypes
```python
@T.prim_func
def add_kernel(
A: T.Tensor((N,), dtype), # dtype could be 'float32' | T.float32 | torch.float32
A: T.Tensor((N,), dtype), # dtype could be 'float32' | T.float32 | torch.float32
B: T.Tensor((N,), dtype),
C: T.Tensor((N,), dtype),
):
... # kernel body
): ... # kernel body
```
- Shapes may be concrete integers or symbolic. For symbolic, you can pass
@@ -39,10 +38,11 @@ def add_kernel(
```python
# Named symbolic dimension (optional)
K = T.dyn['K']
K = T.dyn["K"]
@T.prim_func
def uses_dyn(A: T.Tensor((K,), 'float32')):
...
def uses_dyn(A: T.Tensor((K,), "float32")): ...
```
### Dynamic symbolic dimensions: two ways
@@ -59,17 +59,22 @@ TileLang supports two complementary ways to introduce symbolic (dynamic) dims:
```python
# 1) Annotation-only symbol; read the bound size via shape
K = T.dyn['K'] # dtype defaults to int32
K = T.dyn["K"] # dtype defaults to int32
@T.prim_func
def foo(A: T.Tensor((K,), 'float32')):
def foo(A: T.Tensor((K,), "float32")):
N = A.shape[0]
for i in T.serial(N):
...
# 2) Explicit Var symbol usable in the body
K = T.dynamic('K', 'int32') # or T.dynamic('K') defaults to int32
K = T.dynamic("K", "int32") # or T.dynamic('K') defaults to int32
@T.prim_func
def bar(A: T.Tensor((K,), 'float32')):
def bar(A: T.Tensor((K,), "float32")):
for i in T.serial(K):
...
```
@@ -80,16 +85,40 @@ Notes
## 2. Launching Work with `T.Kernel`
`with T.Kernel(...)` declares a launch context and creates block/thread
bindings. For GPU backends, specify a grid and threads per block.
`with T.Kernel(...)` declares a grid of tile programs. The positional
arguments give the grid extent along each axis and the returned variables are
the program indices along those axes. This is the part of a launch every
target shares: on CUDA a program is a thread block and `bx`/`by` are
`blockIdx.x`/`blockIdx.y`; on CPU the grid becomes the outer loop.
```python
with T.Kernel(grid_x, grid_y, threads=128) as (bx, by):
... # bx/by are blockIdx.x/y
... # bx/by are the program indices (blockIdx.x/y on CUDA)
```
You rarely need raw thread indices; most kernels use structured loops
(`T.serial`, `T.unroll`, `T.Parallel`, `T.Pipelined`) inside a `T.Kernel`.
Keyword arguments are launch annotations that the backend interprets once the
target is known. Each language dialect's `Kernel` declares the annotations its
backend understands as explicit keyword parameters, so hovering or
autocompleting `T.Kernel` shows exactly those and anything else is rejected:
`tilelang.language` (the CUDA dialect) offers `threads`, `prelude` and
`cluster_dims`; `tilelang.rocm.language` / `tilelang.metal.language` offer
`threads` and `prelude`; `tilelang.cpu.language` offers only `prelude`.
`threads` is the SIMT one: how many threads run each tile program on
GPU-style backends. Those backends pick a default (128) when it is omitted;
a kernel written with the CUDA dialect still compiles for CPU, which ignores
the thread count. Code inside `T.Kernel` operates at the tile-program level,
so you rarely need raw thread indices; most kernels use structured loops
(`T.serial`, `T.unroll`, `T.Parallel`, `T.Pipelined`) that the compiler maps
onto threads. `T.get_thread_binding()` exposes the thread index for
thread-level code on SIMT targets; a kernel that uses it is rejected when
compiled for a target without SIMT threads.
`T.ClusterKernel(..., cluster_dims=...)` adds the CUDA thread-block-cluster
annotation (SM90+). A cluster is a `cluster_dims`-shaped tile of the grid, so
`T.get_cluster_id(axis)` is plain program-index arithmetic
(`bx // cluster_dims[axis]`) and works on every target, while
`T.block_rank_in_cluster()` reads the hardware rank and is CUDA-only. Targets
without clusters reject `cluster_dims` at compile time.
## 3. Loops and Control Flow
@@ -125,9 +154,9 @@ TileLang exposes key software‑managed scopes:
(`T.alloc_fragment`, `T.alloc_var`)
```python
A_shared = T.alloc_shared((BM, BK), 'float16')
B_shared = T.alloc_shared((BK, BN), 'float16')
C_local = T.alloc_fragment((BM, BN), 'float32')
A_shared = T.alloc_shared((BM, BK), "float16")
B_shared = T.alloc_shared((BK, BN), "float16")
C_local = T.alloc_fragment((BM, BN), "float32")
T.clear(C_local) # zero accumulators
```
@@ -162,8 +191,9 @@ import tilelang
import tilelang.language as T
from tilelang import jit
@jit # infers target from tensors at first call
def add(N: int, block: int = 256, dtype: str = 'float32'):
def add(N: int, block: int = 256, dtype: str = "float32"):
@T.prim_func
def add_kernel(
@@ -179,12 +209,14 @@ def add(N: int, block: int = 256, dtype: str = 'float32'):
return add_kernel
# Host side (PyTorch shown; NumPy/DLPack also supported)
import torch
N = 1 << 20
A = torch.randn(N, device='cuda', dtype=torch.float32)
B = torch.randn(N, device='cuda', dtype=torch.float32)
C = torch.empty(N, device='cuda', dtype=torch.float32)
A = torch.randn(N, device="cuda", dtype=torch.float32)
B = torch.randn(N, device="cuda", dtype=torch.float32)
C = torch.empty(N, device="cuda", dtype=torch.float32)
kernel = add(N)
kernel(A, B, C) # runs on GPU
@@ -204,14 +236,14 @@ fragment accumulator. It mirrors the quickstart style found in the repository.
```python
@T.prim_func
def gemm(
A: T.Tensor((M, K), 'float16'),
B: T.Tensor((K, N), 'float16'),
C: T.Tensor((M, N), 'float16'),
A: T.Tensor((M, K), "float16"),
B: T.Tensor((K, N), "float16"),
C: T.Tensor((M, N), "float16"),
):
with T.Kernel(T.ceildiv(N, BN), T.ceildiv(M, BM), threads=128) as (bx, by):
A_s = T.alloc_shared((BM, BK), 'float16')
B_s = T.alloc_shared((BK, BN), 'float16')
C_f = T.alloc_fragment((BM, BN), 'float32')
A_s = T.alloc_shared((BM, BK), "float16")
B_s = T.alloc_shared((BK, BN), "float16")
C_f = T.alloc_fragment((BM, BN), "float32")
T.clear(C_f)
for ko in T.Pipelined(T.ceildiv(K, BK), num_stages=3):
@@ -228,9 +260,9 @@ Use `T.print` inside a kernel for quick introspection. TileLang emits printing
from a single thread for shared/fragment scopes to avoid floods.
```python
T.print(C_f, msg='accumulator:')
T.print(A_s, msg='A tile:')
T.print(C[0], msg='C[0] = ')
T.print(C_f, msg="accumulator:")
T.print(A_s, msg="A tile:")
T.print(C[0], msg="C[0] = ")
```
## 9. Where to Go Next
+1 -1
View File
@@ -163,7 +163,7 @@ The same functionality is available in Python:
```python
from tilelang.tools.compile_only import compile_kernel_source
source = compile_kernel_source(add.get_tir()) # default target "c"
source = compile_kernel_source(add.get_tir()) # default target "c"
cuda_source = compile_kernel_source(add.get_tir(), "cuda") # pinned to sm_80
```
+1 -4
View File
@@ -309,10 +309,7 @@ launch = data["launches"][0]
names = data["stringTable"]
store_indices = [
marker["payloadVal"]
for marker in launch["markers"]
if names[marker["markerNameIdx"]] == "store_index"
and "payloadVal" in marker
marker["payloadVal"] for marker in launch["markers"] if names[marker["markerNameIdx"]] == "store_index" and "payloadVal" in marker
]
print(store_indices[:8])
```
+2 -2
View File
@@ -119,8 +119,8 @@ results = lt.lower_trace(func, transform.Simplify(), mode="terminal")
results = lt.lower_trace(
func,
[
("Annotate", tvm.tirx.transform.AnnotateDeviceRegions()),
("Split", tvm.tirx.transform.SplitHostDevice()),
("Annotate", tvm.tirx.transform.AnnotateDeviceRegions()),
("Split", tvm.tirx.transform.SplitHostDevice()),
("ThreadSync", transform.ThreadSync("shared")),
],
mode="both",
+13 -35
View File
@@ -46,22 +46,8 @@ def kernel(
Manually define configurations or use combinatorial generation:
```python
configs = [
{
"block_M": 128,
"block_N": 128,
"block_K": 128,
"num_stages": 3,
"thread_num": 128,
"enable_rasteration": True
},
{
"block_M": 32,
"block_N": 32,
"block_K": 32,
"num_stages": 0,
"thread_num": 32,
"enable_rasteration": False
},
{"block_M": 128, "block_N": 128, "block_K": 128, "num_stages": 3, "thread_num": 128, "enable_rasteration": True},
{"block_M": 32, "block_N": 32, "block_K": 32, "num_stages": 0, "thread_num": 32, "enable_rasteration": False},
# ...additional configurations...
]
```
@@ -83,30 +69,24 @@ _configs = list(
num_stages,
thread_num,
enable_rasterization,
))
)
)
configs = [
{
"block_M": c[0],
"block_N": c[1],
"block_K": c[2],
"num_stages": c[3],
"thread_num": c[4],
"enable_rasteration": c[5]
} for c in _configs
{"block_M": c[0], "block_N": c[1], "block_K": c[2], "num_stages": c[3], "thread_num": c[4], "enable_rasteration": c[5]}
for c in _configs
]
```
### Step 3: Compile and Benchmark
Configure JIT compilation and benchmarking settings:
```python
autotuner = AutoTuner.from_kernel(
kernel=kernel, configs=get_configs(M, N, K, with_roller)).set_compile_args(
out_idx=[-1],
supply_type=tl.TensorSupplyType.Integer,
ref_prog=ref_program,
skip_check=False,
target="auto",
)
autotuner = AutoTuner.from_kernel(kernel=kernel, configs=get_configs(M, N, K, with_roller)).set_compile_args(
out_idx=[-1],
supply_type=tl.TensorSupplyType.Integer,
ref_prog=ref_program,
skip_check=False,
target="auto",
)
result = autotuner.run(warmup=3, rep=20)
out_c = result.kernel(a, b)
```
@@ -135,7 +115,6 @@ roller_hints = carve_template.recommend_hints(topk=10)
# Configure candidate parameters
for hint in roller_hints:
# ...existing code...
config["block_M"] = block_m
@@ -144,5 +123,4 @@ for hint in roller_hints:
config["num_stages"] = hint.pipeline_stage
config["thread_num"] = block_rows * block_cols * 32
config["enable_rasteration"] = hint.rasterization_plan is not NoRasterization
```
+9 -6
View File
@@ -108,20 +108,22 @@ import tilelang.language as T
from tilelang import tvm
from tilelang.engine.callback import register_cuda_postproc_callback
@register_cuda_postproc_callback
def tilelang_callback_cuda_postproc(code, _):
print(code) # print the final CUDA code
print(code) # print the final CUDA code
code = "// modified by tilelang_callback_cuda_postproc\n" + code
return code
kernel = tilelang.compile(matmul, target="cuda")
kernel_source = kernel.get_kernel_source()
print(kernel_source)
'''
"""
// modified by tilelang_callback_cuda_postproc
#include "cuda_runtime.h"
...
'''
"""
```
### Runtime Debug Prints with `T.print`
@@ -262,9 +264,10 @@ The core helpers can also be used directly:
from tilelang.tools.pass_visualizer.viewer import build_pass_data, emit_html
name, stages = build_pass_data(
"path/to/kernel.py", factory=None, target="auto",
kwargs={"M": 1024, "N": 1024, "K": 1024,
"block_M": 128, "block_N": 128, "block_K": 32},
"path/to/kernel.py",
factory=None,
target="auto",
kwargs={"M": 1024, "N": 1024, "K": 1024, "block_M": 128, "block_N": 128, "block_K": 32},
source=open("path/to/kernel.py").read(),
)
html = emit_html(name, stages)
+1 -1
View File
@@ -16,7 +16,7 @@ PASS_CFG = {
@tilelang.jit(
pass_configs={
**PASS_CFG,
tilelang.PassConfigKey.TL_CUDA_AUTO_SCHEDULE: "role_based",
tilelang.PassConfigKey.TL_ENABLE_AUTO_WARP_SPECIALIZATION: "role_based",
}
)
def flash_attention(
+1 -1
View File
@@ -7,7 +7,7 @@ from tilelang.carver.arch import driver
from tilelang.profiler import do_bench
@tilelang.jit(pass_configs={tilelang.PassConfigKey.TL_CUDA_AUTO_SCHEDULE: "role_based"})
@tilelang.jit(pass_configs={tilelang.PassConfigKey.TL_ENABLE_AUTO_WARP_SPECIALIZATION: "role_based"})
def gemm(
A,
B,
+5 -3
View File
@@ -105,9 +105,11 @@ One common strategy to address bank conflicts is shared memory swizzling. This t
Similarly, TileLang also supports shared memory swizzling. Users only need to add a single line of Python code:
```python
T.annotate_layout({
S_shared: TileLang.layout.make_swizzled_layout(S_shared),
})
T.annotate_layout(
{
S_shared: TileLang.layout.make_swizzled_layout(S_shared),
}
)
```
Here, `T.annotate_layout` allows users to specify any desired layout for a buffer. For convenience, TileLang provides the `make_swizzled_layout` primitive to automatically generate a swizzled layout.
+7 -10
View File
@@ -45,9 +45,7 @@ After the matmul, we apply ReLU and aggregate across heads with learned weights:
```python
for bn_i, bq_i, h_i in T.Parallel(block_N, block_Q, heads):
s_reshaped[bn_i, bq_i, h_i] = (
T.max(s[bn_i, bq_i * heads + h_i], 0) * weights[bq_i, h_i]
) * index_k_scale_fragment[bn_i]
s_reshaped[bn_i, bq_i, h_i] = (T.max(s[bn_i, bq_i * heads + h_i], 0) * weights[bq_i, h_i]) * index_k_scale_fragment[bn_i]
T.reduce_sum(s_reshaped, logits, dim=-1, clear=True)
```
@@ -71,7 +69,7 @@ The implementation uses a radix-sort-based approach that processes floats as uns
```python
for s in T.serial(T.ceildiv(seq_len, BLOCK_SIZE)):
input_idx = s*BLOCK_SIZE+tx
input_idx = s * BLOCK_SIZE + tx
if input_idx < l_end_idx and input_idx >= l_start_idx and input_idx < seq_len:
inval_int16 = convert_to_uint16(input[bx, input_idx])
T.atomic_add(s_histogram[inval_int16], 1)
@@ -88,7 +86,7 @@ Elements above the threshold go directly to the output. Elements in the threshol
```python
if l_bin_id32 > l_threshold_bin_id:
pos = T.atomic_add(s_histogram[l_bin_id32+1], 1, return_prev=True)
pos = T.atomic_add(s_histogram[l_bin_id32 + 1], 1, return_prev=True)
index[bx, pos] = input_idx
elif l_bin_id32 == l_threshold_bin_id and l_new_topk > 0:
pos = T.atomic_add(s_num_input[0], 1, return_prev=True)
@@ -107,7 +105,7 @@ Turning dense MLA into sparse MLA requires surprisingly few changes - essentiall
# Dense MLA: iterate over full sequence
loop_range = T.ceildiv(seqlen_kv, block_N)
for k in T.Pipelined(loop_range, num_stages=2):
T.copy(KV[bid, k * block_N:(k + 1) * block_N, cur_kv_head, :], KV_shared)
T.copy(KV[bid, k * block_N : (k + 1) * block_N, cur_kv_head, :], KV_shared)
# ... compute attention over this block
```
@@ -178,8 +176,8 @@ The backward pass consists of three main stages:
```python
for k in T.Pipelined(T.ceildiv(D, block_ND), num_stages=num_stages):
T.copy(O[bz, by * block_ND:(by + 1) * block_ND, bx, k * block_ND:(k + 1) * block_ND], o)
T.copy(dO[bz, by * block_ND:(by + 1) * block_ND, bx, k * block_ND:(k + 1) * block_ND], do)
T.copy(O[bz, by * block_ND : (by + 1) * block_ND, bx, k * block_ND : (k + 1) * block_ND], o)
T.copy(dO[bz, by * block_ND : (by + 1) * block_ND, bx, k * block_ND : (k + 1) * block_ND], do)
for i, j in T.Parallel(block_ND, block_ND):
acc[i, j] += o[i, j] * do[i, j]
T.reduce_sum(acc, delta, 1)
@@ -212,8 +210,7 @@ The key gradient computations are:
```python
# Atomically update dKV at selected indices
for bi_i, d_i in T.Parallel(BI // split_store, D // 4):
T.atomic_addx4(dKV[by, Indices[by, s_i, bz, i_i * BI + bi_i + s * (BI // split_store)], bz, d_i * 4],
acc_dkv_shared[bi_i, d_i * 4])
T.atomic_addx4(dKV[by, Indices[by, s_i, bz, i_i * BI + bi_i + s * (BI // split_store)], bz, d_i * 4], acc_dkv_shared[bi_i, d_i * 4])
```
**Performance**: The sparse MLA backward achieves excellent performance:
+1 -4
View File
@@ -18,10 +18,7 @@ def dequant_matmul(
Ct_local = T.alloc_fragment((block_N, block_M), accum_dtype)
T.clear(Ct_local)
for k in T.Pipelined(
T.ceildiv(K, block_K),
num_stages=num_stages
):
for k in T.Pipelined(T.ceildiv(K, block_K), num_stages=num_stages):
T.copy(A[by * block_M, k * block_K], A_shared)
T.copy(B[bx * block_N, k * block_K // num_elems_per_byte], B_shared)
T.copy(B_shared, B_local)
+17 -17
View File
@@ -46,8 +46,8 @@ fp8 with a per-block f32 scale. Groups `N` K tokens into
**Interface**:
```python
blocked_k, blocked_k_scale = fp8_native_block_mean_pooling_interface(
k, # [N, D] fp8
k_scale, # [N] f32 — per-token scale from indexer_k_quant_and_cache
k, # [N, D] fp8
k_scale, # [N] f32 — per-token scale from indexer_k_quant_and_cache
k_block_size,
)
# blocked_k: [num_blocks, D] fp8
@@ -72,12 +72,12 @@ blocked_k, blocked_k_scale = fp8_native_block_mean_pooling_interface(
**Interface**:
```python
block_k_score = pool_mqa_attn_return_logits_fp8_interface(
q_fp8, # [M, H, D] fp8
blocked_kv_fp8, # [Nb, D] fp8 (from step 1.1)
blocked_kv_scale, # [Nb] f32 (from step 1.1)
weights_f32, # [M, H] f32
cu_seqlen_blocked_ks, # [M] int32 — per-query start in pool-block coords
cu_seqlen_blocked_ke, # [M] int32 — per-query end in pool-block coords
q_fp8, # [M, H, D] fp8
blocked_kv_fp8, # [Nb, D] fp8 (from step 1.1)
blocked_kv_scale, # [Nb] f32 (from step 1.1)
weights_f32, # [M, H] f32
cu_seqlen_blocked_ks, # [M] int32 — per-query start in pool-block coords
cu_seqlen_blocked_ke, # [M] int32 — per-query end in pool-block coords
)
# block_k_score: [M, Nb] f32
```
@@ -104,7 +104,7 @@ zero-init value.
**Interface**:
```python
clean_and_maintain_logits_interface(
logits, # [M, Nb] f32 — stage-1 output; modified in place
logits, # [M, Nb] f32 — stage-1 output; modified in place
cu_seqlen_ks, # [M] int32 — per-row start (inclusive)
cu_seqlen_ke, # [M] int32 — per-row end (exclusive)
)
@@ -131,14 +131,14 @@ auto-dispatched by the factory:
**Interface**:
```python
block_sparse_logits = fp8_native_block_sparse_mqa_attn_return_logits_interface(
q, # [M, H, D] fp8
k, # [N, D] fp8
k_scale, # [N] f32
topk_block_index, # [M, block_topk] int64 — from torch.topk over stage-1 scores
kv_block_size, # == k_block_size
weights, # [M, H] f32
cu_seqlen_ks, # [M] int32 — per-query K start (absolute, in raw tokens)
cu_seqlen_ke, # [M] int32 — per-query K end
q, # [M, H, D] fp8
k, # [N, D] fp8
k_scale, # [N] f32
topk_block_index, # [M, block_topk] int64 — from torch.topk over stage-1 scores
kv_block_size, # == k_block_size
weights, # [M, H] f32
cu_seqlen_ks, # [M] int32 — per-query K start (absolute, in raw tokens)
cu_seqlen_ke, # [M] int32 — per-query K end
)
# block_sparse_logits: [M, block_topk * kv_block_size] f32
```
+2 -7
View File
@@ -34,7 +34,6 @@ def flash_attention(
scores_sum = T.alloc_fragment([block_M], accum_dtype)
logsum = T.alloc_fragment([block_M], accum_dtype)
# Copy a block of Q from global memory to Q_shared
T.copy(Q[bz, bx * block_M : (bx + 1) * block_M, by, :], Q_shared)
@@ -42,9 +41,7 @@ def flash_attention(
T.fill(acc_o, 0)
T.fill(logsum, 0)
T.fill(scores_max, -T.infinity(accum_dtype))
loop_range = (
T.ceildiv((bx + 1) * block_M, block_N) if is_causal else T.ceildiv(seq_len, block_N)
)
loop_range = T.ceildiv((bx + 1) * block_M, block_N) if is_causal else T.ceildiv(seq_len, block_N)
# Pipeline the loop to overlap copies/gemm stages
for k in T.Pipelined(loop_range, num_stages=num_stages):
@@ -53,9 +50,7 @@ def flash_attention(
if is_causal:
for i, j in T.Parallel(block_M, block_N):
acc_s[i, j] = T.if_then_else(
bx * block_M + i >= k * block_N + j, 0, -T.infinity(acc_s.dtype)
)
acc_s[i, j] = T.if_then_else(bx * block_M + i >= k * block_N + j, 0, -T.infinity(acc_s.dtype))
else:
T.clear(acc_s)
+22 -26
View File
@@ -53,6 +53,7 @@ import tilelang
from tilelang import Profiler
import tilelang.language as T
def matmul(M, N, K, block_M, block_N, block_K, dtype=T.float16, accum_dtype=T.float):
@T.prim_func
def main(
@@ -62,7 +63,6 @@ def matmul(M, N, K, block_M, block_N, block_K, dtype=T.float16, accum_dtype=T.fl
):
# Define a grid with enough blocks to cover M×N
with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M), threads=128) as (bx, by):
# Allocate shared memory for the current tile of A and B
A_shared = T.alloc_shared((block_M, block_K), dtype)
B_shared = T.alloc_shared((block_K, block_N), dtype)
@@ -134,6 +134,7 @@ artifact = tilelang.lower(func)
profiler = Profiler(artifact.rt_mod, artifact.params, result_idx=[2])
import torch
a = torch.randn(1024, 1024).cuda().half()
b = torch.randn(1024, 1024).cuda().half()
@@ -172,10 +173,12 @@ Below is a more advanced snippet that showcases how to apply memory layouts, ena
```python
import tilelang.language as T
# `make_mma_swizzle_layout` is a python-defined layout function
# that helps align data for MMA (Matrix Multiply-Accumulate) operations.
from tilelang.cuda.intrinsics import make_mma_swizzle_layout as make_swizzle_layout
def matmul(M, N, K, block_M, block_N, block_K, dtype=T.float16, accum_dtype=T.float):
@T.prim_func
def main(
@@ -187,13 +190,15 @@ def matmul(M, N, K, block_M, block_N, block_K, dtype=T.float16, accum_dtype=T.fl
# Allocate shared and local fragments
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), accum_dtype)
C_local = T.alloc_fragment((block_M, block_N), accum_dtype)
# Annotate memory layout
T.annotate_layout({
A_shared: make_swizzle_layout(A_shared),
B_shared: make_swizzle_layout(B_shared),
})
T.annotate_layout(
{
A_shared: make_swizzle_layout(A_shared),
B_shared: make_swizzle_layout(B_shared),
}
)
# Enable swizzle-based rasterization for better L2 locality
T.use_swizzle(panel_size=10, enable=True)
@@ -329,12 +334,11 @@ def tl_matmul(
@T.prim_func
def main(
A: T.Tensor(A_shape, in_dtype),
B: T.Tensor(B_shape, in_dtype),
C: T.Tensor((M, N), out_dtype),
A: T.Tensor(A_shape, in_dtype),
B: T.Tensor(B_shape, in_dtype),
C: T.Tensor((M, N), out_dtype),
):
with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M), threads=threads) as (bx, by):
A_shared = T.alloc_shared(A_shared_shape, in_dtype, scope=shared_scope)
B_shared = T.alloc_shared(B_shared_shape, in_dtype, scope=shared_scope)
C_shared = T.alloc_shared(C_shared_shape, out_dtype, scope=shared_scope)
@@ -342,10 +346,12 @@ def tl_matmul(
B_local = T.alloc_local((warp_cols * local_size_b), in_dtype)
C_local = T.alloc_local((warp_rows * warp_cols * local_size_c), accum_dtype)
T.annotate_layout({
A_shared: make_swizzle_layout(A_shared),
B_shared: make_swizzle_layout(B_shared),
})
T.annotate_layout(
{
A_shared: make_swizzle_layout(A_shared),
B_shared: make_swizzle_layout(B_shared),
}
)
# Improve L2 Cache
T.use_swizzle(panel_size=10)
@@ -353,7 +359,6 @@ def tl_matmul(
T.clear(C_local)
for ko in T.Pipelined((K // block_K), num_stages=stage):
# Load A into shared memory
for i, k in T.Parallel(block_M, block_K):
A_shared[i, k] = A[by * block_M + i, ko * block_K + k]
@@ -363,20 +368,11 @@ def tl_matmul(
B_shared[j, k] = B[bx * block_N + j, ko * block_K + k]
for ki in T.serial(0, (block_K // micro_size_k)):
# Load A into fragment
mma_emitter.ldmatrix_a(
A_local,
A_shared,
ki
)
mma_emitter.ldmatrix_a(A_local, A_shared, ki)
# Load B into fragment
mma_emitter.ldmatrix_b(
B_local,
B_shared,
ki
)
mma_emitter.ldmatrix_b(B_local, B_shared, ki)
# Perform Matrix Multiplication
mma_emitter.mma(A_local, B_local, C_local)
+20 -15
View File
@@ -17,8 +17,8 @@ matching `mbarrier_wait_parity(...)` automatically after TCGEN5MMA issue.
TCGEN5MMA is asynchronous and requires explicit synchronization:
```python
mbar = T.alloc_barrier(1) # expect-arrive-count = 1
T.tcgen05_gemm(A_shared, B_shared, C_tmem, trans_A, trans_B, mbar=mbar, clear_accum=k==0)
T.mbarrier_wait_parity(mbar, k%2) # Manual phase calculation required
T.tcgen05_gemm(A_shared, B_shared, C_tmem, trans_A, trans_B, mbar=mbar, clear_accum=k == 0)
T.mbarrier_wait_parity(mbar, k % 2) # Manual phase calculation required
```
TileLang now has a conservative `InjectTcgen05Fence` pass on TCGEN05-capable targets that can
@@ -60,6 +60,7 @@ import torch
import tilelang
import tilelang.language as T
@T.prim_func
def main(
A: T.Tensor((M, K), T.bfloat16),
@@ -70,11 +71,11 @@ def main(
# 1. Allocate memory buffers
A_shared = T.alloc_shared((block_M, block_K), T.bfloat16) # A matrix shared memory
B_shared = T.alloc_shared((block_N, block_K), T.bfloat16) # B matrix shared memory
C_tmem = T.alloc_tmem([block_M, block_N], T.float) # TCGEN5MMA output to Tensor Memory
mbar = T.alloc_barrier(1) # mbarrier synchronization primitive
C_tmem = T.alloc_tmem([block_M, block_N], T.float) # TCGEN5MMA output to Tensor Memory
mbar = T.alloc_barrier(1) # mbarrier synchronization primitive
C_local = T.alloc_fragment((block_M, block_N), T.float) # Register storage
C_shared = T.alloc_shared((block_M, block_N), T.bfloat16) # Output shared memory
C_local = T.alloc_fragment((block_M, block_N), T.float) # Register storage
C_shared = T.alloc_shared((block_M, block_N), T.bfloat16) # Output shared memory
# 2. Main computation loop
for k in T.Pipelined(T.ceildiv(K, block_K), num_stages=1):
@@ -83,15 +84,14 @@ def main(
T.copy(B[bx * block_N, k * block_K], B_shared)
# TCGEN5MMA computation: asynchronous launch, output to Tensor Memory
T.tcgen05_gemm(A_shared, B_shared, C_tmem, trans_A=False, trans_B=True,
mbar=mbar, clear_accum=k==0)
T.tcgen05_gemm(A_shared, B_shared, C_tmem, trans_A=False, trans_B=True, mbar=mbar, clear_accum=k == 0)
# Critical: wait for TCGEN5MMA completion
T.mbarrier_wait_parity(mbar, k%2)
T.mbarrier_wait_parity(mbar, k % 2)
# 3. Output processing (only subset of threads)
T.copy(C_tmem, C_local) # Tensor Memory → registers
T.copy(C_local, C_shared) # registers → shared memory
T.copy(C_tmem, C_local) # Tensor Memory → registers
T.copy(C_local, C_shared) # registers → shared memory
# 4. Write back to global memory
T.copy(C_shared, C[by * block_M, bx * block_N])
@@ -105,9 +105,14 @@ M, N, K = 4096, 4096, 8192
block_M, block_N, block_K = 128, 256, 128
# Compile kernel
jit_kernel = tilelang.compile(func, out_idx=[2], target="cuda", pass_configs={
tilelang.PassConfigKey.TL_DISABLE_WARP_SPECIALIZED: True, # Required
})
jit_kernel = tilelang.compile(
func,
out_idx=[2],
target="cuda",
pass_configs={
tilelang.PassConfigKey.TL_DISABLE_WARP_SPECIALIZED: True, # Required
},
)
# Run test
a = torch.randn(M, K, device="cuda", dtype=torch.bfloat16)
@@ -122,5 +127,5 @@ torch.testing.assert_close(c, ref_c, rtol=1e-2, atol=1e-2)
profiler = jit_kernel.get_profiler()
latency = profiler.do_bench()
print(f"Latency: {latency} ms")
print(f"Performance: {2 * M * N * K / (latency/1e3) / 1e12:.2f} TFLOPS")
print(f"Performance: {2 * M * N * K / (latency / 1e3) / 1e12:.2f} TFLOPS")
```
@@ -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])
+38 -5
View File
@@ -21,6 +21,15 @@ OLD_PYTHON = os.environ.get("OLD_PYTHON", "./old/bin/python")
NEW_PYTHON = os.environ.get("NEW_PYTHON", "./new/bin/python")
OUT_MD = os.environ.get("PERF_REGRESSION_MD", "regression_result.md")
OUT_PNG = os.environ.get("PERF_REGRESSION_PNG", "regression_result.png")
# Fail the run only when at least MIN_COUNT benchmarks drop below THRESHOLD.
# Calibrated on the 33 bot reports posted to PRs up to #3176 (1901 samples):
# 20 reports are clean, 6 (all on PR #2464) show a real regression with 14-23
# benchmarks collapsing at once, and 7 carry exactly one outlier as low as 0.044
# while every other benchmark sits at 1.000. Requiring two simultaneous drops
# separates those cases exactly -- 6/6 real regressions caught, 0 false alarms --
# for any threshold in 0.80..0.95.
THRESHOLD = float(os.environ.get("PERF_REGRESSION_THRESHOLD", "0.90"))
MIN_COUNT = int(os.environ.get("PERF_REGRESSION_MIN_COUNT", "2"))
_RESULTS_JSON_PREFIX = "__TILELANG_PERF_RESULTS_JSON__="
@@ -226,11 +235,35 @@ if not table:
table.sort(key=lambda x: x[-1])
headers = ["File", "Original Latency", "Current Latency", "Speedup"]
with open(OUT_MD, "w") as f:
f.write(tabulate(table, headers=headers, tablefmt="github", stralign="left", numalign="decimal"))
f.write("\n")
df = pd.DataFrame(table, columns=headers)
df = df.sort_values("Speedup", ascending=False).reset_index(drop=True)
draw(df)
# Speedup is old/new latency, so anything below the threshold is a slowdown.
slow = [row for row in table if row[-1] < THRESHOLD]
regressed = slow if len(slow) >= MIN_COUNT else []
if regressed:
summary = f"**Regression: {len(regressed)}/{len(table)} benchmarks below {THRESHOLD:g}x**\n\n"
summary += "".join(f"- `{row[0]}`: {row[-1]:.3f}x\n" for row in regressed)
elif slow:
summary = (
f"No regression. {len(slow)}/{len(table)} benchmarks landed below {THRESHOLD:g}x, "
f"under the {MIN_COUNT}-benchmark bar that separates a real regression from a single noisy run: "
+ ", ".join(f"`{row[0]}` {row[-1]:.3f}x" for row in slow)
+ "\n"
)
else:
summary = f"No regression: all {len(table)} benchmarks kept speedup >= {THRESHOLD:g}.\n"
marked = [[f"{row[0]} :warning:" if row[-1] < THRESHOLD else row[0], *row[1:]] for row in table]
with open(OUT_MD, "w") as f:
f.write(summary)
f.write("\n")
f.write(tabulate(marked, headers=headers, tablefmt="github", stralign="left", numalign="decimal"))
f.write("\n")
print(summary)
if regressed:
# Exit last so the markdown table and the plot are still produced for the report.
exit(1)
+1 -1
View File
@@ -15,7 +15,7 @@ namespace tl {
using namespace tirx;
TVM_REGISTER_PASS_CONFIG_OPTION(kDisableWarpSpecialized, Bool);
TVM_REGISTER_PASS_CONFIG_OPTION(kCudaAutoSchedule, ffi::String);
TVM_REGISTER_PASS_CONFIG_OPTION(kEnableAutoWarpSpecialization, ffi::String);
TVM_REGISTER_PASS_CONFIG_OPTION(kDisableTMALower, Bool);
TVM_REGISTER_PASS_CONFIG_OPTION(kPtxasRegisterUsageLevel, Integer);
TVM_REGISTER_PASS_CONFIG_OPTION(kDisableVectorize256, Bool);
+2 -3
View File
@@ -35,9 +35,8 @@ static constexpr const char *kHasTMA = "tl.has_tma";
// because they are part of the Python PassContext interface.
static constexpr const char *kDisableWarpSpecialized =
"tl.disable_warp_specialized";
// Keep the CUDA scheduler name separate from Ascend's existing boolean
// tl.enable_auto_schedule option; pass-config types are process-global.
static constexpr const char *kCudaAutoSchedule = "tl.cuda_auto_schedule";
static constexpr const char *kEnableAutoWarpSpecialization =
"tl.enable_auto_warp_specialization";
static constexpr const char *kDisableTMALower = "tl.disable_tma_lower";
static constexpr const char *kPtxasRegisterUsageLevel =
"tl.ptxas_register_usage_level";
+59 -59
View File
@@ -510,18 +510,18 @@ typedef enum CUstreamWaitValue_flags_enum {
0x3, /**< Wait until ~(*addr | value) != 0. Support for this operation can
be queried with ::cuDeviceGetAttribute() and
::CU_DEVICE_ATTRIBUTE_CAN_USE_STREAM_WAIT_VALUE_NOR.*/
CU_STREAM_WAIT_VALUE_FLUSH =
1 << 30 /**< Follow the wait operation with a flush of outstanding remote
writes. This means that, if a remote write operation is
guaranteed to have reached the device before the wait can be
satisfied, that write is guaranteed to be visible to downstream
device work. The device is permitted to reorder remote writes
internally. For example, this flag would be required if two
remote writes arrive in a defined order, the wait is satisfied
by the second write, and downstream work needs to observe the
first write. Support for this operation is restricted to
selected platforms and can be queried with
::CU_DEVICE_ATTRIBUTE_CAN_FLUSH_REMOTE_WRITES.*/
CU_STREAM_WAIT_VALUE_FLUSH = 1
<< 30 /**< Follow the wait operation with a flush of outstanding remote
writes. This means that, if a remote write operation is
guaranteed to have reached the device before the wait can be
satisfied, that write is guaranteed to be visible to downstream
device work. The device is permitted to reorder remote writes
internally. For example, this flag would be required if two
remote writes arrive in a defined order, the wait is satisfied
by the second write, and downstream work needs to observe the
first write. Support for this operation is restricted to
selected platforms and can be queried with
::CU_DEVICE_ATTRIBUTE_CAN_FLUSH_REMOTE_WRITES.*/
} CUstreamWaitValue_flags;
/**
@@ -1801,8 +1801,8 @@ typedef enum CUjit_target_enum {
CU_TARGET_COMPUTE_90 = 90, /**< Compute device class 9.0.*/
/**< Compute device class 9.0. with accelerated features.*/
CU_TARGET_COMPUTE_90A =
CU_COMPUTE_ACCELERATED_TARGET_BASE + CU_TARGET_COMPUTE_90,
CU_TARGET_COMPUTE_90A = CU_COMPUTE_ACCELERATED_TARGET_BASE +
CU_TARGET_COMPUTE_90,
} CUjit_target;
/**
@@ -2240,7 +2240,7 @@ typedef struct CUgraphEdgeData_st {
cases means the entirety of the
downstream node is dependent on the
upstream work. <br> Currently no
node types define non-zero ports. Accordingly,
node types define non-zero ports. Accordingly,
this field must be set to zero. */
unsigned char type; /**< This should be populated with a value from
::CUgraphDependencyType. (It is typed as char due to
@@ -2659,10 +2659,10 @@ typedef CUstreamAttrValue_v1 CUstreamAttrValue;
typedef enum CUdriverProcAddress_flags_enum {
CU_GET_PROC_ADDRESS_DEFAULT =
0, /**< Default search mode for driver symbols. */
CU_GET_PROC_ADDRESS_LEGACY_STREAM =
1 << 0, /**< Search for legacy versions of driver symbols. */
CU_GET_PROC_ADDRESS_PER_THREAD_DEFAULT_STREAM =
1 << 1 /**< Search for per-thread versions of driver symbols. */
CU_GET_PROC_ADDRESS_LEGACY_STREAM = 1
<< 0, /**< Search for legacy versions of driver symbols. */
CU_GET_PROC_ADDRESS_PER_THREAD_DEFAULT_STREAM = 1
<< 1 /**< Search for per-thread versions of driver symbols. */
} CUdriverProcAddress_flags;
/**
@@ -5040,13 +5040,13 @@ typedef struct CUgraphNodeParams_st {
* Bitmasks for ::CU_DEVICE_ATTRIBUTE_GPU_DIRECT_RDMA_FLUSH_WRITES_OPTIONS
*/
typedef enum CUflushGPUDirectRDMAWritesOptions_enum {
CU_FLUSH_GPU_DIRECT_RDMA_WRITES_OPTION_HOST =
1 << 0, /**< ::cuFlushGPUDirectRDMAWrites() and its CUDA Runtime API
counterpart are supported on the device. */
CU_FLUSH_GPU_DIRECT_RDMA_WRITES_OPTION_MEMOPS =
1 << 1 /**< The ::CU_STREAM_WAIT_VALUE_FLUSH flag and the
::CU_STREAM_MEM_OP_FLUSH_REMOTE_WRITES MemOp are supported on
the device. */
CU_FLUSH_GPU_DIRECT_RDMA_WRITES_OPTION_HOST = 1
<< 0, /**< ::cuFlushGPUDirectRDMAWrites() and its CUDA Runtime API
counterpart are supported on the device. */
CU_FLUSH_GPU_DIRECT_RDMA_WRITES_OPTION_MEMOPS = 1
<< 1 /**< The ::CU_STREAM_WAIT_VALUE_FLUSH flag and the
::CU_STREAM_MEM_OP_FLUSH_REMOTE_WRITES MemOp are supported on
the device. */
} CUflushGPUDirectRDMAWritesOptions;
/**
@@ -5089,41 +5089,41 @@ typedef enum CUflushGPUDirectRDMAWritesTarget_enum {
* The additional write options for ::cuGraphDebugDotPrint
*/
typedef enum CUgraphDebugDot_flags_enum {
CU_GRAPH_DEBUG_DOT_FLAGS_VERBOSE =
1 << 0, /**< Output all debug data as if every debug flag is enabled */
CU_GRAPH_DEBUG_DOT_FLAGS_RUNTIME_TYPES =
1 << 1, /**< Use CUDA Runtime structures for output */
CU_GRAPH_DEBUG_DOT_FLAGS_KERNEL_NODE_PARAMS =
1 << 2, /**< Adds CUDA_KERNEL_NODE_PARAMS values to output */
CU_GRAPH_DEBUG_DOT_FLAGS_MEMCPY_NODE_PARAMS =
1 << 3, /**< Adds CUDA_MEMCPY3D values to output */
CU_GRAPH_DEBUG_DOT_FLAGS_MEMSET_NODE_PARAMS =
1 << 4, /**< Adds CUDA_MEMSET_NODE_PARAMS values to output */
CU_GRAPH_DEBUG_DOT_FLAGS_HOST_NODE_PARAMS =
1 << 5, /**< Adds CUDA_HOST_NODE_PARAMS values to output */
CU_GRAPH_DEBUG_DOT_FLAGS_EVENT_NODE_PARAMS =
1 << 6, /**< Adds CUevent handle from record and wait nodes to output */
CU_GRAPH_DEBUG_DOT_FLAGS_EXT_SEMAS_SIGNAL_NODE_PARAMS =
1 << 7, /**< Adds CUDA_EXT_SEM_SIGNAL_NODE_PARAMS values to output */
CU_GRAPH_DEBUG_DOT_FLAGS_EXT_SEMAS_WAIT_NODE_PARAMS =
1 << 8, /**< Adds CUDA_EXT_SEM_WAIT_NODE_PARAMS values to output */
CU_GRAPH_DEBUG_DOT_FLAGS_KERNEL_NODE_ATTRIBUTES =
1 << 9, /**< Adds CUkernelNodeAttrValue values to output */
CU_GRAPH_DEBUG_DOT_FLAGS_HANDLES =
1 << 10, /**< Adds node handles and every kernel function handle to output
*/
CU_GRAPH_DEBUG_DOT_FLAGS_MEM_ALLOC_NODE_PARAMS =
1 << 11, /**< Adds memory alloc node parameters to output */
CU_GRAPH_DEBUG_DOT_FLAGS_MEM_FREE_NODE_PARAMS =
1 << 12, /**< Adds memory free node parameters to output */
CU_GRAPH_DEBUG_DOT_FLAGS_BATCH_MEM_OP_NODE_PARAMS =
1 << 13 /**< Adds batch mem op node parameters to output */
CU_GRAPH_DEBUG_DOT_FLAGS_VERBOSE = 1
<< 0, /**< Output all debug data as if every debug flag is enabled */
CU_GRAPH_DEBUG_DOT_FLAGS_RUNTIME_TYPES = 1
<< 1, /**< Use CUDA Runtime structures for output */
CU_GRAPH_DEBUG_DOT_FLAGS_KERNEL_NODE_PARAMS = 1
<< 2, /**< Adds CUDA_KERNEL_NODE_PARAMS values to output */
CU_GRAPH_DEBUG_DOT_FLAGS_MEMCPY_NODE_PARAMS = 1
<< 3, /**< Adds CUDA_MEMCPY3D values to output */
CU_GRAPH_DEBUG_DOT_FLAGS_MEMSET_NODE_PARAMS = 1
<< 4, /**< Adds CUDA_MEMSET_NODE_PARAMS values to output */
CU_GRAPH_DEBUG_DOT_FLAGS_HOST_NODE_PARAMS = 1
<< 5, /**< Adds CUDA_HOST_NODE_PARAMS values to output */
CU_GRAPH_DEBUG_DOT_FLAGS_EVENT_NODE_PARAMS = 1
<< 6, /**< Adds CUevent handle from record and wait nodes to output */
CU_GRAPH_DEBUG_DOT_FLAGS_EXT_SEMAS_SIGNAL_NODE_PARAMS = 1
<< 7, /**< Adds CUDA_EXT_SEM_SIGNAL_NODE_PARAMS values to output */
CU_GRAPH_DEBUG_DOT_FLAGS_EXT_SEMAS_WAIT_NODE_PARAMS = 1
<< 8, /**< Adds CUDA_EXT_SEM_WAIT_NODE_PARAMS values to output */
CU_GRAPH_DEBUG_DOT_FLAGS_KERNEL_NODE_ATTRIBUTES = 1
<< 9, /**< Adds CUkernelNodeAttrValue values to output */
CU_GRAPH_DEBUG_DOT_FLAGS_HANDLES = 1
<< 10, /**< Adds node handles and every kernel function handle to output
*/
CU_GRAPH_DEBUG_DOT_FLAGS_MEM_ALLOC_NODE_PARAMS = 1
<< 11, /**< Adds memory alloc node parameters to output */
CU_GRAPH_DEBUG_DOT_FLAGS_MEM_FREE_NODE_PARAMS = 1
<< 12, /**< Adds memory free node parameters to output */
CU_GRAPH_DEBUG_DOT_FLAGS_BATCH_MEM_OP_NODE_PARAMS = 1
<< 13 /**< Adds batch mem op node parameters to output */
,
CU_GRAPH_DEBUG_DOT_FLAGS_EXTRA_TOPO_INFO =
1 << 14 /**< Adds edge numbering information */
CU_GRAPH_DEBUG_DOT_FLAGS_EXTRA_TOPO_INFO = 1
<< 14 /**< Adds edge numbering information */
,
CU_GRAPH_DEBUG_DOT_FLAGS_CONDITIONAL_NODE_PARAMS =
1 << 15 /**< Adds conditional node parameters to output */
CU_GRAPH_DEBUG_DOT_FLAGS_CONDITIONAL_NODE_PARAMS = 1
<< 15 /**< Adds conditional node parameters to output */
} CUgraphDebugDot_flags;
/**
+27 -24
View File
@@ -2,18 +2,17 @@
* \file auto_schedule.cc
* \brief Generic entrypoint for automatic warp-specialization schedulers.
*
* The pass config "tl.cuda_auto_schedule" names the scheduler to run
* (see SchedulerRegistry; currently "role_based"); when unset the pass
* is a no-op. For each eligible kernel — a tilelang_root block with a
* known threadIdx.x extent, no existing schedule, and no manual warp
* specialization — the entrypoint
* gives every schedulable statement a stable "tl.ws_op_id" marker and
* checks the schedulability contract (preprocess_ir), then asks the
* The pass config "tl.enable_auto_warp_specialization" names the scheduler to
* run (see SchedulerRegistry; currently "role_based"); when unset the pass is a
* no-op. For each eligible kernel — a tilelang_root block with a known
* threadIdx.x extent, no existing schedule, and no manual warp specialization —
* the entrypoint gives every schedulable statement a stable "tl.ws_op_id"
* marker and checks the schedulability contract (preprocess_ir), then asks the
* scheduler for a typed WSSchedule, which MaterializeWSSchedule then
* lowers. The kernel body itself only gains the id markers; any kernel
* preprocessing or the scheduler declines is left byte-for-byte
* unchanged, with the reason emitted as an on-site warning (auto
* scheduling is opt-in, so the user expects it to fire).
* unchanged, with the reason emitted as an on-site warning (auto warp
* specialization is opt-in, so the user expects it to fire).
*
* TODO: verify dependence coverage of USER-PROVIDED schedules with a real
* dependence analysis (synthesized schedules cover exactly the cross-role
@@ -43,10 +42,12 @@ namespace tl {
using namespace tirx;
using namespace tirx::transform;
using namespace cuda;
namespace {
// Available schedulers, by the name passed in tl.cuda_auto_schedule.
// Available schedulers, by the name passed in
// tl.enable_auto_warp_specialization.
const std::map<std::string, SchedulerFn> &SchedulerRegistry() {
static const std::map<std::string, SchedulerFn> registry = {
{"role_based", RoleBasedSchedule},
@@ -61,15 +62,15 @@ SchedulerFn FindScheduler(const ffi::String &name) {
std::string known;
for (const auto &[known_name, fn] : registry)
known += (known.empty() ? "" : ", ") + known_name;
LOG(FATAL) << "unknown auto-schedule scheduler '" << name
LOG(FATAL) << "unknown auto-warp-specialization scheduler '" << name
<< "'; available: " << known;
}
return it->second;
}
class AutoScheduleRewriter : public StmtMutator {
class AutoWarpSpecializationRewriter : public StmtMutator {
public:
AutoScheduleRewriter(SchedulerFn scheduler, Target target)
AutoWarpSpecializationRewriter(SchedulerFn scheduler, Target target)
: scheduler_(scheduler), target_(std::move(target)) {}
private:
@@ -120,12 +121,13 @@ private:
int original_threads_{0};
};
PrimFunc AutoScheduleImpl(PrimFunc func, SchedulerFn scheduler) {
PrimFunc AutoWarpSpecializationImpl(PrimFunc func, SchedulerFn scheduler) {
auto target = func->GetAttr<Target>(tvm::attr::kTarget);
ICHECK(target.defined()) << "AutoSchedule: PrimFunc has no bound target "
"(BindTarget must run before AutoSchedule)";
ICHECK(target.defined())
<< "AutoWarpSpecialization: PrimFunc has no bound target "
"(BindTarget must run before AutoWarpSpecialization)";
AutoScheduleRewriter rewriter(scheduler, target.value());
AutoWarpSpecializationRewriter rewriter(scheduler, target.value());
Stmt body = rewriter(func->body);
if (body.same_as(func->body))
return func;
@@ -135,21 +137,22 @@ PrimFunc AutoScheduleImpl(PrimFunc func, SchedulerFn scheduler) {
} // namespace
tvm::transform::Pass CudaAutoSchedule() {
tvm::transform::Pass AutoWarpSpecialization() {
auto pass_func = [](PrimFunc func, const IRModule &, const PassContext &ctx) {
auto scheduler_name =
ctx->GetConfig(kCudaAutoSchedule, ffi::Optional<ffi::String>());
auto scheduler_name = ctx->GetConfig(kEnableAutoWarpSpecialization,
ffi::Optional<ffi::String>());
if (!scheduler_name.has_value())
return func;
return AutoScheduleImpl(std::move(func),
FindScheduler(scheduler_name.value()));
return AutoWarpSpecializationImpl(std::move(func),
FindScheduler(scheduler_name.value()));
};
return CreatePrimFuncPass(pass_func, 0, "tl.cuda.AutoSchedule", {});
return CreatePrimFuncPass(pass_func, 0, "tl.AutoWarpSpecialization", {});
}
TVM_FFI_STATIC_INIT_BLOCK() {
namespace refl = tvm::ffi::reflection;
refl::GlobalDef().def("tl.cuda.transform.AutoSchedule", CudaAutoSchedule);
refl::GlobalDef().def("tl.cuda.transform.AutoWarpSpecialization",
AutoWarpSpecialization);
}
} // namespace tl
+5 -2
View File
@@ -1,6 +1,7 @@
/*!
* \file common.h
* \brief Shared declarations of the AutoSchedule entrypoint and schedulers.
* \brief Shared declarations of the AutoWarpSpecialization entrypoint and
* schedulers.
*
* A scheduler receives one eligible kernel — the root block and its body
* normalized so every schedulable statement carries a "tl.ws_op_id" marker —
@@ -19,6 +20,7 @@
namespace tvm {
namespace tl {
namespace cuda {
using SchedulerFn = ffi::Optional<WSSchedule> (*)(const tirx::SBlock &block,
const tirx::Stmt &body,
@@ -32,9 +34,10 @@ inline ffi::String ExtractOpId(const ffi::Any &value) {
return string.value();
if (const auto *imm = value.as<tirx::StringImmNode>())
return imm->value;
TVM_FFI_THROW(ValueError) << "AutoSchedule op id must be a string";
TVM_FFI_THROW(ValueError) << "AutoWarpSpecialization op id must be a string";
return "";
}
} // namespace cuda
} // namespace tl
} // namespace tvm
@@ -23,6 +23,7 @@
namespace tvm {
namespace tl {
namespace cuda {
using namespace tirx;
using ffi::Array;
@@ -459,5 +460,6 @@ private:
}
};
} // namespace cuda
} // namespace tl
} // namespace tvm
@@ -37,6 +37,7 @@
namespace tvm {
namespace tl {
namespace cuda {
using namespace tirx;
using namespace ffi;
@@ -62,23 +63,26 @@ bool CheckCalls(const Stmt &body) {
if (call == nullptr || !ok)
return;
if (call->op.same_as(tl::loop_break())) {
LOG(WARNING) << "AutoSchedule skipped: cannot schedule loop_break";
LOG(WARNING)
<< "AutoWarpSpecialization skipped: cannot schedule loop_break";
ok = false;
return;
}
if (IsBarrierOrTmaControlCall(call)) {
LOG(WARNING) << "AutoSchedule skipped: cannot schedule hand-written "
"synchronization or TMA control ('"
<< call->op << "'): block-wide syncs cannot be duplicated "
<< "into roles, and hand-managed barrier protocols are "
<< "invisible to the schedule";
LOG(WARNING)
<< "AutoWarpSpecialization skipped: cannot schedule hand-written "
"synchronization or TMA control ('"
<< call->op << "'): block-wide syncs cannot be duplicated "
<< "into roles, and hand-managed barrier protocols are "
<< "invisible to the schedule";
ok = false;
return;
}
if (const auto *op = call->op.as<OpNode>()) {
if (std::string(op->name).find("atomic") != std::string::npos) {
LOG(WARNING) << "AutoSchedule skipped: cannot schedule atomic op '"
<< op->name << "'";
LOG(WARNING)
<< "AutoWarpSpecialization skipped: cannot schedule atomic op '"
<< op->name << "'";
ok = false;
}
}
@@ -101,7 +105,7 @@ bool CheckHostsNoAsync(const Stmt &stmt, const Target &target) {
return;
if (const auto *copy = tile_op.as<CopyNode>()) {
if (ClassifyCopy(copy, target) == TileStmtKind::kTmaProducer) {
LOG(WARNING) << "AutoSchedule skipped: an asynchronous "
LOG(WARNING) << "AutoWarpSpecialization skipped: an asynchronous "
"global->shared copy is nested inside a compound "
"statement; write it as its own statement so its "
"completion barrier can be wired";
@@ -111,9 +115,10 @@ bool CheckHostsNoAsync(const Stmt &stmt, const Target &target) {
}
if (auto gemm = GetGemmInfo(tile_op)) {
if (IsTmemBuffer(gemm->accumulator) || gemm->wg_wait != 0) {
LOG(WARNING) << "AutoSchedule skipped: an asynchronous gemm is "
"nested inside a compound statement; write it as "
"its own statement";
LOG(WARNING)
<< "AutoWarpSpecialization skipped: an asynchronous gemm is "
"nested inside a compound statement; write it as "
"its own statement";
ok = false;
}
}
@@ -258,5 +263,6 @@ ffi::Optional<Stmt> PreprocessIR(Stmt body, const Target &target) {
return OpIdNormalizer::Rewrite(std::move(body), target);
}
} // namespace cuda
} // namespace tl
} // namespace tvm
@@ -8,11 +8,13 @@
namespace tvm {
namespace tl {
namespace cuda {
// Give every schedulable statement a stable "tl.ws_op_id" marker in the
// forms MaterializeWSSchedule consumes; idempotent. Unschedulable
// constructs decline the kernel (warning + nullopt).
ffi::Optional<tirx::Stmt> PreprocessIR(tirx::Stmt body, const Target &target);
} // namespace cuda
} // namespace tl
} // namespace tvm
@@ -63,6 +63,7 @@
namespace tvm {
namespace tl {
namespace cuda {
using namespace tirx;
using namespace ffi;
@@ -332,8 +333,8 @@ private:
if (!ok_)
return;
if (op->else_case.defined()) {
LOG(WARNING)
<< "AutoSchedule skipped: if-else branches are not supported";
LOG(WARNING) << "AutoWarpSpecialization skipped: if-else branches are "
"not supported";
ok_ = false;
return;
}
@@ -346,8 +347,8 @@ private:
if (!ok_)
return;
auto id = op->annotations.Get(kWSOpIdKey);
ICHECK(id.has_value()) << "AutoSchedule normalized loop is missing "
<< kWSOpIdKey;
ICHECK(id.has_value())
<< "AutoWarpSpecialization normalized loop is missing " << kWSOpIdKey;
// Sequential loops are scopes; parallel / vectorized loops are one op.
if (op->kind != ForKind::kSerial && op->kind != ForKind::kUnrolled) {
MakeOp(ExtractOpId(id.value()), GetRef<For>(op));
@@ -413,22 +414,22 @@ private:
return;
}
}
LOG(FATAL) << "AutoSchedule: Evaluate carries no ws op id";
LOG(FATAL) << "AutoWarpSpecialization: Evaluate carries no ws op id";
}
// The normalizer wraps these statement forms with an id; reaching one
// bare is its bug.
void VisitStmt_(const BufferStoreNode *op) final {
LOG(FATAL) << "AutoSchedule: BufferStore carries no ws op id";
LOG(FATAL) << "AutoWarpSpecialization: BufferStore carries no ws op id";
}
void VisitStmt_(const BindNode *op) final {
LOG(FATAL) << "AutoSchedule: Bind carries no ws op id";
LOG(FATAL) << "AutoWarpSpecialization: Bind carries no ws op id";
}
void VisitStmt_(const WhileNode *op) final {
LOG(FATAL) << "AutoSchedule: while loop carries no ws op id";
LOG(FATAL) << "AutoWarpSpecialization: while loop carries no ws op id";
}
void VisitStmt_(const SBlockNode *op) final {
LOG(FATAL) << "AutoSchedule: block carries no ws op id";
LOG(FATAL) << "AutoWarpSpecialization: block carries no ws op id";
}
Target target_;
@@ -439,8 +440,8 @@ private:
};
// Plans the schedule for one normalized kernel body. Every unsupported
// shape declines with a warning: auto scheduling is opt-in, so the user
// expects it to fire.
// shape declines with a warning: auto warp specialization is opt-in, so the
// user expects it to fire.
class RoleBasedScheduler {
public:
RoleBasedScheduler(Target target, int worker_threads)
@@ -448,12 +449,14 @@ public:
Optional<WSSchedule> Run(const SBlock &block, const Stmt &body) {
if (!TargetHasBulkCopy(target_)) {
LOG(WARNING) << "AutoSchedule skipped: target has no bulk-copy (TMA) "
"support";
LOG(WARNING)
<< "AutoWarpSpecialization skipped: target has no bulk-copy (TMA) "
"support";
return std::nullopt;
}
if (worker_threads_ % 128 != 0) {
LOG(WARNING) << "AutoSchedule skipped: worker threads must be a multiple "
LOG(WARNING) << "AutoWarpSpecialization skipped: worker threads must be "
"a multiple "
"of 128 so issuer warps start a fresh warpgroup";
return std::nullopt;
}
@@ -473,8 +476,9 @@ public:
if (!BuildPipelines(block))
return std::nullopt;
if (pipelines_.empty()) {
LOG(WARNING) << "AutoSchedule skipped: no cross-role on-chip handoff to "
"pipeline";
LOG(WARNING)
<< "AutoWarpSpecialization skipped: no cross-role on-chip handoff to "
"pipeline";
return std::nullopt;
}
return Emit();
@@ -522,7 +526,7 @@ private:
if (auto mask = copy->annotations.Get("cluster_mask")) {
if (const auto *imm = mask.value().as<IntImmNode>();
imm == nullptr || imm->value != 0) {
LOG(WARNING) << "AutoSchedule skipped: op '" << op.id
LOG(WARNING) << "AutoWarpSpecialization skipped: op '" << op.id
<< "' is a cluster multicast copy";
return false;
}
@@ -539,7 +543,7 @@ private:
continue;
case TileStmtKind::kCpAsyncRaw:
// Raw cp.async carries its own thread-local completion protocol.
LOG(WARNING) << "AutoSchedule skipped: op '" << op.id
LOG(WARNING) << "AutoWarpSpecialization skipped: op '" << op.id
<< "': raw cp.async statements carry their own "
"thread-local completion protocol";
return false;
@@ -565,7 +569,7 @@ private:
// release right after the gemm would race with them.
if (auto gemm = GetGemmInfo(op.tile_op)) {
if (gemm->wg_wait != 0) {
LOG(WARNING) << "AutoSchedule skipped: op '" << op.id
LOG(WARNING) << "AutoWarpSpecialization skipped: op '" << op.id
<< "': gemm with wg_wait != 0 completes "
"asynchronously; delayed wgmma waits are not "
"supported yet";
@@ -619,9 +623,9 @@ private:
if (def->roles.Contains(role))
return true;
if (!def->roleless) {
LOG(WARNING) << "AutoSchedule skipped: '" << user << "' needs op '"
<< def->id << "' in role " << RoleName(role)
<< ", but that op is fixed to role "
LOG(WARNING) << "AutoWarpSpecialization skipped: '" << user
<< "' needs op '" << def->id << "' in role "
<< RoleName(role) << ", but that op is fixed to role "
<< RoleName(RoleOf(*def));
return false;
}
@@ -684,7 +688,8 @@ private:
for (const auto &op : ops_)
active.Add(op->roles);
if (active.Empty()) {
LOG(WARNING) << "AutoSchedule skipped: kernel has no schedulable work";
LOG(WARNING)
<< "AutoWarpSpecialization skipped: kernel has no schedulable work";
return false;
}
std::vector<SchedOp *> leftovers;
@@ -797,7 +802,7 @@ private:
two_roles = two_roles && (role == producer || role == consumer);
}
if (!two_roles) {
LOG(WARNING) << "AutoSchedule skipped: storage '"
LOG(WARNING) << "AutoWarpSpecialization skipped: storage '"
<< resolution.allocation->name
<< "' is handed between more than two roles in scope '"
<< scope.id << "'";
@@ -879,7 +884,7 @@ private:
ok = false;
}
if (!ok) {
LOG(WARNING) << "AutoSchedule skipped: storage '"
LOG(WARNING) << "AutoWarpSpecialization skipped: storage '"
<< resolution.allocation->name << "' in scope '"
<< scope.id << "': " << reason;
resolution.failed = true;
@@ -934,7 +939,8 @@ private:
// Under versioning, a guard-skipped write would expose the slot from
// `depth` iterations ago instead of the previous value.
if (guarded_writer && pipeline->depth > 1) {
LOG(WARNING) << "AutoSchedule skipped: storage '" << pipeline->name
LOG(WARNING) << "AutoWarpSpecialization skipped: storage '"
<< pipeline->name
<< "' is written under a guard and would be "
<< pipeline->depth
<< "-way versioned; a skipped write would expose a "
@@ -976,7 +982,7 @@ private:
for (const SchedOp *op : buffer_uses_[buffer->data].touches)
others.Add(op->roles);
if (!writer->roles.ContainsAll(others)) {
LOG(WARNING) << "AutoSchedule skipped: global buffer '"
LOG(WARNING) << "AutoWarpSpecialization skipped: global buffer '"
<< buffer->name << "' is written by op '" << writer->id
<< "' and touched by another role";
return false;
@@ -1161,7 +1167,7 @@ private:
int max_threads = static_cast<int>(
target_->GetAttr<Integer>("max_num_threads").value_or(1024)->value);
if (num_warps * 32 > max_threads) {
LOG(WARNING) << "AutoSchedule skipped: " << num_warps
LOG(WARNING) << "AutoWarpSpecialization skipped: " << num_warps
<< " warps exceed the target's " << max_threads
<< "-thread block limit";
return std::nullopt;
@@ -1257,5 +1263,6 @@ ffi::Optional<WSSchedule> RoleBasedSchedule(const SBlock &block,
return RoleBasedScheduler(target, worker_threads).Run(block, body);
}
} // namespace cuda
} // namespace tl
} // namespace tvm
@@ -8,6 +8,7 @@
namespace tvm {
namespace tl {
namespace cuda {
// Fixed-role heuristic: classify ops by lowering eligibility (Load / MMA /
// Store / Worker), pull warp-private def-use chains into their consumers'
@@ -18,5 +19,6 @@ ffi::Optional<WSSchedule> RoleBasedSchedule(const tirx::SBlock &block,
int worker_threads,
const Target &target);
} // namespace cuda
} // namespace tl
} // namespace tvm
@@ -100,6 +100,7 @@ namespace tl {
using namespace tirx;
using namespace ffi;
using namespace cuda;
namespace {
+96 -22
View File
@@ -30,10 +30,10 @@ class SimtVFFrameNode;
class SimtVFFrame;
// Build a ForFrame that emits a target-neutral kThreadBinding loop for one
// kernel-launch dimension. The launch nest is materialized into the
// target-specific form (thread_extent AttrStmt on GPU, serial For on CPU) by
// the tl.MaterializeKernelLaunch pass once the Target is known at compile
// time.
// grid (program index) axis of a kernel launch. The launch nest is
// materialized into the target-specific form (thread_extent AttrStmt on GPU,
// serial For on CPU) by the tl.MaterializeKernelLaunch pass once the Target is
// known at compile time.
static ForFrame MakeThreadBindingFrame(const std::string &name,
const String &thread_tag,
const PrimExpr &extent) {
@@ -269,6 +269,38 @@ ForFrame PersistentFor(const Array<PrimExpr> &domain, const PrimExpr &wave_size,
return ForFrame(n);
}
// Build a frame whose exit prefixes the body with
// `tx = tl.launch_thread_idx(0); ty = ...; tz = ...` Bind statements. The
// launch nest is traced before the Target is known, so the thread indices are
// only placeholders here: the Vars keep their identity through
// tl.MaterializeKernelLaunch, which rebinds them as threadIdx.* thread_extent
// scopes on SIMT backends and drops them elsewhere.
static ForFrame MakeLaunchThreadFrame() {
using namespace tvm::tirx;
static const char *kThreadVarNames[3] = {"tx", "ty", "tz"};
DataType dtype = DataType::Int(32);
ObjectPtr<ForFrameNode> n = make_object<ForFrameNode>();
for (int axis = 0; axis < 3; axis++) {
n->vars.push_back(Var(kThreadVarNames[axis], dtype));
// The extent is decided by the backend at materialization; this dom only
// keeps the ForFrame invariants satisfied.
n->doms.push_back(Range(make_const(dtype, 0), make_const(dtype, 1)));
}
n->f_make_for_loop = [](const Array<Var> &vars, const Array<Range> &doms,
const Array<Optional<PrimExpr>> &steps,
Stmt body) -> Stmt {
Array<Stmt> seq;
for (int axis = 0; axis < static_cast<int>(vars.size()); axis++) {
PrimExpr thread_idx = Call(vars[axis]->dtype, launch_thread_idx(),
{IntImm(DataType::Int(32), axis)});
seq.push_back(tvm::tirx::Bind(vars[axis], thread_idx));
}
seq.push_back(body);
return SeqStmt::Flatten(seq);
};
return ForFrame(n);
}
/*!
* \brief A frame that represents a kernel launch.
*
@@ -276,12 +308,26 @@ ForFrame PersistentFor(const Array<PrimExpr> &domain, const PrimExpr &wave_size,
*/
class KernelLaunchFrameNode : public TIRFrameNode {
public:
/*! \brief Grid loops, thread placeholders and the root block, outer to
* inner. */
Array<TIRFrame> frames;
/*! \brief Program (grid) index vars, one per launch axis. */
Array<tvm::tirx::Var> grid_vars;
/*! \brief Grid extents, one per launch axis. */
Array<PrimExpr> grid_extents;
/*! \brief Placeholder thread index vars for the x, y and z axes. */
Array<tvm::tirx::Var> thread_vars;
/*! \brief Requested SIMT thread-block extents, when threads= was given. */
Optional<Array<PrimExpr>> thread_extents;
static void RegisterReflection() {
namespace refl = reflection;
refl::ObjectDef<KernelLaunchFrameNode>().def_ro(
"frames", &KernelLaunchFrameNode::frames);
refl::ObjectDef<KernelLaunchFrameNode>()
.def_ro("frames", &KernelLaunchFrameNode::frames)
.def_ro("grid_vars", &KernelLaunchFrameNode::grid_vars)
.def_ro("grid_extents", &KernelLaunchFrameNode::grid_extents)
.def_ro("thread_vars", &KernelLaunchFrameNode::thread_vars)
.def_ro("thread_extents", &KernelLaunchFrameNode::thread_extents);
}
TVM_FFI_DECLARE_OBJECT_INFO_FINAL("tl.KernelLaunchFrame",
@@ -323,30 +369,40 @@ KernelLaunchFrame KernelLaunch(const Array<PrimExpr> &grid_size,
const Map<String, Any> &attrs) {
ObjectPtr<KernelLaunchFrameNode> n = make_object<KernelLaunchFrameNode>();
auto block_size = block_size_opt.value_or(Array<PrimExpr>());
ICHECK(grid_size.size() <= 3);
ICHECK(block_size.size() <= 3);
static const char *kBlockVarNames[3] = {"bx", "by", "bz"};
static const char *kBlockTags[3] = {"blockIdx.x", "blockIdx.y", "blockIdx.z"};
static const char *kThreadVarNames[3] = {"tx", "ty", "tz"};
static const char *kThreadTags[3] = {"threadIdx.x", "threadIdx.y",
"threadIdx.z"};
for (size_t i = 0; i < grid_size.size(); i++) {
n->frames.push_back(
MakeThreadBindingFrame(kBlockVarNames[i], kBlockTags[i], grid_size[i]));
ForFrame frame =
MakeThreadBindingFrame(kBlockVarNames[i], kBlockTags[i], grid_size[i]);
n->grid_vars.push_back(frame->vars[0]);
n->grid_extents.push_back(grid_size[i]);
n->frames.push_back(frame);
}
for (size_t i = 0; i < block_size.size(); i++) {
n->frames.push_back(MakeThreadBindingFrame(kThreadVarNames[i],
kThreadTags[i], block_size[i]));
// Thread placeholders are always emitted so the body may reference a thread
// index regardless of whether threads= was given; the backend decides what
// they mean.
ForFrame thread_frame = MakeLaunchThreadFrame();
n->thread_vars = thread_frame->vars;
n->frames.push_back(thread_frame);
Map<String, Any> block_annotations =
attrs.defined() ? attrs : Map<String, Any>{};
if (block_size_opt.defined()) {
Array<PrimExpr> block_size = block_size_opt.value();
ICHECK(block_size.size() <= 3);
while (block_size.size() < 3) {
block_size.push_back(IntImm(DataType::Int(32), 1));
}
n->thread_extents = block_size;
block_annotations.Set(attr::kLaunchThreads, block_size);
}
auto empty_block = tvm::script::ir_builder::tirx::Block(DeviceMainBlockName);
empty_block->reads = Array<tvm::tirx::BufferRegion>();
empty_block->writes = Array<tvm::tirx::BufferRegion>();
Map<String, Any> block_annotations =
attrs.defined() ? attrs : Map<String, Any>{};
empty_block->annotations = block_annotations;
n->frames.push_back(empty_block);
@@ -368,14 +424,32 @@ KernelLaunchFrame MixedKernelLaunch(const Array<PrimExpr> &grid_size,
// Frame 0: bx = blockIdx.x. Emit a target-neutral thread_binding For loop;
// tl.MaterializeKernelLaunch turns it into a thread_extent AttrStmt.
n->frames.push_back(MakeThreadBindingFrame("bx", "blockIdx.x", grid_size[0]));
ForFrame bx_frame = MakeThreadBindingFrame("bx", "blockIdx.x", grid_size[0]);
n->grid_vars.push_back(bx_frame->vars[0]);
n->grid_extents.push_back(grid_size[0]);
n->frames.push_back(bx_frame);
// Frame 1: sid = get_subblockid() via the Ascend "cthread" binding. Also
// emitted as a thread_binding For loop; MaterializeKernelLaunch recognizes
// the "cthread" tag and materializes it into a thread_extent AttrStmt.
n->frames.push_back(MakeThreadBindingFrame("sid", "cthread", cthread_extent));
// the "cthread" tag (declared by the Ascend pipeline in launch_dim_tags) and
// materializes it into a thread_extent AttrStmt. It is a block-level launch
// dimension of the mixed kernel, so it is reported alongside bx as one of the
// vars the launch yields (`with T.MixedKernel(...) as (bx, sid)`).
ForFrame sid_frame =
MakeThreadBindingFrame("sid", "cthread", cthread_extent);
n->grid_vars.push_back(sid_frame->vars[0]);
n->grid_extents.push_back(cthread_extent);
n->frames.push_back(sid_frame);
// Frame 2: MainBlock with NPU marker
// Frame 2: thread placeholders, dropped by the Ascend pipeline
// (lower_thread_binding=false) exactly as for T.Kernel. They exist so a body
// that references a thread index gets the actionable "no SIMT threads"
// diagnostic instead of an out-of-range frame lookup.
ForFrame thread_frame = MakeLaunchThreadFrame();
n->thread_vars = thread_frame->vars;
n->frames.push_back(thread_frame);
// Frame 3: MainBlock with NPU marker
auto main_block = tvm::script::ir_builder::tirx::Block(DeviceMainBlockName);
main_block->reads = Array<tvm::tirx::BufferRegion>();
main_block->writes = Array<tvm::tirx::BufferRegion>();
+5
View File
@@ -87,6 +87,11 @@ TIR_DEFINE_TL_BUILTIN(access_ptr)
TIR_DEFINE_TL_BUILTIN(region).set_num_inputs(-1).set_attr<TCallEffectKind>(
"TCallEffectKind", Integer(CallEffectKind::kPure));
TIR_DEFINE_TL_BUILTIN(launch_thread_idx)
.set_num_inputs(1)
.set_attr<TCallEffectKind>("TCallEffectKind",
Integer(CallEffectKind::kOpaque));
TIR_DEFINE_TL_BUILTIN(add2).set_num_inputs(2).set_attr<TCallEffectKind>(
"TCallEffectKind", Integer(CallEffectKind::kPure));
+19
View File
@@ -32,6 +32,12 @@ static constexpr const char *kLocalVarInit = "tl.local_var_init";
static constexpr const char *kNonRestrictParams = "tl.non_restrict_params";
static constexpr const char *kLexicalAllocScope = "lexical_alloc_scope";
// Annotation on the tilelang_root block recording the SIMT thread-block
// extents requested by T.Kernel(threads=...). It is a launch hint: SIMT
// backends materialize it as threadIdx.* thread_extent scopes, other
// backends ignore it.
static constexpr const char *kLaunchThreads = "tl.launch_threads";
} // namespace attr
inline ffi::Optional<PrimExpr> GetAnnotatedMbarPhaseExpr(
@@ -185,6 +191,19 @@ TVM_DLL const Op &access_ptr();
*/
TVM_DLL const Op &region();
/*!
* \brief Placeholder for the thread index along one launch axis.
*
* T.Kernel binds each thread variable as `LetStmt(tx, launch_thread_idx(axis))`
* so the kernel body can reference a thread index before the target is known.
* tl.MaterializeKernelLaunch replaces the binding with a real threadIdx.*
* thread_extent scope on SIMT backends and rejects any use on backends
* without SIMT. It must never reach codegen.
*
* int32 launch_thread_idx(axis)
*/
TVM_DLL const Op &launch_thread_idx();
// Packed x2 element-wise math (float32x2, bfloat16x2, float16x2)
TVM_DLL const Op &add2();
TVM_DLL const Op &sub2();
+3 -5
View File
@@ -24,11 +24,9 @@ inline bool IsDeviceMainBlock(const tirx::SBlockNode *node) {
return node->name_hint == DeviceMainBlockName;
}
constexpr const char *tilelang_is_cpu_kernel_frame =
"tilelang.is_cpu_kernel_frame";
constexpr const char *tilelang_is_npu_kernel_frame =
"tilelang.is_npu_kernel_frame";
// The kernel-launch frame no longer needs a per-backend marker annotation:
// KernelLaunchFrame exposes its grid_vars/thread_vars explicitly and each
// backend pipeline decides what the launch means (MaterializeKernelLaunch).
constexpr const char *tilelang_simt_vf_captures = "tl.simt_vf_captures";
+264 -57
View File
@@ -3,40 +3,83 @@
* \brief Materialize the target-neutral kernel launch nest emitted by
* T.Kernel into a backend-specific form.
*
* T.Kernel traces into a nest of For loops with ForKind::kThreadBinding
* tagged blockIdx.* / threadIdx.*. This pass runs right after BindTarget
* and rewrites the nest according to `lower_thread_binding`, which each
* backend pipeline chooses for itself (no target dispatch happens here):
* T.Kernel traces into
*
* for bx in thread_binding(blockIdx.x): # one per grid axis
* tx = tl.launch_thread_idx(0) # x, y, z placeholders
* ty = tl.launch_thread_idx(1)
* tz = tl.launch_thread_idx(2)
* block tilelang_root { annotations: tl.launch_threads?, ... }
*
* The grid loops are the program-index space every backend shares. The
* thread placeholders only reserve Var identities the body may reference;
* what they mean, and how many threads run, is decided here once the Target
* is bound. Each backend pipeline chooses the mode for itself (no target
* dispatch happens in this pass).
*
* "How a launch dimension is lowered" and "whether SIMT threads exist" are
* independent axes, so they are two flags rather than one:
* - lower_grid_binding = true: every launch loop becomes an AttrStmt
* thread_extent scope carrying that loop's own thread tag. This is what
* targets with a real block/core-level launch want (CUDA, Ascend, ...).
* false (e.g. CPU): the launch loops become plain serial For loops, since
* such targets have no program-index space at all.
* - lower_thread_binding = true (SIMT backends, e.g. CUDA/ROCm/Metal):
* each launch loop becomes an AttrStmt thread_extent, reusing the loop
* variable so body references stay valid.
* - lower_thread_binding = false (backends without SIMT, e.g. CPU):
* blockIdx.* loops become serial For loops over the grid extent;
* threadIdx.* loops are ignored; they become unit serial loops so the
* loop variable stays defined (pinned to 0) while the requested thread
* extent (e.g. the default threads=128) is dropped.
* each thread placeholder is rebound as a threadIdx.* thread_extent over
* the same Var, with the extent taken from the `tl.launch_threads`
* annotation (T.Kernel threads=...) or, failing that, from
* `default_threads`.
* false (backends without SIMT, e.g. CPU and Ascend): thread placeholders
* are dropped and `tl.launch_threads` is ignored. A body that references a
* thread index has no meaning on such a target and is rejected. Ascend
* pairs this with lower_grid_binding = true: its NPU launch is a real 1-D
* core grid, while thread domains only exist inside T.SimtVF, which emits
* its own thread scopes below the launch nest.
*
* `launch_dim_tags` names extra thread_binding tags that belong to the launch
* nest rather than to the thread domain, so a backend can extend the launch
* vocabulary without this pass knowing about it. Ascend's T.MixedKernel uses
* it for its `cthread` sub-block-id dimension, which its codegen reads back
* off the thread_extent AttrStmt.
*
* Launch annotations listed in `unsupported_annotations` (e.g. `cluster_dims`
* on a target without thread block clusters) are rejected instead of being
* silently dropped further down the pipeline.
*
* Only the outermost contiguous launch nest is converted; thread_binding
* loops deeper inside the kernel body (separated by the tilelang_root
* block) are left for LowerOpaqueBlock to handle at its usual stage.
*/
#include "../op/builtin.h"
#include "common/attr.h"
#include "support/check.h"
#include <tvm/ir/transform.h>
#include <tvm/runtime/logging.h>
#include <tvm/target/target.h>
#include <tvm/tirx/analysis.h>
#include <tvm/tirx/op.h>
#include <tvm/tirx/stmt_functor.h>
#include <tvm/tirx/transform.h>
#include <algorithm>
#include <utility>
#include <vector>
namespace tvm {
namespace tl {
using namespace tirx;
using ffi::Array;
using ffi::GetRef;
using ffi::Optional;
namespace {
constexpr int kNumThreadAxes = 3;
constexpr const char *kThreadTags[kNumThreadAxes] = {
"threadIdx.x", "threadIdx.y", "threadIdx.z"};
bool IsBlockBinding(const ForNode *op) {
if (op->kind != ForKind::kThreadBinding || !op->thread_binding.defined())
return false;
@@ -44,78 +87,242 @@ bool IsBlockBinding(const ForNode *op) {
return tag.rfind("blockIdx.", 0) == 0;
}
bool IsThreadBinding(const ForNode *op) {
if (op->kind != ForKind::kThreadBinding || !op->thread_binding.defined())
// `v = tl.launch_thread_idx(axis)` emitted by T.Kernel.
bool IsLaunchThreadPlaceholder(const Stmt &stmt) {
const BindNode *bind = stmt.as<BindNode>();
if (!bind)
return false;
std::string tag = op->thread_binding.value()->thread_tag;
return tag.rfind("threadIdx.", 0) == 0;
const CallNode *call = bind->value.as<CallNode>();
return call && call->op.same_as(launch_thread_idx());
}
// Ascend mixed-kernel sub-block id binding. It behaves like a block-level
// launch dimension (materialized into a thread_extent AttrStmt on the SIMT
// path) and is emitted at the top launch nest by T.MixedKernel.
bool IsCthreadBinding(const ForNode *op) {
if (op->kind != ForKind::kThreadBinding || !op->thread_binding.defined())
return false;
return op->thread_binding.value()->thread_tag == "cthread";
int ThreadAxisOf(const BindNode *bind) {
const CallNode *call = bind->value.as<CallNode>();
ICHECK(call && call->args.size() == 1);
const int64_t *axis = as_const_int(call->args[0]);
ICHECK(axis && *axis >= 0 && *axis < kNumThreadAxes)
<< "tl.launch_thread_idx expects a constant axis in [0, 3), got "
<< call->args[0];
return static_cast<int>(*axis);
}
bool IsLaunchBinding(const ForNode *op) {
return IsBlockBinding(op) || IsThreadBinding(op) || IsCthreadBinding(op);
// The tilelang_root block that carries the launch annotations, if `body` is
// the kernel body directly below the launch nest.
const SBlockNode *GetLaunchBlock(const Stmt &body) {
const SBlockRealizeNode *realize = body.as<SBlockRealizeNode>();
if (!realize || !IsDeviceMainBlock(realize->block.get()))
return nullptr;
return realize->block.get();
}
Optional<Array<PrimExpr>> GetLaunchThreads(const Stmt &body) {
const SBlockNode *block = GetLaunchBlock(body);
if (!block)
return std::nullopt;
if (auto threads = block->annotations.Get(attr::kLaunchThreads)) {
if (auto arr = threads.value().try_cast<Array<PrimExpr>>())
return arr.value();
LOG(FATAL) << "Expected `" << attr::kLaunchThreads
<< "` to be an Array<PrimExpr>, but got "
<< threads.value().GetTypeKey();
}
return std::nullopt;
}
class KernelLaunchMaterializer : public StmtMutator {
public:
explicit KernelLaunchMaterializer(bool lower_thread_binding)
: lower_thread_binding_(lower_thread_binding) {}
KernelLaunchMaterializer(bool lower_grid_binding, bool lower_thread_binding,
Optional<Array<PrimExpr>> default_threads,
Array<ffi::String> unsupported_annotations,
Array<ffi::String> launch_dim_tags,
ffi::String target_name)
: lower_grid_binding_(lower_grid_binding),
lower_thread_binding_(lower_thread_binding),
default_threads_(std::move(default_threads)),
unsupported_annotations_(std::move(unsupported_annotations)),
launch_dim_tags_(std::move(launch_dim_tags)),
target_name_(std::move(target_name)) {}
Stmt VisitStmt_(const ForNode *op) final {
if (IsLaunchBinding(op)) {
return ConvertNest(op);
if (IsLaunchDimBinding(op)) {
return ConvertNest(GetRef<Stmt>(op));
}
return StmtMutator::VisitStmt_(op);
}
// A launch without grid axes starts directly at the thread placeholders.
Stmt VisitStmt_(const SeqStmtNode *op) final {
if (op->size() > 0 && IsLaunchThreadPlaceholder(op->seq[0])) {
return ConvertNest(GetRef<Stmt>(op));
}
return StmtMutator::VisitStmt_(op);
}
private:
// Peel the contiguous launch nest rooted at `op` without descending into
// the kernel body below it.
Stmt ConvertNest(const ForNode *op) {
Stmt body;
if (const ForNode *inner = op->body.as<ForNode>();
inner && IsLaunchBinding(inner)) {
body = ConvertNest(inner);
} else {
body = op->body;
}
if (lower_thread_binding_) {
ffi::String tag = op->thread_binding.value()->thread_tag;
IterVar iter_var(Range::FromMinExtent(op->min, op->extent), op->loop_var,
IterVarType::kThreadIndex, tag);
return AttrStmt(std::move(iter_var), tirx::attr::thread_extent,
op->extent, std::move(body), op->span);
}
// No SIMT: grid dims run as plain serial loops; thread dims are ignored
// (a unit loop keeps the loop variable defined and pinned to 0).
PrimExpr extent = IsThreadBinding(op)
? PrimExpr(IntImm(op->extent.dtype(), 1))
: op->extent;
return For(op->loop_var, op->min, std::move(extent), ForKind::kSerial,
std::move(body),
/*thread_binding=*/std::nullopt, op->annotations, op->step,
op->span);
// A launch dimension: the blockIdx.* grid axes every target shares, plus the
// tags this backend declared in `launch_dim_tags` (e.g. Ascend's `cthread`).
bool IsLaunchDimBinding(const ForNode *op) const {
if (IsBlockBinding(op))
return true;
if (op->kind != ForKind::kThreadBinding || !op->thread_binding.defined())
return false;
const std::string tag = op->thread_binding.value()->thread_tag;
return std::find(launch_dim_tags_.begin(), launch_dim_tags_.end(), tag) !=
launch_dim_tags_.end();
}
// Peel the contiguous launch nest rooted at `root` without descending into
// the kernel body below it, then rebuild it in the backend's form.
Stmt ConvertNest(const Stmt &root) {
std::vector<const ForNode *> grid_loops;
Stmt body = root;
while (const ForNode *loop = body.as<ForNode>()) {
if (!IsLaunchDimBinding(loop))
break;
grid_loops.push_back(loop);
body = loop->body;
}
std::vector<const BindNode *> thread_binds;
body = PeelThreadPlaceholders(body, &thread_binds);
RejectUnsupportedAnnotations(body);
body = lower_thread_binding_ ? BindThreads(thread_binds, body)
: DropThreads(thread_binds, body);
for (auto it = grid_loops.rbegin(); it != grid_loops.rend(); ++it) {
const ForNode *loop = *it;
if (lower_grid_binding_) {
ffi::String tag = loop->thread_binding.value()->thread_tag;
IterVar iter_var(Range::FromMinExtent(loop->min, loop->extent),
loop->loop_var, IterVarType::kThreadIndex, tag);
body = AttrStmt(std::move(iter_var), tirx::attr::thread_extent,
loop->extent, std::move(body), loop->span);
} else {
body = For(loop->loop_var, loop->min, loop->extent, ForKind::kSerial,
std::move(body),
/*thread_binding=*/std::nullopt, loop->annotations,
loop->step, loop->span);
}
}
return body;
}
// Split the leading `v = tl.launch_thread_idx(axis)` binds off `stmt` and
// return what follows them.
static Stmt PeelThreadPlaceholders(const Stmt &stmt,
std::vector<const BindNode *> *binds) {
const SeqStmtNode *seq = stmt.as<SeqStmtNode>();
if (!seq)
return stmt;
size_t i = 0;
while (i < seq->size() && IsLaunchThreadPlaceholder(seq->seq[i])) {
binds->push_back(seq->seq[i].as<BindNode>());
++i;
}
if (i == 0)
return stmt;
ICHECK_LT(i, seq->size())
<< "T.Kernel launch has thread placeholders but no body";
if (i + 1 == seq->size())
return seq->seq[i];
return SeqStmt(Array<Stmt>(seq->seq.begin() + i, seq->seq.end()));
}
// SIMT: every placeholder becomes a threadIdx.* thread_extent scope over
// the same Var so body references stay valid.
Stmt BindThreads(const std::vector<const BindNode *> &thread_binds,
Stmt body) {
if (thread_binds.empty())
return body;
Array<PrimExpr> extents = ResolveThreadExtents(body);
for (auto it = thread_binds.rbegin(); it != thread_binds.rend(); ++it) {
const BindNode *bind = *it;
int axis = ThreadAxisOf(bind);
PrimExpr extent = extents[axis];
IterVar iter_var(
Range::FromMinExtent(make_zero(bind->var.dtype()), extent), bind->var,
IterVarType::kThreadIndex, kThreadTags[axis]);
body = AttrStmt(std::move(iter_var), tirx::attr::thread_extent, extent,
std::move(body), bind->span);
}
return body;
}
// No SIMT: thread placeholders carry no meaning, so they are removed. A body
// that reads one would otherwise silently run as a single thread.
Stmt DropThreads(const std::vector<const BindNode *> &thread_binds,
Stmt body) {
for (const BindNode *bind : thread_binds) {
const VarNode *var = bind->var.get();
if (UsesVar(body, [var](const VarNode *v) { return v == var; })) {
LOG(FATAL) << "T.Kernel body references thread index `"
<< bind->var->name_hint << "`, but target `" << target_name_
<< "` has no SIMT threads. Express the computation with "
"tile-level operators (T.Parallel, T.copy, ...) instead "
"of per-thread indexing.";
}
}
return body;
}
void RejectUnsupportedAnnotations(const Stmt &body) const {
const SBlockNode *block = GetLaunchBlock(body);
if (!block)
return;
for (const ffi::String &key : unsupported_annotations_) {
if (block->annotations.count(key)) {
LOG(FATAL) << "T.Kernel launch annotation `" << key
<< "` is not supported on target `" << target_name_ << "`";
}
}
}
Array<PrimExpr> ResolveThreadExtents(const Stmt &body) {
Optional<Array<PrimExpr>> threads = GetLaunchThreads(body);
if (!threads.defined())
threads = default_threads_;
ICHECK(threads.defined())
<< "T.Kernel did not specify threads= and target `" << target_name_
<< "` provides no default thread-block size";
Array<PrimExpr> extents = threads.value();
ICHECK_LE(extents.size(), static_cast<size_t>(kNumThreadAxes));
while (extents.size() < static_cast<size_t>(kNumThreadAxes)) {
extents.push_back(IntImm(DataType::Int(32), 1));
}
return extents;
}
bool lower_grid_binding_;
bool lower_thread_binding_;
Optional<Array<PrimExpr>> default_threads_;
Array<ffi::String> unsupported_annotations_;
Array<ffi::String> launch_dim_tags_;
ffi::String target_name_;
};
} // namespace
tvm::transform::Pass MaterializeKernelLaunch(bool lower_thread_binding) {
tvm::transform::Pass
MaterializeKernelLaunch(bool lower_grid_binding, bool lower_thread_binding,
Optional<Array<PrimExpr>> default_threads,
Optional<Array<ffi::String>> unsupported_annotations,
Optional<Array<ffi::String>> launch_dim_tags) {
using namespace tirx::transform;
auto pass_func = [lower_thread_binding](
Array<ffi::String> unsupported =
unsupported_annotations.value_or(Array<ffi::String>());
Array<ffi::String> dim_tags = launch_dim_tags.value_or(Array<ffi::String>());
auto pass_func = [lower_grid_binding, lower_thread_binding, default_threads,
unsupported, dim_tags](
PrimFunc func, const IRModule &mod,
const tvm::transform::PassContext &ctx) -> PrimFunc {
KernelLaunchMaterializer mutator(lower_thread_binding);
ffi::String target_name = "<unbound>";
if (auto target = func->GetAttr<Target>(tvm::attr::kTarget)) {
target_name = target.value()->kind->name;
}
KernelLaunchMaterializer mutator(lower_grid_binding, lower_thread_binding,
default_threads, unsupported, dim_tags,
target_name);
func.CopyOnWrite()->body = mutator(func->body);
return func;
};
@@ -24,13 +24,13 @@ def test_cuda_and_ascend_auto_schedule_options_coexist(enabled):
with tilelang.transform.PassContext(
config={
tilelang.PassConfigKey.TL_ENABLE_AUTO_SCHEDULE: enabled,
tilelang.PassConfigKey.TL_CUDA_AUTO_SCHEDULE: "role_based",
tilelang.PassConfigKey.TL_ENABLE_AUTO_WARP_SPECIALIZATION: "role_based",
}
) as context:
assert bool(allow_autoschedule(context)) is enabled
assert context.config[tilelang.PassConfigKey.TL_CUDA_AUTO_SCHEDULE] == "role_based"
assert context.config[tilelang.PassConfigKey.TL_ENABLE_AUTO_WARP_SPECIALIZATION] == "role_based"
assert tilelang.ascend.transform.AutoSchedule().info.name == "tl.AutoSchedule"
assert tilelang.cuda.transform.AutoSchedule().info.name == "tl.cuda.AutoSchedule"
assert tilelang.cuda.transform.AutoWarpSpecialization().info.name == "tl.AutoWarpSpecialization"
if __name__ == "__main__":
+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
+12 -12
View File
@@ -49,9 +49,9 @@ def _simple_program():
@T.prim_func
def program(A: T.Tensor((128,), "float32"), B: T.Tensor((128,), "float32")):
with T.Kernel(threads=128):
tid = T.get_thread_binding()
B[tid] = A[tid] + 1.0
with T.Kernel(1):
for i in T.serial(128):
B[i] = A[i] + 1.0
return program
@@ -186,9 +186,9 @@ def test_multiple_pipelines_share_one_compile_session(monkeypatch, tmp_path):
@T.prim_func
def tiny(A: T.Tensor((32,), "float32"), B: T.Tensor((32,), "float32")):
with T.Kernel(32):
tid = T.get_thread_binding()
B[tid] = A[tid] + 1.0
with T.Kernel(1):
for i in T.serial(32):
B[i] = A[i] + 1.0
mod = tvm.IRModule({"main": tiny})
context = create_backend_context("c", "c", "cython")
@@ -279,9 +279,9 @@ def test_no_skipped_phantom_records(monkeypatch, tmp_path):
@T.prim_func
def tiny(A: T.Tensor((32,), "float32"), B: T.Tensor((32,), "float32")):
with T.Kernel(32):
tid = T.get_thread_binding()
B[tid] = A[tid] + 1.0
with T.Kernel(1):
for i in T.serial(32):
B[i] = A[i] + 1.0
tilelang.lower(tiny, target="c")
@@ -325,9 +325,9 @@ def test_terminal_mode_no_html(monkeypatch, tmp_path):
@T.prim_func
def tiny(A: T.Tensor((32,), "float32"), B: T.Tensor((32,), "float32")):
with T.Kernel(32):
tid = T.get_thread_binding()
B[tid] = A[tid] + 1.0
with T.Kernel(1):
for i in T.serial(32):
B[i] = A[i] + 1.0
tilelang.lower(tiny, target="c")
@@ -44,6 +44,22 @@ def get_cta_rank_in_cluster(cluster_size=4):
return main
@tilelang.jit(out_idx=-1)
def get_cluster_id_kernel(cluster_size=4):
assert 128 % cluster_size == 0
@T.prim_func
def main(A: T.Tensor((128, 2), T.int32)):
with T.ClusterKernel(128, cluster_dims=(cluster_size, 1, 1)) as bx:
if T.get_thread_binding() == 0:
A[bx, 0] = T.get_cluster_id()
# Program-space cluster id and the hardware rank must agree on
# which programs form a cluster.
A[bx, 1] = T.get_cluster_id() * T.get_cluster_size() + T.block_rank_in_cluster()
return main
@tilelang.jit(out_idx=-1)
def barrier_kernel():
@T.prim_func
@@ -110,6 +126,15 @@ def test_cluster_launch_intrinsics(cluster_size=4):
assert torch.all(result == ref)
@tilelang.testing.requires_cuda
@tilelang.testing.requires_cuda_compute_version_ge(9, 0)
def test_cluster_id_matches_hardware_rank(cluster_size=4):
result = get_cluster_id_kernel(cluster_size)()
bx = torch.arange(128, dtype=torch.int32, device="cuda")
assert torch.all(result[:, 0] == bx // cluster_size)
assert torch.all(result[:, 1] == bx)
@tilelang.testing.requires_cuda
@tilelang.testing.requires_cuda_compute_version_ge(9, 0)
def test_cluster_barrier():
@@ -238,5 +238,83 @@ def test_jit2_compile_with_consts():
transpose.compile(M=1024, N=1024, block_M=64, block_N=64)
def _kernel_body_is_skipped_in_phase1(func) -> bool:
"""Whether the eager rewriter guards the launch body with skip_kernel_ctx,
which is what keeps phase-1 signature inference from executing it."""
from tilelang.language.eager.ast import mutate
return "skip_kernel_ctx" in mutate(func).source
def test_jit2_recognizes_launch_from_every_dialect():
from tilelang.cpu import language as Tcpu
from tilelang.cuda import language as Tcuda
from tilelang.rocm import language as Trocm
Launch = T.Kernel
def default_facade(A):
with T.Kernel(1):
pass
def cuda_dialect(A):
with Tcuda.Kernel(1, threads=128):
pass
def rocm_dialect(A):
with Trocm.Kernel(1, threads=64):
pass
def cpu_dialect(A):
with Tcpu.Kernel(1):
pass
def cluster(A):
with T.ClusterKernel(2, cluster_dims=2):
pass
def aliased(A):
with Launch(1):
pass
def not_a_launch(A):
with T.ws(0):
pass
for func in (default_facade, cuda_dialect, rocm_dialect, cpu_dialect, cluster, aliased):
assert _kernel_body_is_skipped_in_phase1(func), func.__name__
assert not _kernel_body_is_skipped_in_phase1(not_a_launch)
@tilelang.testing.requires_cuda
def test_jit2_phase1_does_not_execute_kernel_body():
"""Phase 1 infers the signature with symbolic T.const values, so the launch
body must not run then: a gemm tile that depends on a const is only valid
once the values are bound in phase 2."""
@tilelang.jit
def gemm_full_n(A, B, block_M, block_K):
M, N, K = T.const("M, N, K")
A: T.Tensor[[M, K], T.float16]
B: T.Tensor[[K, N], T.float16]
C = T.empty((M, N), T.float16)
with T.Kernel(T.ceildiv(M, block_M), threads=128) as bx:
A_shared = T.alloc_shared((block_M, block_K), T.float16)
B_shared = T.alloc_shared((block_K, N), T.float16)
C_local = T.alloc_fragment((block_M, N), T.float32)
T.clear(C_local)
for k in T.Pipelined(T.ceildiv(K, block_K), num_stages=2):
T.copy(A[bx * block_M, k * block_K], A_shared)
T.copy(B[k * block_K, 0], B_shared)
T.gemm(A_shared, B_shared, C_local)
T.copy(C_local, C[bx * block_M, 0])
return C
a = torch.randn(256, 128, device="cuda", dtype=torch.float16)
b = torch.randn(128, 64, device="cuda", dtype=torch.float16)
c = gemm_full_n(a, b, 64, 32)
torch.testing.assert_close(c, (a.float() @ b.float()).half(), rtol=1e-2, atol=1e-2)
if __name__ == "__main__":
tilelang.testing.main()
@@ -20,7 +20,6 @@ def test_normalize_threads_rejects_non_positive(threads):
@pytest.mark.parametrize(
"threads, expected",
[
(None, [128, 1, 1]),
(256, [256, 1, 1]),
([32, 4], [32, 4, 1]),
((32, 2, 2), [32, 2, 2]),
@@ -31,5 +30,11 @@ def test_normalize_threads_accepts_positive(threads, expected):
assert _normalize_threads(threads) == expected
def test_normalize_threads_leaves_default_to_backend():
"""No threads= means no SIMT hint: the backend picks its default when it
materializes the launch, the frontend does not guess one."""
assert _normalize_threads(None) is None
if __name__ == "__main__":
tilelang.testing.main()
@@ -1,4 +1,4 @@
"""Tests for AutoSchedule's "role_based" scheduler.
"""Tests for AutoWarpSpecialization's "role_based" scheduler.
The pass consumes plain (schedule-free) kernels, assigns fixed roles from
lowering eligibility (Load / MMA / Store / Worker), pulls warp-private
@@ -28,15 +28,15 @@ def _prepare(func):
return mod
def _auto_schedule(mod, scheduler="role_based"):
"""Apply AutoSchedule with the scheduler opted in via pass config."""
with tvm.transform.PassContext(config={"tl.cuda_auto_schedule": scheduler}):
return tilelang.cuda.transform.AutoSchedule()(mod)
def _auto_warp_specialization(mod, scheduler="role_based"):
"""Apply AutoWarpSpecialization with the scheduler opted in via pass config."""
with tvm.transform.PassContext(config={"tl.enable_auto_warp_specialization": scheduler}):
return tilelang.cuda.transform.AutoWarpSpecialization()(mod)
def _schedule(func):
"""Run the pass; returns (scheduled module, root WSSchedule or None)."""
scheduled = _auto_schedule(_prepare(func))
scheduled = _auto_warp_specialization(_prepare(func))
return scheduled, _root_schedule(scheduled["main"])
@@ -622,7 +622,7 @@ def test_guarded_write_into_versioned_pipeline_declines():
than silently re-versioned; guarded writes to single-buffered
pipelines (FA's rescale) keep source semantics and stay schedulable."""
prepared = _prepare(_guarded_producer_kernel())
scheduled = _auto_schedule(prepared)
scheduled = _auto_warp_specialization(prepared)
assert tvm.ir.structural_equal(scheduled["main"], prepared["main"])
@@ -746,7 +746,7 @@ def test_rmw_accumulator_numerical():
_rmw_accumulator_kernel(),
target="cuda",
out_idx=[3],
pass_configs={"tl.cuda_auto_schedule": "role_based"},
pass_configs={"tl.enable_auto_warp_specialization": "role_based"},
)
a = torch.randn((4, 128, 64), device="cuda", dtype=torch.float16)
b = torch.randn((4, 64, 64), device="cuda", dtype=torch.float16)
@@ -875,7 +875,7 @@ def test_read_before_nested_cycle_declines():
T.copy(F, B[w, k, 0, 0])
prepared = _prepare(kernel)
scheduled = _auto_schedule(prepared)
scheduled = _auto_warp_specialization(prepared)
assert tvm.ir.structural_equal(scheduled["main"], prepared["main"])
@@ -886,7 +886,7 @@ def test_post_loop_read_numerical():
_post_loop_read_kernel(),
target="cuda",
out_idx=[1, 2],
pass_configs={"tl.cuda_auto_schedule": "role_based"},
pass_configs={"tl.enable_auto_warp_specialization": "role_based"},
)
a = torch.randn((2, 4, 64, 64), device="cuda", dtype=torch.float16)
b, c = kernel(a)
@@ -900,7 +900,7 @@ def test_pipeline_opt_in_runs_automatic_ws():
mod = tvm.IRModule.from_expr(func)
with (
_TARGET,
tvm.transform.PassContext(config={"tl.cuda_auto_schedule": "role_based"}),
tvm.transform.PassContext(config={"tl.enable_auto_warp_specialization": "role_based"}),
):
out = tilelang.cuda.pipeline.CUDAPassPipelineBodyPrologue(mod, _TARGET)["main"]
@@ -920,7 +920,7 @@ def test_no_shared_handoff_is_left_unchanged():
T.copy(F, B)
prepared = _prepare(kernel)
scheduled = _auto_schedule(prepared)
scheduled = _auto_warp_specialization(prepared)
assert tvm.ir.structural_equal(scheduled["main"], prepared["main"])
@@ -993,7 +993,7 @@ def test_while_scope_numerical():
_persistent_while_kernel(),
target="cuda",
out_idx=[1],
pass_configs={"tl.cuda_auto_schedule": "role_based"},
pass_configs={"tl.enable_auto_warp_specialization": "role_based"},
)
a = torch.randn((2, 64, 64), device="cuda", dtype=torch.float16)
torch.testing.assert_close(kernel(a), a)
@@ -1005,7 +1005,7 @@ def test_pipelined_load_numerical():
_pipelined_load_kernel(),
target="cuda",
out_idx=[1],
pass_configs={"tl.cuda_auto_schedule": "role_based"},
pass_configs={"tl.enable_auto_warp_specialization": "role_based"},
)
a = torch.randn((4, 64, 64), device="cuda", dtype=torch.float16)
torch.testing.assert_close(kernel(a), a)
@@ -1017,7 +1017,7 @@ def test_two_cycles_numerical():
_two_cycle_kernel(),
target="cuda",
out_idx=[1],
pass_configs={"tl.cuda_auto_schedule": "role_based"},
pass_configs={"tl.enable_auto_warp_specialization": "role_based"},
)
a = torch.randn((2, 64, 64), device="cuda", dtype=torch.float16)
torch.testing.assert_close(kernel(a), a)
@@ -1029,7 +1029,7 @@ def test_gather_bind_numerical():
_gather_kernel(worker_uses_index=True),
target="cuda",
out_idx=[2],
pass_configs={"tl.cuda_auto_schedule": "role_based"},
pass_configs={"tl.enable_auto_warp_specialization": "role_based"},
)
indices = torch.tensor([5, 2, 7, 0], device="cuda", dtype=torch.int32)
a = torch.randn((8, 64, 64), device="cuda", dtype=torch.float16)
@@ -1045,7 +1045,7 @@ def test_local_chain_numerical():
_local_chain_kernel(),
target="cuda",
out_idx=[1],
pass_configs={"tl.cuda_auto_schedule": "role_based"},
pass_configs={"tl.enable_auto_warp_specialization": "role_based"},
)
a = torch.randn((4, 64, 64), device="cuda", dtype=torch.float16)
expected = torch.stack([a[k * 2 % 4] for k in range(4)])
@@ -1058,7 +1058,7 @@ def test_worker_gemm_numerical():
_local_accumulator_gemm(),
target="cuda",
out_idx=[2],
pass_configs={"tl.cuda_auto_schedule": "role_based"},
pass_configs={"tl.enable_auto_warp_specialization": "role_based"},
)
a = torch.randn((64, 128), device="cuda", dtype=torch.float16)
b = torch.randn((128, 64), device="cuda", dtype=torch.float16)
@@ -1097,7 +1097,7 @@ def test_storage_cycling_in_sibling_loops_declines():
nothing chains — the second loop's pre-armed acquire could overwrite
data the first loop's consumer still reads. The kernel declines."""
prepared = _prepare(_two_loop_kernel())
scheduled = _auto_schedule(prepared)
scheduled = _auto_warp_specialization(prepared)
assert tvm.ir.structural_equal(scheduled["main"], prepared["main"])
@@ -1169,7 +1169,7 @@ def test_non_tma_layout_numerical():
_non_tma_layout_kernel(),
target="cuda",
out_idx=[1],
pass_configs={"tl.cuda_auto_schedule": "role_based"},
pass_configs={"tl.enable_auto_warp_specialization": "role_based"},
)
a = torch.randn((4, 64, 64), device="cuda", dtype=torch.float16)
torch.testing.assert_close(kernel(a), a + a)
@@ -1229,7 +1229,7 @@ def test_cp_async_load_numerical():
_cp_async_load_kernel(),
target="cuda",
out_idx=[1],
pass_configs={"tl.cuda_auto_schedule": "role_based"},
pass_configs={"tl.enable_auto_warp_specialization": "role_based"},
)
a = torch.randn((4, 64, 64), device="cuda", dtype=torch.float16)
torch.testing.assert_close(kernel(a), a)
@@ -1261,7 +1261,7 @@ def test_async_wgmma_wait_kernel_is_left_unchanged():
T.copy(C_local, C[0, 0])
prepared = _prepare(kernel)
scheduled = _auto_schedule(prepared)
scheduled = _auto_warp_specialization(prepared)
assert tvm.ir.structural_equal(scheduled["main"], prepared["main"])
@@ -1288,7 +1288,7 @@ def test_raw_cp_async_kernel_is_left_unchanged():
T.copy(F, A[k, 0, 0])
prepared = _prepare(kernel)
scheduled = _auto_schedule(prepared)
scheduled = _auto_warp_specialization(prepared)
assert tvm.ir.structural_equal(scheduled["main"], prepared["main"])
@@ -1322,7 +1322,7 @@ def test_pointer_table_bind_follows_freshened_buffer():
_pointer_table_kernel(),
target="cuda",
out_idx=[1],
pass_configs={"tl.cuda_auto_schedule": "role_based"},
pass_configs={"tl.enable_auto_warp_specialization": "role_based"},
)
src = torch.randn((64, 64), device="cuda", dtype=torch.float16)
ptrs = torch.tensor([src.data_ptr()], device="cuda", dtype=torch.int64)
@@ -1331,8 +1331,8 @@ def test_pointer_table_bind_follows_freshened_buffer():
@tilelang.testing.requires_cuda
def test_unknown_scheduler_rejected():
with pytest.raises(Exception, match="unknown auto-schedule scheduler"):
_auto_schedule(_prepare(_pipelined_load_kernel()), scheduler="nonexistent")
with pytest.raises(Exception, match="unknown auto-warp-specialization scheduler"):
_auto_warp_specialization(_prepare(_pipelined_load_kernel()), scheduler="nonexistent")
@tilelang.testing.requires_cuda
@@ -1371,7 +1371,7 @@ def test_unschedulable_constructs_decline():
for kernel in (atomic_kernel, hosted_async_kernel, sync_threads_kernel):
prepared = _prepare(kernel)
scheduled = _auto_schedule(prepared)
scheduled = _auto_warp_specialization(prepared)
assert tvm.ir.structural_equal(scheduled["main"], prepared["main"])
@@ -1382,7 +1382,7 @@ def test_tmem_gemm_numerical():
_tmem_gemm(),
target="cuda",
out_idx=[2],
pass_configs={"tl.cuda_auto_schedule": "role_based"},
pass_configs={"tl.enable_auto_warp_specialization": "role_based"},
)
a = torch.randn((2, 128, 64), device="cuda", dtype=torch.float16)
b = torch.randn((2, 128, 64), device="cuda", dtype=torch.float16)
@@ -1408,7 +1408,7 @@ def test_thread_budget_exceeded_declines():
T.copy(F, B[k, 0, 0])
prepared = _prepare(kernel)
scheduled = _auto_schedule(prepared)
scheduled = _auto_warp_specialization(prepared)
assert tvm.ir.structural_equal(scheduled["main"], prepared["main"])
@@ -1501,7 +1501,7 @@ def test_annotated_ws_pipeline_depth_on_inner_binding_declines():
rescale): versioning it would expose a stale slot, so the annotated
depth declines."""
prepared = _prepare(_waved_rmw_accumulator_kernel(inner_depth=2))
scheduled = _auto_schedule(prepared)
scheduled = _auto_warp_specialization(prepared)
assert tvm.ir.structural_equal(scheduled["main"], prepared["main"])
@@ -1512,7 +1512,7 @@ def test_annotated_ws_pipeline_depth_ring_numerical():
_waved_rmw_accumulator_kernel(depth=2),
target="cuda",
out_idx=[3],
pass_configs={"tl.cuda_auto_schedule": "role_based"},
pass_configs={"tl.enable_auto_warp_specialization": "role_based"},
)
a = torch.randn((2, 4, 128, 64), device="cuda", dtype=torch.float16)
b = torch.randn((2, 4, 64, 64), device="cuda", dtype=torch.float16)
@@ -26,6 +26,15 @@ def _strip_block_reads_writes(stmt, strip_annotations: bool = False):
return ir_transform(stmt, None, _postorder)
def _materialize_launch(func):
"""Run the launch materialization the pipeline performs before
LegalizeSafeMemoryAccess, so thread indices carry their extents."""
mod = tvm.IRModule({func.attrs["global_symbol"]: func})
mod = tvm.tirx.transform.BindTarget(tvm.target.Target("cuda"))(mod)
mod = tl.transform.MaterializeKernelLaunch()(mod)
return mod[func.attrs["global_symbol"]]
def _collect_call_nodes(stmt, op_names):
if isinstance(op_names, str):
op_names = {op_names}
@@ -113,6 +122,7 @@ def vectorize_access_legalize(M: int = 64, N: int = 64, M_offset: int = 2, N_off
def assert_vectorize_access(M: int = 64, N: int = 64):
func, expected = vectorize_access_legalize(M, N)
func, expected = _materialize_launch(func), _materialize_launch(expected)
mod = tvm.IRModule({func.attrs["global_symbol"]: func})
transformed = tl.transform.LegalizeSafeMemoryAccess()(mod)
@@ -156,6 +166,7 @@ def vectorize_access_with_atmoic_add_legalize(M: int = 64, N: int = 64, M_offset
def assert_vectorize_access_with_atmoic_add(M: int = 64, N: int = 64):
func, expected = vectorize_access_with_atmoic_add_legalize(M, N)
func, expected = _materialize_launch(func), _materialize_launch(expected)
mod = tvm.IRModule({func.attrs["global_symbol"]: func})
transformed = tl.transform.LegalizeSafeMemoryAccess()(mod)
print(transformed)
@@ -193,6 +204,7 @@ def oob_store_legalize(M: int = 64, N: int = 64, M_offset: int = 2, N_offset: in
def assert_oob_store_legalize(M: int = 64, N: int = 64):
func, expected = oob_store_legalize(M, N)
func, expected = _materialize_launch(func), _materialize_launch(expected)
mod = tvm.IRModule({func.attrs["global_symbol"]: func})
transformed = tl.transform.LegalizeSafeMemoryAccess()(mod)
tvm.ir.assert_structural_equal(
@@ -0,0 +1,337 @@
"""Tests for the target-neutral T.Kernel encoding and MaterializeKernelLaunch.
T.Kernel is traced before the Target is known, so it only records the grid
loops, thread-index placeholders and launch annotations; the backend pipeline
decides what the thread placeholders mean when it runs MaterializeKernelLaunch.
"""
import importlib
import inspect
import pytest
import tilelang as tl
import tilelang.language as T
import tilelang
import tilelang.testing
from tilelang import tvm
from tvm.tirx.stmt_functor import post_order_visit
def _collect(root, kind):
found = []
def _visit(node):
if isinstance(node, kind):
found.append(node)
post_order_visit(root.body if hasattr(root, "body") else root, _visit)
return found
def _launch_placeholders(func):
return [
stmt
for stmt in _collect(func, tvm.tirx.Bind)
if isinstance(stmt.value, tvm.tirx.Call) and str(stmt.value.op.name) == "tl.launch_thread_idx"
]
def _thread_extents(func):
"""{thread_tag: extent} of every thread_extent AttrStmt in `func`."""
extents = {}
for attr in _collect(func, tvm.tirx.AttrStmt):
if attr.attr_key == "thread_extent":
extents[str(attr.node.thread_tag)] = int(attr.value)
return extents
def _root_block(func):
blocks = [b for b in _collect(func, tvm.tirx.SBlock) if b.name_hint == "tilelang_root"]
assert len(blocks) == 1
return blocks[0]
def _materialize(func, target: str, **kwargs):
mod = tvm.IRModule.from_expr(func)
mod = tvm.tirx.transform.BindTarget(tvm.target.Target(target))(mod)
mod = tl.transform.MaterializeKernelLaunch(**kwargs)(mod)
return mod[func.attrs["global_symbol"]]
def _parallel_kernel(threads=None):
@T.prim_func
def main(A: T.Tensor((256,), "float32"), B: T.Tensor((256,), "float32")):
with T.Kernel(2, threads=threads) as bx:
for i in T.Parallel(128):
B[bx * 128 + i] = A[bx * 128 + i] + 1.0
return main
def _thread_indexed_kernel():
@T.prim_func
def main(A: T.Tensor((256,), "float32"), B: T.Tensor((256,), "float32")):
with T.Kernel(2, threads=128) as bx:
tx = T.get_thread_binding()
B[bx * 128 + tx] = A[bx * 128 + tx] + 1.0
return main
def test_traced_launch_records_grid_and_thread_placeholders():
func = _parallel_kernel()
grid = [f for f in _collect(func, tvm.tirx.For) if f.kind == tvm.tirx.ForKind.THREAD_BINDING]
assert [str(f.thread_binding.thread_tag) for f in grid] == ["blockIdx.x"]
placeholders = _launch_placeholders(func)
assert [p.var.name for p in placeholders] == ["tx", "ty", "tz"]
assert [int(p.value.args[0]) for p in placeholders] == [0, 1, 2]
# No threads= means no SIMT hint is recorded; the backend picks.
assert "tl.launch_threads" not in _root_block(func).annotations
assert _thread_extents(func) == {}
def test_traced_launch_records_requested_threads_as_annotation():
func = _parallel_kernel(threads=(64, 2))
threads = _root_block(func).annotations["tl.launch_threads"]
assert [int(x) for x in threads] == [64, 2, 1]
# Still only placeholders: the frontend does not bind threadIdx itself.
assert _thread_extents(func) == {}
def test_kernel_launch_annotations_are_recorded_on_the_root_block():
@T.prim_func
def main(A: T.Tensor((16,), "int32")):
with T.Kernel(1, threads=64, prelude="// hi", cluster_dims=2):
A[0] = 0
annotations = _root_block(main).annotations
assert [int(x) for x in annotations["tl.launch_threads"]] == [64, 1, 1]
assert [int(x) for x in annotations["cluster_dims"]] == [2, 1, 1]
assert str(annotations["pragma_import_c"]) == "// hi"
def test_kernel_rejects_unknown_launch_annotation():
"""A dialect's Kernel declares its launch annotations as explicit keyword
parameters, so a misspelled or foreign key fails at trace time."""
with pytest.raises(TypeError, match="unexpected keyword argument 'thread'"):
@T.prim_func
def typo(A: T.Tensor((16,), "int32")):
with T.Kernel(1, thread=128):
A[0] = 0
with pytest.raises(TypeError, match="unexpected keyword argument 'core_type'"):
@T.prim_func
def foreign(A: T.Tensor((16,), "int32")):
with T.Kernel(1, core_type="aiv"):
A[0] = 0
def _launch_annotations(kernel) -> set[str]:
return {name for name, p in inspect.signature(kernel).parameters.items() if p.kind is inspect.Parameter.KEYWORD_ONLY}
def test_each_dialect_declares_its_own_launch_annotations():
expected = {
"tilelang.language.common": set(),
"tilelang.cuda.language": {"threads", "prelude", "cluster_dims"},
"tilelang.rocm.language": {"threads", "prelude"},
"tilelang.metal.language": {"threads", "prelude"},
"tilelang.webgpu.language": {"threads"},
"tilelang.cpu.language": {"prelude"},
}
for module, keys in expected.items():
dialect = importlib.import_module(module)
assert _launch_annotations(dialect.Kernel) == keys, module
assert dialect.Kernel.__module__.startswith(module.removesuffix(".common")), module
# The default facade is the CUDA dialect.
assert T.Kernel is importlib.import_module("tilelang.cuda.language").Kernel
def test_cpu_dialect_kernel_has_no_threads():
from tilelang.cpu import language as Tcpu
with pytest.raises(TypeError, match="unexpected keyword argument 'threads'"):
@Tcpu.prim_func
def main(A: Tcpu.Tensor((16,), "int32")):
with Tcpu.Kernel(1, threads=128):
A[0] = 0
@Tcpu.prim_func
def ok(A: Tcpu.Tensor((16,), "int32")):
with Tcpu.Kernel(1, prelude="// cpu"):
A[0] = 0
assert str(_root_block(ok).annotations["pragma_import_c"]) == "// cpu"
def test_simt_binds_requested_threads():
func = _materialize(_parallel_kernel(threads=64), "cuda")
assert _thread_extents(func) == {"blockIdx.x": 2, "threadIdx.x": 64, "threadIdx.y": 1, "threadIdx.z": 1}
assert _launch_placeholders(func) == []
def test_simt_uses_backend_default_when_threads_omitted():
func = _materialize(_parallel_kernel(), "cuda", default_threads=256)
assert _thread_extents(func)["threadIdx.x"] == 256
func = _materialize(_parallel_kernel(), "cuda")
assert _thread_extents(func)["threadIdx.x"] == tl.transform.DEFAULT_SIMT_THREADS
def test_simt_requires_threads_when_backend_has_no_default():
with pytest.raises(Exception, match="did not specify threads="):
_materialize(_parallel_kernel(), "cuda", default_threads=None)
def test_simt_preserves_thread_var_identity():
"""The Var handed out by T.get_thread_binding() at trace time must be the
Var bound by the threadIdx.x thread_extent after materialization."""
func = _thread_indexed_kernel()
(placeholder,) = [p for p in _launch_placeholders(func) if p.var.name == "tx"]
lowered = _materialize(func, "cuda")
(attr,) = [a for a in _collect(lowered, tvm.tirx.AttrStmt) if str(a.node.thread_tag) == "threadIdx.x"]
assert attr.node.var.same_as(placeholder.var)
body_vars = [v for v in _collect(lowered, tvm.tirx.Var) if v.name == "tx"]
assert body_vars and all(v.same_as(placeholder.var) for v in body_vars)
def test_non_simt_drops_thread_placeholders():
func = _materialize(_parallel_kernel(threads=128), "c", lower_thread_binding=False)
assert _thread_extents(func) == {}
assert _launch_placeholders(func) == []
grid = [f for f in _collect(func, tvm.tirx.For) if f.loop_var.name == "bx"]
assert len(grid) == 1 and grid[0].kind == tvm.tirx.ForKind.SERIAL and int(grid[0].extent) == 2
assert not any(v.name in ("tx", "ty", "tz") for v in _collect(func, tvm.tirx.Var))
def test_non_simt_rejects_thread_index_use():
with pytest.raises(Exception, match="references thread index `tx`"):
_materialize(_thread_indexed_kernel(), "c", lower_thread_binding=False)
def test_get_thread_extent_requires_threads_at_trace_time():
with pytest.raises(ValueError, match="not known at trace time"):
@T.prim_func
def main(A: T.Tensor((16,), "int32")):
with T.Kernel(1):
A[0] = T.get_thread_extent()
def test_get_thread_extent_with_threads_at_trace_time():
@T.prim_func
def main(A: T.Tensor((16,), "int32")):
with T.Kernel(1, threads=(32, 4)):
A[0] = T.get_thread_extent(0) * T.get_thread_extent(1)
(store,) = _collect(main, tvm.tirx.BufferStore)
assert int(store.value) == 128
class _TraceFailure(Exception):
pass
def test_failed_trace_unwinds_launch_frames():
"""An exception inside T.Kernel must leave no stale launch frame behind,
otherwise the next trace sees the previous kernel's KernelLaunchFrame."""
with pytest.raises(_TraceFailure):
@T.prim_func
def failing(A: T.Tensor((16,), "int32")):
with T.Kernel(1, threads=128):
raise _TraceFailure()
assert T.KernelLaunchFrame.Current() is None
@tilelang.jit
def failing_jit(A):
A: T.Tensor[[16], T.int32]
with T.Kernel(1, threads=128):
raise _TraceFailure()
import torch
with pytest.raises(_TraceFailure):
failing_jit.get_tir(torch.zeros(16, dtype=torch.int32))
assert T.KernelLaunchFrame.Current() is None
def _cluster_kernel():
@T.prim_func
def main(A: T.Tensor((16,), "int32")):
with T.ClusterKernel(8, 4, threads=128, cluster_dims=2) as (bx, by):
A[0] = 0
return main
def test_cluster_dims_is_a_launch_annotation():
dims = _root_block(_cluster_kernel()).annotations["cluster_dims"]
assert [int(d) for d in dims] == [2, 1, 1]
def test_cluster_id_is_program_space_arithmetic():
"""Cluster identity is derived from the program index and cluster_dims at
trace time, so it needs no target-specific intrinsic."""
captured = {}
@T.prim_func
def main(A: T.Tensor((16,), "int32")):
with T.ClusterKernel(8, 4, threads=128, cluster_dims=2) as (bx, by):
captured["bx"], captured["by"] = bx, by
captured["ids"] = T.get_cluster_ids()
captured["dims"] = T.get_cluster_dims()
captured["size"] = T.get_cluster_size()
captured["extents"] = T.get_cluster_extents()
A[0] = T.get_cluster_id(0)
cx, cy = captured["ids"]
assert isinstance(cx, tvm.tirx.FloorDiv) and cx.a.same_as(captured["bx"]) and int(cx.b) == 2
# A unit cluster axis is the program index itself.
assert cy.same_as(captured["by"])
assert captured["dims"] == [2, 1, 1]
assert captured["size"] == 2
assert captured["extents"] == [4, 4, 1]
(store,) = _collect(main, tvm.tirx.BufferStore)
assert isinstance(store.value, tvm.tirx.FloorDiv)
def test_cluster_id_without_clusters_is_the_program_index():
captured = {}
@T.prim_func
def main(A: T.Tensor((16,), "int32")):
with T.Kernel(8) as bx:
captured["bx"] = bx
captured["id"] = T.get_cluster_id()
captured["dims"] = T.get_cluster_dims()
captured["extents"] = T.get_cluster_extents()
A[0] = 0
assert captured["id"].same_as(captured["bx"])
assert captured["dims"] == [1, 1, 1]
# Axes beyond the launched grid have a single cluster.
assert captured["extents"] == [8, 1, 1]
def test_cluster_dims_accepted_by_default_and_rejected_when_unsupported():
func = _materialize(_cluster_kernel(), "cuda")
assert "cluster_dims" in _root_block(func).annotations
with pytest.raises(Exception, match="`cluster_dims` is not supported on target `c`"):
_materialize(_cluster_kernel(), "c", lower_thread_binding=False, unsupported_annotations=["cluster_dims"])
# A launch without the annotation is unaffected by the rejection list.
_materialize(_parallel_kernel(), "c", lower_thread_binding=False, unsupported_annotations=["cluster_dims"])
if __name__ == "__main__":
tilelang.testing.main()
@@ -58,9 +58,9 @@ def _marker_line(marker: str) -> int:
def _make_vector_add():
@T.prim_func
def main(A: T.Tensor((1024,), "float32"), B: T.Tensor((1024,), "float32")):
with T.Kernel(1024):
tid = T.get_thread_binding()
B[tid] = A[tid] + 1.0 # span_marker_vadd_store
with T.Kernel(8) as bx:
for i in T.Parallel(128):
B[bx * 128 + i] = A[bx * 128 + i] + 1.0 # span_marker_vadd_store
return main
@@ -96,8 +96,9 @@ def _lower_with_recorder(func, target: str) -> _SpanCoverageRecorder:
return recorder
# Passes whose span propagation was fixed; they must never drop a span
# (statement *deletion* is fine — it also reduces the total).
# Passes whose span propagation was fixed; they must never strip a span from
# a statement or introduce a new statement without one (statement *deletion*
# is fine — it reduces the spanned and total counts alike).
_SPAN_SAFE_PASSES = {
"tl.MaterializeKernelLaunch",
"tl.AddWrapperForSingleBufStore",
@@ -112,11 +113,13 @@ _SPAN_SAFE_PASSES = {
def _assert_no_span_loss(recorder: _SpanCoverageRecorder):
prev_w = None
for name, w, _t in recorder.rows:
if prev_w is not None and name in _SPAN_SAFE_PASSES:
assert w >= prev_w, f"pass {name} dropped spans: {prev_w} -> {w}"
prev_w = w
prev = None
for name, w, t in recorder.rows:
if prev is not None and name in _SPAN_SAFE_PASSES:
prev_w, prev_t = prev
unspanned, prev_unspanned = t - w, prev_t - prev_w
assert unspanned <= prev_unspanned, f"pass {name} dropped spans: {prev_w}/{prev_t} spanned -> {w}/{t} spanned"
prev = (w, t)
def test_span_survives_lowering_cpu():
+11 -1
View File
@@ -12,7 +12,17 @@ from .annotations import ( # noqa: F401
)
from .copy_op import dual_copy # noqa: F401
from .gemm_op import blockscaled_gemm # noqa: F401
from .kernel import MixedKernel # noqa: F401
# Ascend owns its launch and its thread-scope accessors. These deliberately
# shadow the common surface imported above: `T.Kernel` here is the 1-D NPU core
# grid (no threads=), and `T.get_thread_binding()` resolves inside T.SimtVF.
from .kernel import ( # noqa: F401
Kernel,
MixedKernel,
get_thread_binding,
get_thread_bindings,
get_thread_extent,
get_thread_extents,
)
from .schedule_hint import PerCoreTask, Stage, Task, assume_no_conflict # noqa: F401
from .tile_schedule import ( # noqa: F401
AscendBaseTileScheduler,
+2 -2
View File
@@ -88,7 +88,7 @@ class SimtVFFrame(TIRFrame):
def __enter__(self):
super().__enter__()
from tilelang.language.kernel import SimtVFContext, push_simtvf_context
from .kernel import SimtVFContext, push_simtvf_context
ctx = SimtVFContext(
thread_vars=list(self.thread_vars),
@@ -98,7 +98,7 @@ class SimtVFFrame(TIRFrame):
return self
def __exit__(self, ptype, value, trace):
from tilelang.language.kernel import pop_simtvf_context
from .kernel import pop_simtvf_context
pop_simtvf_context()
super().__exit__(ptype, value, trace)
+171 -9
View File
@@ -1,14 +1,181 @@
"""Ascend NPU mixed-kernel (AIC + AIV) launch frame."""
"""Ascend NPU dialect of ``T.Kernel``.
Ascend owns its launch surface. The NPU launch is a 1-D grid of AI cores with
no SIMT thread domain at kernel scope, so this dialect's ``Kernel`` declares
``prelude`` and nothing else: passing ``threads=`` or ``cluster_dims=`` is
rejected by Python itself rather than by a runtime probe inside the shared
launch path. Thread domains are declared explicitly *inside* the kernel body by
``T.SimtVF(threads=...)`` (real threadIdx scopes) or ``T.SimdVF()`` (register
level, no threads); both emit their own thread scopes below the launch nest, so
the kernel-level ``tx/ty/tz`` placeholders are dropped by the Ascend pipeline.
"""
from __future__ import annotations
from tilelang import _ffi_api
from tilelang.jit.exceptions import JITNoBuilderError
import threading
from tvm import tirx
__all__ = ["MixedKernel"]
from tilelang import _ffi_api
from tilelang.jit.exceptions import JITNoBuilderError
from tilelang.language.kernel import (
FrameStack,
KernelLaunchFrame,
get_block_binding,
get_block_bindings,
get_block_extent,
get_block_extents,
kernel_launch_factory,
launch_kernel,
)
__all__ = [
"Kernel",
"MixedKernel",
"SimtVFContext",
"get_block_binding",
"get_block_bindings",
"get_block_extent",
"get_block_extents",
"get_thread_binding",
"get_thread_bindings",
"get_thread_extent",
"get_thread_extents",
"pop_simtvf_context",
"push_simtvf_context",
]
# ---------------------------------------------------------------------------
# SIMT thread scopes
#
# ``T.SimtVF`` owns a thread domain that is nested *inside* the kernel launch,
# so ``T.get_thread_binding()`` and friends have to resolve against the active
# SimtVF scope when there is one and against the launch frame otherwise. Both
# the scope state and the accessors are Ascend's: CUDA/ROCm/Metal declare their
# thread domain on the launch itself and never need the indirection. Keeping
# them here is what lets the shared ``tilelang.language.kernel`` stay free of
# backend branches.
# ---------------------------------------------------------------------------
class SimtVFContext:
"""Stores thread binding info for an active SimtVF scope."""
__slots__ = ("thread_vars", "thread_extents")
def __init__(self, thread_vars, thread_extents):
self.thread_vars = thread_vars
self.thread_extents = thread_extents
_simtvf_local = threading.local()
def _get_simtvf_stack() -> FrameStack:
if not hasattr(_simtvf_local, "simtvf_stack"):
_simtvf_local.simtvf_stack = FrameStack()
return _simtvf_local.simtvf_stack
def _get_current_simtvf() -> SimtVFContext | None:
stack = _get_simtvf_stack()
return stack.top() if stack else None
def push_simtvf_context(ctx: SimtVFContext):
"""Enter a SimtVF thread scope, making its thread vars the current ones."""
_get_simtvf_stack().push(ctx)
def pop_simtvf_context():
"""Leave the innermost SimtVF thread scope."""
_get_simtvf_stack().pop()
def get_thread_binding(dim: int = 0):
"""Returns the thread binding for the given dimension.
Inside a ``T.SimtVF`` block this is the SimtVF thread var; otherwise it is
the kernel launch's placeholder, which only has a meaning if the pipeline
materialized SIMT threads (Ascend never does at kernel scope).
"""
simtvf = _get_current_simtvf()
if simtvf is not None:
return simtvf.thread_vars[dim]
assert KernelLaunchFrame.Current() is not None, "KernelLaunchFrame is not initialized"
return KernelLaunchFrame.Current().get_thread_binding(dim)
def get_thread_bindings() -> list:
"""Returns all three thread bindings."""
simtvf = _get_current_simtvf()
if simtvf is not None:
return list(simtvf.thread_vars)
assert KernelLaunchFrame.Current() is not None, "KernelLaunchFrame is not initialized"
return KernelLaunchFrame.Current().get_thread_bindings()
def get_thread_extent(dim: int = 0) -> int:
"""Returns the thread extent for the given dimension."""
simtvf = _get_current_simtvf()
if simtvf is not None:
return simtvf.thread_extents[dim]
assert KernelLaunchFrame.Current() is not None, "KernelLaunchFrame is not initialized"
return KernelLaunchFrame.Current().get_thread_extent(dim)
def get_thread_extents() -> list:
"""Returns all three thread extents."""
simtvf = _get_current_simtvf()
if simtvf is not None:
return list(simtvf.thread_extents)
assert KernelLaunchFrame.Current() is not None, "KernelLaunchFrame is not initialized"
return KernelLaunchFrame.Current().get_thread_extents()
# ---------------------------------------------------------------------------
# Launch frames
# ---------------------------------------------------------------------------
@kernel_launch_factory
def Kernel(
*blocks: int | tirx.PrimExpr,
prelude: str | None = None,
) -> KernelLaunchFrame:
"""Construct a kernel launch frame for Ascend: a 1-D grid of AI cores.
The grid becomes the NPU core index (``blockIdx.x``). There is no SIMT
thread domain at this scope, so this dialect has no ``threads`` parameter:
``with T.Kernel(N) as bx`` yields one program index, and ``bx`` is iterable
as ``(bx,)``. Use ``T.SimtVF(threads=...)`` inside the body to run
thread-parallel code, or ``T.MixedKernel`` for an AIC+AIV mixed kernel.
Parameters
----------
*blocks : int | PrimExpr
Extent of the 1-D core grid. Exactly one dimension is allowed; a
multi-dimensional launch is rejected here rather than silently
flattened downstream.
prelude : str, optional
AscendC source injected before the generated kernel, e.g. ``#include``
lines or helper functions.
Examples
--------
.. code-block:: python
with T.Kernel(NUM_CORES) as bx:
with T.SimtVF(threads=128):
for i in T.Parallel(128):
out[bx * 128 + i] = x[bx * 128 + i] * 2.0
"""
if len(blocks) != 1:
raise ValueError(f"Ascend targets a 1-D core grid: T.Kernel(N) takes exactly one grid extent. Got {len(blocks)}-D: {blocks}.")
return launch_kernel(blocks, prelude=prelude)
@kernel_launch_factory
def MixedKernel(
*blocks: int | tirx.PrimExpr,
sids: int = 2,
@@ -46,14 +213,10 @@ def MixedKernel(
...
"""
from tilelang.language.eager.builder import Builder
from tilelang.ascend.target import check_ascend_availability
if Builder.current() is None:
raise JITNoBuilderError("T.MixedKernel() can only be used inside @tilelang.jit or @T.prim_func context. No Builder is available.")
if not check_ascend_availability():
raise RuntimeError("T.MixedKernel() requires an Ascend NPU environment (torch.npu.is_available() must return True).")
if len(blocks) != 1:
raise ValueError(f"T.MixedKernel() only supports 1-D block grid. Got {len(blocks)}-D: {blocks}")
@@ -61,7 +224,6 @@ def MixedKernel(
raise ValueError(f"T.MixedKernel() sids must be 1 or 2. Got {sids}")
attrs: dict = {}
attrs["tilelang.is_npu_kernel_frame"] = True
if prelude is not None:
attrs["pragma_import_c"] = prelude
+14 -6
View File
@@ -20,12 +20,20 @@ from . import transform as ascend_transform
def AscendPassPipelineBody(mod: IRModule, target: Target) -> IRModule:
mod = tirx.transform.BindTarget(target)(mod)
# Materialize the target-neutral kernel-launch nest (thread_binding For
# loops emitted by T.Kernel) into thread_extent AttrStmts. Ascend's NPU
# launch is a 1-D blockIdx.x grid with no threadIdx, so SIMT-style
# materialization (lower_thread_binding=True) reproduces the previous
# LaunchThread(blockIdx.x) behavior.
mod = tilelang.transform.MaterializeKernelLaunch()(mod)
# Materialize the target-neutral kernel-launch nest emitted by T.Kernel.
# Ascend's NPU launch is a real 1-D blockIdx.x core grid, so the grid loop
# becomes a thread_extent AttrStmt; there is no threadIdx at kernel scope,
# so the `tx/ty/tz = tl.launch_thread_idx(...)` placeholders are dropped.
# Thread domains only exist inside T.SimtVF, which emits its own thread
# scopes below the launch nest. `cthread` is Ascend's sub-block-id launch
# dimension, emitted by T.MixedKernel and read back by the Ascend codegen.
mod = tilelang.transform.MaterializeKernelLaunch(
lower_grid_binding=True,
lower_thread_binding=False,
default_threads=None,
unsupported_annotations=["cluster_dims"],
launch_dim_tags=["cthread"],
)(mod)
pass_ctx = tilelang.transform.get_pass_context()
if should_force_let_inline(pass_ctx=pass_ctx):
+1 -1
View File
@@ -174,7 +174,7 @@ freely compose that surface with extensions under
`tilelang/<backend>/language`:
```python
from tilelang import language as T # common + CUDA compatibility facade
from tilelang import language as T # common + CUDA compatibility facade
from tilelang.cuda import language as T # common + CUDA extensions
from tilelang.rocm import language as T # common + ROCm extensions
```
+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
@@ -865,13 +865,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
+2 -8
View File
@@ -153,7 +153,7 @@ class CUDA(TileDevice):
self.warp_size = device.warp_size
...
self.transaction_size = [32, 128] # bytes
self.bandwidth = [750, 12080] # MB/s, approximate
self.bandwidth = [750, 12080] # MB/s, approximate
self.available_tensor_instructions = None
def get_avaliable_tensorintrin_shapes(self):
@@ -173,13 +173,7 @@ One of Carver’s main benefits is its adaptability. Here are a examples for tri
Given a Carver hint like:
```python
{
'block': [32, 64],
'warp': [16, 32],
'rstep': [128],
'use_tc': True,
'vectorize': {'A_reindex': 8, 'B_reindex': 8}
}
{"block": [32, 64], "warp": [16, 32], "rstep": [128], "use_tc": True, "vectorize": {"A_reindex": 8, "B_reindex": 8}}
```
You might interpret this in **Triton** as:
- `block_m = 32, block_n = 64, block_k = 128`
+6 -3
View File
@@ -5,7 +5,10 @@ from __future__ import annotations
from tilelang.language.common import * # noqa: F401,F403
from tilelang.language.common import __all__ as _COMMON_ALL
__tilelang_dialect__ = "cpu"
__all__ = tuple(_COMMON_ALL)
from .kernel import * # noqa: F401,F403
from .kernel import __all__ as _KERNEL_ALL
del _COMMON_ALL
__tilelang_dialect__ = "cpu"
__all__ = tuple(dict.fromkeys((*_COMMON_ALL, *_KERNEL_ALL)))
del _COMMON_ALL, _KERNEL_ALL
+41
View File
@@ -0,0 +1,41 @@
"""CPU dialect of ``T.Kernel``: the common launch plus CPU launch annotations."""
from __future__ import annotations
from tvm import tirx
from tilelang.language.kernel import KernelLaunchFrame, kernel_launch_factory, launch_kernel
__all__ = ["Kernel"]
@kernel_launch_factory
def Kernel(
*blocks: int | tirx.PrimExpr,
prelude: str | None = None,
) -> KernelLaunchFrame:
"""Construct a kernel launch frame for CPU: a grid of tile programs.
The grid becomes the outer loop nest of the generated function and each
tile program runs as a plain serial body; ``T.Parallel`` loops are lowered
to serial loops. There are no SIMT threads, so this dialect has no
``threads`` and ``T.get_thread_binding()`` is rejected at compile time.
Parameters
----------
*blocks : int | PrimExpr
Grid extent along each axis (1-3 dimensions). The launch yields one
program index per axis.
prelude : str, optional
C source injected before the generated kernel, e.g. ``#include`` lines
or helper functions.
Examples
--------
.. code-block:: python
with T.Kernel(T.ceildiv(N, 128)) as bx:
for i in T.Parallel(128):
...
"""
return launch_kernel(blocks, prelude=prelude)
+3 -1
View File
@@ -14,7 +14,9 @@ from tilelang.backend.pass_pipeline.pipeline_utils import (
def CPUPassPipelineBody(mod: IRModule, target: Target) -> IRModule:
mod = tirx.transform.BindTarget(target)(mod)
mod = tilelang.transform.MaterializeKernelLaunch(lower_thread_binding=False)(mod)
mod = tilelang.transform.MaterializeKernelLaunch(
lower_grid_binding=False, lower_thread_binding=False, default_threads=None, unsupported_annotations=["cluster_dims"]
)(mod)
pass_ctx = tilelang.transform.get_pass_context()
if should_force_let_inline():
+4
View File
@@ -46,6 +46,9 @@ from tilelang.language.builtin import ( # noqa: F401
from tilelang.language.copy_op import copy_cluster, tma_copy, tma_gather4, tma_gather4_bytes, tma_scatter4 # noqa: F401
from tilelang.language.kernel import ClusterKernel, CUDASourceCodeKernel # noqa: F401
# The CUDA dialect's T.Kernel shadows the target-neutral one from common: same
# launch, plus the CUDA launch annotations (threads, prelude, cluster_dims).
from .kernel import Kernel # noqa: F401
from .cluster import * # noqa: F401,F403
from .cluster import __all__ as _CLUSTER_ALL
from .intrinsics import * # noqa: F401,F403
@@ -66,6 +69,7 @@ from .warpgroup import __all__ as _WARPGROUP_ALL
_CUDA_API_ALL = (
"ClusterKernel",
"CUDASourceCodeKernel",
"Kernel",
"alloc_cluster_barrier",
"alloc_descriptor",
"alloc_tmem",
+55
View File
@@ -0,0 +1,55 @@
"""CUDA dialect of ``T.Kernel``: the common launch plus CUDA launch annotations."""
from __future__ import annotations
from tvm import tirx
from tilelang.language.kernel import KernelLaunchFrame, kernel_launch_factory, launch_kernel
__all__ = ["Kernel"]
@kernel_launch_factory
def Kernel(
*blocks: int | tirx.PrimExpr,
threads: int | list[int] | tuple[int, ...] | None = None,
prelude: str | None = None,
cluster_dims: int | tuple[int, int, int] | list[int] | None = None,
) -> KernelLaunchFrame:
"""Construct a kernel launch frame for CUDA: a grid of thread blocks.
Code inside the launch operates at the block level: ``T.Parallel``,
``T.copy`` and friends are mapped onto threads by the compiler.
``T.get_thread_binding()`` exposes ``threadIdx`` for thread-level code.
The keyword arguments are recorded at trace time and materialized by the
CUDA pipeline once the target is known.
Parameters
----------
*blocks : int | PrimExpr
Grid extent along each axis (1-3 dimensions, ``gridDim.(x|y|z)``). The
launch yields one block index per axis (``blockIdx.(x|y|z)``).
threads : int | list[int] | tuple[int, ...], optional
Threads per block: a count for ``blockDim.x`` or up to three
per-dimension extents for ``blockDim.(x|y|z)``. Defaults to 128 when
omitted.
prelude : str, optional
CUDA source injected before the generated kernel, e.g. ``#include``
lines or helper functions.
cluster_dims : int | tuple[int, int, int] | list[int], optional
Thread block cluster shape (SM90+). ``2`` or ``(2, 1, 1)`` launches
2-CTA clusters via ``cudaLaunchKernelEx``. ``T.ClusterKernel`` is the
same launch with a required ``cluster_dims``.
Examples
--------
.. code-block:: python
with T.Kernel(T.ceildiv(N, 128), threads=128) as bx:
...
with T.Kernel(grid_x, grid_y, threads=(64, 2)) as (bx, by):
tx, ty = T.get_thread_bindings()
...
"""
return launch_kernel(blocks, threads=threads, prelude=prelude, cluster_dims=cluster_dims)
+3 -3
View File
@@ -67,7 +67,7 @@ def _module_has_shared_barrier(mod: IRModule) -> bool:
def CUDAPassPipelineBodyPrologue(mod: IRModule, target: Target) -> IRModule:
mod = tirx.transform.BindTarget(target)(mod)
mod = tilelang.transform.MaterializeKernelLaunch()(mod)
mod = tilelang.transform.MaterializeKernelLaunch(lower_thread_binding=True, default_threads=128)(mod)
# Record body-bound global bases before optional let inlining obscures
# their provenance. CopyAnalysis consumes this marker for every lowering
# path, independently of whether warp specialization is enabled.
@@ -100,11 +100,11 @@ def CUDAPassPipelineBodyPrologue(mod: IRModule, target: Target) -> IRModule:
# @CUDA-specific
# Tile-level warp specialization: runs before layout inference so that
# producer/consumer split happens at the high-level tile-op IR.
# AutoSchedule: derive a WSSchedule for kernels without one, using the scheduler named by tl.cuda_auto_schedule (opt-in; no-op when unset).
# AutoWarpSpecialization: derive a WSSchedule for kernels without one, using the scheduler named by tl.enable_auto_warp_specialization (opt-in; no-op when unset).
# MaterializeWSSchedule: Materialize a user-provided warp-specialization schedule (T.annotate_ws_schedule).
# ProducerConsumerWarpSpecialized: The pass classifies copy ops as TMA/cp.async/sync inline. Shared buffers are multi-versioned internally only for functions where the WS transformation actually applies.
if allow_warp_specialized(target=target):
mod = tilelang.cuda.transform.AutoSchedule()(mod)
mod = tilelang.cuda.transform.AutoWarpSpecialization()(mod)
mod = tilelang.cuda.transform.MaterializeWSSchedule()(mod)
mod = tilelang.cuda.transform.ProducerConsumerWarpSpecialized()(mod)
+4 -4
View File
@@ -3,9 +3,9 @@
from .. import _ffi_api
def AutoSchedule():
def AutoWarpSpecialization():
"""Derive a warp-specialization schedule with the scheduler named by
the ``tl.cuda_auto_schedule`` pass config (currently
the ``tl.enable_auto_warp_specialization`` pass config (currently
``"role_based"``); a no-op when the config is unset.
Eligible kernels gain stable ``tl.ws_op_id`` markers and a typed
@@ -18,7 +18,7 @@ def AutoSchedule():
fpass : tvm.transform.Pass
The result pass
"""
return _ffi_api.AutoSchedule() # type: ignore
return _ffi_api.AutoWarpSpecialization() # type: ignore
def AnnotateDeviceBoundTmaCopies():
@@ -197,7 +197,7 @@ def PersistThreadblock():
__all__ = [
"AnnotateDeviceBoundTmaCopies",
"AutoSchedule",
"AutoWarpSpecialization",
"AnnotateWarpGroupRegAlloc",
"FuseMBarrierArriveExpectTx",
"InjectFenceProxy",
-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")
+4
View File
@@ -11,4 +11,8 @@ from __future__ import annotations
from tilelang.ascend.language import * # noqa: F401,F403
from tilelang.ascend.language import __all__ as __all__ # noqa: F401
# Imported by name so static type checkers resolve the Ascend-typed launch
# signature through this facade (they cannot evaluate the dynamic __all__).
from tilelang.ascend.language import Kernel # noqa: F401
__tilelang_dialect__ = "ascend"
+12
View File
@@ -42,6 +42,12 @@ from .kernel import (
get_block_bindings, # noqa: F401
get_block_extent, # noqa: F401
get_block_extents, # noqa: F401
get_cluster_dims, # noqa: F401
get_cluster_size, # noqa: F401
get_cluster_id, # noqa: F401
get_cluster_ids, # noqa: F401
get_cluster_extent, # noqa: F401
get_cluster_extents, # noqa: F401
)
from .allocate import (
alloc_var, # noqa: F401
@@ -265,6 +271,12 @@ _LOCAL_EXPORTS = (
"get_block_bindings",
"get_block_extent",
"get_block_extents",
"get_cluster_dims",
"get_cluster_extent",
"get_cluster_extents",
"get_cluster_id",
"get_cluster_ids",
"get_cluster_size",
"get_lane_idx",
"get_let_value",
"get_thread_binding",
+14 -4
View File
@@ -621,11 +621,21 @@ class DSLMutator(ast.NodeTransformer):
is_kernel_ctx = False
for expr in node.items:
cexpr = expr.context_expr
if isinstance(cexpr, ast.Call) and isinstance(cexpr.func, ast.Attribute) and cexpr.func.attr in ("Kernel", "ClusterKernel"):
eval_res = self._try_eval(cexpr.func)
from tilelang.language.kernel import ClusterKernel, Kernel
if isinstance(cexpr, ast.Call) and isinstance(cexpr.func, (ast.Attribute, ast.Name)):
# Only resolve plain names and module attribute chains, so no
# factory expression such as make_scope().context() is executed
# at rewrite time.
root = cexpr.func
while isinstance(root, ast.Attribute):
root = root.value
if not isinstance(root, ast.Name):
continue
# Every dialect's Kernel (and ClusterKernel) is marked as a launch
# factory; identity against one implementation would miss the
# others, and aliases such as `K = T.Kernel`.
from tilelang.language.kernel import is_kernel_launch_factory
if eval_res is Kernel or eval_res is ClusterKernel:
if is_kernel_launch_factory(self._try_eval(cexpr.func)):
is_kernel_ctx = True
break
node = self.generic_visit(node)
+7 -1
View File
@@ -395,7 +395,13 @@ class Builder(BaseBuilder):
self.current_file,
self.current_line,
)
yield self.enter_frame(frame)
try:
yield self.enter_frame(frame)
except BaseException as exc:
# Unwind Python frame state without finalizing incomplete TIR.
while len(self.frames) > pop_idx:
self.frames.pop().__exit__(type(exc), exc, exc.__traceback__)
raise
if self._spans_enabled:
# Flush leaf stmts of the frame being exited before its __exit__
# moves them into the produced node.
+217 -189
View File
@@ -3,12 +3,12 @@
from __future__ import annotations
from collections import deque
import os
from typing import Any
from tvm import tirx
from tvm.tirx import Var
from tvm.tirx.script.builder import evaluate as T_evaluate
from tvm.tirx.script.builder.frame import TIRFrame
from tvm.tirx.script.builder.frame import SBlockFrame
from tvm.target import Target
from tvm.ffi import register_object
from tilelang import _ffi_api
from tilelang.jit.exceptions import JITNoBuilderError
@@ -85,38 +85,6 @@ def _get_current_stack() -> FrameStack:
return _local.kernel_launch_frame_stack
class SimtVFContext:
"""Stores thread binding info for an active SimtVF scope."""
__slots__ = ("thread_vars", "thread_extents")
def __init__(self, thread_vars, thread_extents):
self.thread_vars = thread_vars
self.thread_extents = thread_extents
_simtvf_local = threading.local()
def _get_simtvf_stack() -> FrameStack:
if not hasattr(_simtvf_local, "simtvf_stack"):
_simtvf_local.simtvf_stack = FrameStack()
return _simtvf_local.simtvf_stack
def _get_current_simtvf() -> SimtVFContext | None:
stack = _get_simtvf_stack()
return stack.top() if stack else None
def push_simtvf_context(ctx: SimtVFContext):
_get_simtvf_stack().push(ctx)
def pop_simtvf_context():
_get_simtvf_stack().pop()
def _normalize_bindings(bindings: list[Var]) -> Var | list[Var]:
"""
Return a bare Var when we only have a single binding so that users may write either
@@ -130,21 +98,25 @@ def _normalize_bindings(bindings: list[Var]) -> Var | list[Var]:
def _normalize_threads(
threads: int | list[int] | tuple | None,
) -> list[int]:
) -> list[int] | None:
"""Normalize a thread-block specification into a 3-D extent list.
Args:
threads: A thread count, a per-dimension extent list/tuple, or None for the default.
threads: A thread count, a per-dimension extent list/tuple, or None to
leave the choice to the backend.
Returns:
The extents as ``[x, y, z]``, padding missing dimensions with 1.
The extents as ``[x, y, z]``, padding missing dimensions with 1, or
None when no thread count was requested. The frontend does not pick a
default: the thread count is a SIMT launch hint whose default (if any)
belongs to the backend that materializes the launch.
Raises:
ValueError: If ``threads`` has an unsupported type, or any concrete extent
is not positive.
"""
if threads is None:
threads = 128 # default thread number
return None
if isinstance(threads, int):
normalized = [threads, 1, 1]
@@ -179,30 +151,23 @@ def _normalize_cluster_dims(
return None if cluster_dims == [1, 1, 1] else cluster_dims
def _current_target_is_ascend() -> bool:
current_target = Target.current(allow_none=True)
if current_target is None:
return False
try:
from tilelang.ascend.target import target_is_ascend
return bool(target_is_ascend(current_target))
except Exception:
return current_target.kind.name == "ascend"
@register_object("tl.KernelLaunchFrame")
class KernelLaunchFrame(TIRFrame):
"""
KernelLaunchFrame is a custom TIRFrame that manages block/thread indices
and handles the entry and exit of the kernel launch scope.
Grid (program index) vars are bound by the frame itself. Thread vars are
placeholders: they have an identity so the body can reference them, but
their extent is only known once a backend materializes the launch. Thread
extents are therefore available at trace time only when ``threads=`` was
passed to :func:`Kernel`.
"""
def __enter__(self) -> Var | list[Var]:
"""
Enters the KernelLaunchFrame scope and pushes this frame onto the stack.
Returns one Var if we detect exactly 5 frames (meaning there is a single
block dimension), or a list of Vars otherwise.
Returns one Var for a single grid dimension, or a list of Vars otherwise.
"""
super().__enter__()
_get_current_stack().push(self)
@@ -210,20 +175,7 @@ class KernelLaunchFrame(TIRFrame):
last_block_frame = self.frames[-1]
assert isinstance(last_block_frame, SBlockFrame), f"Last frame must be a block frame, got {last_block_frame}"
maybe_cpu = last_block_frame.annotations.get("tilelang.is_cpu_kernel_frame", False)
maybe_npu = last_block_frame.annotations.get("tilelang.is_npu_kernel_frame", False)
# All launch dimensions are target-neutral thread_binding For frames
# (regular T.Kernel and T.MixedKernel alike), so the loop var is
# frame.vars[0].
if maybe_cpu or maybe_npu:
# CPU/NPU kernels have no threadIdx frames; only the trailing block
# frame (with attributes) follows the grid frames.
return _normalize_bindings([frame.vars[0] for frame in self.frames[0:-1]])
else:
# GPU: exclude the last 4 frames (threadIdx.x/y/z and the block frame
# with attributes).
return _normalize_bindings([frame.vars[0] for frame in self.frames[0:-4]])
return _normalize_bindings(list(self.grid_vars))
def __exit__(self, ptype, value, trace):
"""
@@ -248,9 +200,11 @@ class KernelLaunchFrame(TIRFrame):
"""
Returns the block extent for the given dimension.
dim=0 corresponds to blockIdx.x, dim=1 to blockIdx.y, and dim=2 to blockIdx.z.
Grid axes that were not launched have extent 1.
"""
iter_var = self.frames[dim].doms[0]
return int(iter_var.extent)
if dim >= len(self.grid_extents):
return 1
return int(self.grid_extents[dim])
def get_block_extents(self) -> list[int]:
"""
@@ -262,9 +216,17 @@ class KernelLaunchFrame(TIRFrame):
"""
Returns the thread extent for the given dimension.
dim=0 corresponds to threadIdx.x, dim=1 to threadIdx.y, and dim=2 to threadIdx.z.
Raises:
ValueError: If the kernel was launched without ``threads=``. The
extent is then chosen by the backend and is not known at trace time.
"""
iter_var = self.frames[-4 + dim].doms[0]
return int(iter_var.extent)
if self.thread_extents is None:
raise ValueError(
"The thread extent is not known at trace time: T.Kernel(...) was called without "
"threads=. Pass threads= explicitly when the kernel body needs the thread-block size."
)
return int(self.thread_extents[dim])
def get_thread_extents(self) -> list[int]:
"""
@@ -277,14 +239,14 @@ class KernelLaunchFrame(TIRFrame):
Returns the thread binding for the given dimension.
dim=0 corresponds to threadIdx.x, dim=1 to threadIdx.y, and dim=2 to threadIdx.z.
"""
return self.frames[-4 + dim].vars[0]
return self.thread_vars[dim]
def get_thread_bindings(self) -> list[Var]:
"""
Returns the thread binding for the given dimension.
dim=0 corresponds to threadIdx.x, dim=1 to threadIdx.y, and dim=2 to threadIdx.z.
"""
return [frame.vars[0] for frame in self.frames[-4:-1]]
return list(self.thread_vars)
def get_num_threads(self) -> int:
"""
@@ -300,27 +262,94 @@ class KernelLaunchFrame(TIRFrame):
Returns the block binding for the given dimension.
dim=0 corresponds to blockIdx.x, dim=1 to blockIdx.y, and dim=2 to blockIdx.z.
"""
return self.frames[dim].vars[0]
return self.grid_vars[dim]
def get_block_bindings(self) -> list[Var]:
"""
Returns all three block bindings.
"""
return [frame.vars[0] for frame in self.frames[0:-4]]
return list(self.grid_vars)
def get_launch_annotation(self, key: str, default=None):
"""
Returns the launch annotation ``key`` recorded by T.Kernel (e.g. ``cluster_dims``),
or ``default`` when it was not given.
"""
annotations = self.frames[-1].annotations
if annotations is None or key not in annotations:
return default
return annotations[key]
def get_cluster_dims(self) -> list[int]:
"""
Returns the cluster dimensions as ``[x, y, z]``. A launch without
``cluster_dims`` has clusters of a single program, i.e. ``[1, 1, 1]``.
"""
dims = self.get_launch_annotation("cluster_dims")
if dims is None:
return [1, 1, 1]
dims = [int(d) for d in dims]
return dims + [1] * (3 - len(dims))
def get_cluster_size(self) -> int:
"""
Returns the number of programs per cluster (product of the cluster dimensions).
"""
size = 1
for dim in self.get_cluster_dims():
size *= dim
return size
def get_cluster_id(self, dim: int = 0) -> Var | tirx.PrimExpr:
"""
Returns the index of the cluster the current program belongs to along
``dim``, in program-space arithmetic: ``block_id // cluster_dims[dim]``.
A cluster is a ``cluster_dims``-shaped tile of the grid, so this is the
same on every target (clusterIdx on CUDA, a group of consecutive
programs elsewhere) and stays consistent with threadblock swizzling,
which permutes the grid at cluster granularity.
"""
if dim >= len(self.grid_vars):
return tirx.IntImm("int32", 0)
block = self.grid_vars[dim]
cluster_dim = self.get_cluster_dims()[dim]
if cluster_dim == 1:
return block
return tirx.floordiv(block, tirx.IntImm(block.dtype, cluster_dim))
def get_cluster_ids(self) -> list[Var | tirx.PrimExpr]:
"""
Returns the cluster index along every launched grid axis.
"""
return [self.get_cluster_id(dim) for dim in range(len(self.grid_vars))]
def get_cluster_extent(self, dim: int = 0) -> int:
"""
Returns the number of clusters along ``dim``: ``ceil(grid_extent / cluster_dims[dim])``.
"""
cluster_dim = self.get_cluster_dims()[dim]
return -(-self.get_block_extent(dim) // cluster_dim)
def get_cluster_extents(self) -> list[int]:
"""
Returns the number of clusters along all three dimensions.
"""
return [self.get_cluster_extent(dim) for dim in range(3)]
@property
def blocks(self) -> list[Var]:
"""
Returns the block indices from the topmost frame.
"""
return [frame.vars[0] for frame in self.frames[0:-4]]
return list(self.grid_vars)
@property
def threads(self) -> list[Var]:
"""
Returns the thread indices from the topmost frame.
"""
return [frame.vars[0] for frame in self.frames[-4:-1]]
return list(self.thread_vars)
@property
def num_threads(self) -> int:
@@ -330,65 +359,51 @@ class KernelLaunchFrame(TIRFrame):
return self.get_num_threads()
def Kernel(
*blocks: int | tirx.PrimExpr,
threads: int | list[int] | tuple | None = None,
cluster_dims: int | tuple[int, int, int] | list[int] | None = None,
is_cpu: bool = False,
# ---------------------------------------------------------------------------
# Launch annotations
#
# T.Kernel(*grid) is the launch every target shares. Everything a backend may
# additionally need (thread count, clusters, ...) is a *launch annotation*: the
# frontend records it verbatim and the backend interprets it once the target
# is known (MaterializeKernelLaunch). Which annotations exist is declared per
# language dialect as the explicit keyword parameters of its own `Kernel`
# (see tilelang/<backend>/language/kernel.py), so `tilelang.cuda.language.Kernel`
# shows, autocompletes and accepts exactly what CUDA understands. Every
# dialect's `Kernel` funnels into `launch_kernel` below.
# ---------------------------------------------------------------------------
_KERNEL_LAUNCH_FACTORY_ATTR = "__tilelang_kernel_launch__"
def kernel_launch_factory(func):
"""Mark ``func`` as a launch factory: a callable used as ``with func(...)``
to open a kernel launch. Every dialect's ``Kernel`` (and ``ClusterKernel``)
carries this mark so the eager JIT rewriter can find the launch regardless
of which dialect or alias the user went through."""
setattr(func, _KERNEL_LAUNCH_FACTORY_ATTR, True)
return func
def is_kernel_launch_factory(obj) -> bool:
return getattr(obj, _KERNEL_LAUNCH_FACTORY_ATTR, False) is True
def launch_kernel(
blocks: tuple[int | tirx.PrimExpr, ...],
*,
threads: int | list[int] | tuple[int, ...] | None = None,
prelude: str | None = None,
):
"""Tools to quickly construct a kernel launch frame.
cluster_dims: int | tuple[int, int, int] | list[int] | None = None,
**annotations: Any,
) -> KernelLaunchFrame:
"""Shared implementation behind every dialect's ``T.Kernel``.
The launch nest is emitted in a target-neutral form (thread_binding
For loops); each backend pipeline materializes it via
MaterializeKernelLaunch. Backends without SIMT (e.g. CPU) simply
ignore the thread extents at compile time, so the same kernel can be
compiled for any target.
Parameters
----------
blocks : int
A list of extent, can be 1-3 dimension, representing gridDim.(x|y|z)
threads : int
A integer representing blockDim.x
Or a list of integers representing blockDim.(x|y|z)
if the value is -1, we skip the threadIdx.x binding.
cluster_dims : int | tuple[int, int, int] | list[int] | None
The cluster dimensions for SM90+ cluster launch.
For example, use 2 or (2, 1, 1) to create 2-CTA clusters.
When specified, the kernel will be launched using cudaLaunchKernelEx
with cudaLaunchAttributeClusterDimension.
is_cpu : bool
Whether the kernel is running on CPU.
is_ascend : bool
Explicitly force Ascend kernel-frame semantics when target inference is
not active.
prelude : str
The import c code of the kernel,
will be injected before the generated kernel code.
Returns
-------
res : Tuple[frame.LaunchThreadFrame]
The result LaunchThreadFrame.
Examples
--------
Create a 1-D CUDA kernel launch and unpack the single block index:
.. code-block:: python
with T.Kernel(T.ceildiv(N, 128), threads=128) as bx:
# bx is the blockIdx.x binding (also iterable as (bx,))
...
Launch a 2-D grid while requesting two thread dimensions:
.. code-block:: python
with T.Kernel(grid_x, grid_y, threads=(64, 2)) as (bx, by):
tx, ty = T.get_thread_bindings()
...
The well-known launch annotations are normalized here; any other keyword
a dialect forwards is recorded verbatim on the launch block for that
backend's pipeline to consume. Dialects, not this function, decide which
keywords exist: they only forward what their own ``Kernel`` signature
declares.
"""
# In eager mode, we construct AST directly without prim_func,
# so there must be a Builder available. If not, this function
@@ -399,44 +414,47 @@ def Kernel(
if Builder.current() is None:
raise JITNoBuilderError("T.Kernel() can only be used inside @tilelang.jit or @T.prim_func context. No Builder is available.")
from tilelang.ascend.target import check_ascend_availability
is_ascend = check_ascend_availability()
attrs: dict = {}
if is_ascend and is_cpu:
raise ValueError("Uncertain backend behavior for `is_cpu=True` and `is_ascend=True`")
if (is_ascend or _current_target_is_ascend()) and not is_cpu:
if threads is not None:
raise ValueError(
"Ascend backend does not support `threads=` in `T.Kernel(...)`. Use `T.SimtVF(threads=...)` to define thread domains."
)
if len(blocks) != 1:
raise ValueError(f"Ascend backend only supports 1-D grid in `T.Kernel(N)`. Got {len(blocks)}-D grid.")
attrs["tilelang.is_npu_kernel_frame"] = True
elif not is_cpu and threads is None:
# Keep backward compatibility when target is not Ascend (or unknown).
threads = 128 # default thread number
if threads is None:
normalized_threads = None
else:
normalized_threads = _normalize_threads(threads)
if is_cpu:
attrs["tilelang.is_cpu_kernel_frame"] = True
if prelude is not None:
attrs["pragma_import_c"] = prelude
cluster_dims = _normalize_cluster_dims(cluster_dims)
if cluster_dims is not None:
attrs["cluster_dims"] = cluster_dims
for key, value in annotations.items():
if value is not None:
attrs[key] = value
return _ffi_api.KernelLaunch(blocks, normalized_threads, attrs)
return _ffi_api.KernelLaunch(blocks, _normalize_threads(threads), attrs)
@kernel_launch_factory
def Kernel(*blocks: int | tirx.PrimExpr) -> KernelLaunchFrame:
"""Construct a kernel launch frame: a grid of tile programs.
This is the target-neutral launch: the part every backend shares. Backend
dialects offer their own ``T.Kernel`` with the launch annotations that
backend understands, e.g. ``tilelang.cuda.language.Kernel(..., threads=128)``;
``tilelang.language`` is the CUDA dialect.
Parameters
----------
*blocks : int | PrimExpr
Extent of the grid along each axis (1-3 dimensions). The launch yields
one program index per axis (``blockIdx`` on CUDA, the outer loop on
CPU, the core index on an NPU).
Examples
--------
.. code-block:: python
with T.Kernel(T.ceildiv(N, 128)) as bx:
# bx is the program index along x; also iterable as (bx,)
...
"""
return launch_kernel(blocks)
@kernel_launch_factory
def ClusterKernel(
*blocks: int | tirx.PrimExpr,
cluster_dims: int | tuple[int, int, int] | list[int],
@@ -472,22 +490,7 @@ def ClusterKernel(
with T.ClusterKernel(grid_x, grid_y, cluster_dims=2, threads=128) as (bx, by):
...
"""
from tilelang.language.eager.builder import Builder
if Builder.current() is None:
raise JITNoBuilderError("T.ClusterKernel() can only be used inside @tilelang.jit or @T.prim_func context. No Builder is available.")
attrs: dict = {}
threads = _normalize_threads(threads)
if prelude is not None:
attrs["pragma_import_c"] = prelude
cluster_dims = _normalize_cluster_dims(cluster_dims)
if cluster_dims is not None:
attrs["cluster_dims"] = cluster_dims
return _ffi_api.KernelLaunch(blocks, threads, attrs)
return launch_kernel(blocks, threads=threads, prelude=prelude, cluster_dims=cluster_dims)
# For CUDA source kernels, we need to load the source code from a file or string.
@@ -583,18 +586,12 @@ def CUDASourceCodeKernel(
def get_thread_binding(dim: int = 0) -> Var:
"""Returns the thread binding for the given dimension."""
simtvf = _get_current_simtvf()
if simtvf is not None:
return simtvf.thread_vars[dim]
assert KernelLaunchFrame.Current() is not None, "KernelLaunchFrame is not initialized"
return KernelLaunchFrame.Current().get_thread_binding(dim)
def get_thread_bindings() -> list[Var]:
"""Returns all three thread bindings."""
simtvf = _get_current_simtvf()
if simtvf is not None:
return list(simtvf.thread_vars)
assert KernelLaunchFrame.Current() is not None, "KernelLaunchFrame is not initialized"
return KernelLaunchFrame.Current().get_thread_bindings()
@@ -613,18 +610,12 @@ def get_block_bindings() -> list[Var]:
def get_thread_extent(dim: int = 0) -> int:
"""Returns the thread extent for the given dimension."""
simtvf = _get_current_simtvf()
if simtvf is not None:
return simtvf.thread_extents[dim]
assert KernelLaunchFrame.Current() is not None, "KernelLaunchFrame is not initialized"
return KernelLaunchFrame.Current().get_thread_extent(dim)
def get_thread_extents() -> list[int]:
"""Returns all three thread extents."""
simtvf = _get_current_simtvf()
if simtvf is not None:
return list(simtvf.thread_extents)
assert KernelLaunchFrame.Current() is not None, "KernelLaunchFrame is not initialized"
return KernelLaunchFrame.Current().get_thread_extents()
@@ -639,3 +630,40 @@ def get_block_extents() -> list[int]:
"""Returns all three block extents."""
assert KernelLaunchFrame.Current() is not None, "KernelLaunchFrame is not initialized"
return KernelLaunchFrame.Current().get_block_extents()
def get_cluster_dims() -> list[int]:
"""Returns the cluster dimensions ``[x, y, z]`` of the current launch (``[1, 1, 1]`` without clusters)."""
assert KernelLaunchFrame.Current() is not None, "KernelLaunchFrame is not initialized"
return KernelLaunchFrame.Current().get_cluster_dims()
def get_cluster_size() -> int:
"""Returns the number of programs per cluster of the current launch."""
assert KernelLaunchFrame.Current() is not None, "KernelLaunchFrame is not initialized"
return KernelLaunchFrame.Current().get_cluster_size()
def get_cluster_id(dim: int = 0) -> Var | tirx.PrimExpr:
"""Returns the cluster index of the current program along ``dim``
(``block_id // cluster_dims[dim]``). See :meth:`KernelLaunchFrame.get_cluster_id`."""
assert KernelLaunchFrame.Current() is not None, "KernelLaunchFrame is not initialized"
return KernelLaunchFrame.Current().get_cluster_id(dim)
def get_cluster_ids() -> list[Var | tirx.PrimExpr]:
"""Returns the cluster index along every launched grid axis."""
assert KernelLaunchFrame.Current() is not None, "KernelLaunchFrame is not initialized"
return KernelLaunchFrame.Current().get_cluster_ids()
def get_cluster_extent(dim: int = 0) -> int:
"""Returns the number of clusters along ``dim``."""
assert KernelLaunchFrame.Current() is not None, "KernelLaunchFrame is not initialized"
return KernelLaunchFrame.Current().get_cluster_extent(dim)
def get_cluster_extents() -> list[int]:
"""Returns the number of clusters along all three dimensions."""
assert KernelLaunchFrame.Current() is not None, "KernelLaunchFrame is not initialized"
return KernelLaunchFrame.Current().get_cluster_extents()
+4 -1
View File
@@ -11,6 +11,8 @@ from tilelang.language.builtin import ( # noqa: F401
cooperative_tensor_store,
)
from .kernel import * # noqa: F401,F403
from .kernel import __all__ as _KERNEL_ALL
from .tir import * # noqa: F401,F403
from .tir import __all__ as _TIR_ALL
@@ -19,6 +21,7 @@ __all__ = tuple(
dict.fromkeys(
(
*_COMMON_ALL,
*_KERNEL_ALL,
*_TIR_ALL,
"cooperative_tensor_fill",
"cooperative_tensor_load",
@@ -28,4 +31,4 @@ __all__ = tuple(
)
)
del _COMMON_ALL, _TIR_ALL
del _COMMON_ALL, _KERNEL_ALL, _TIR_ALL
+45
View File
@@ -0,0 +1,45 @@
"""Metal dialect of ``T.Kernel``: the common launch plus Metal launch annotations."""
from __future__ import annotations
from tvm import tirx
from tilelang.language.kernel import KernelLaunchFrame, kernel_launch_factory, launch_kernel
__all__ = ["Kernel"]
@kernel_launch_factory
def Kernel(
*blocks: int | tirx.PrimExpr,
threads: int | list[int] | tuple[int, ...] | None = None,
prelude: str | None = None,
) -> KernelLaunchFrame:
"""Construct a kernel launch frame for Metal: a grid of threadgroups.
Code inside the launch operates at the threadgroup level: ``T.Parallel``,
``T.copy`` and friends are mapped onto threads by the compiler.
``T.get_thread_binding()`` exposes the thread index for thread-level code.
The keyword arguments are recorded at trace time and materialized by the
Metal pipeline once the target is known.
Parameters
----------
*blocks : int | PrimExpr
Grid extent along each axis (1-3 dimensions). The launch yields one
threadgroup index per axis.
threads : int | list[int] | tuple[int, ...], optional
Threads per threadgroup: a count or up to three per-dimension extents.
Defaults to 128 when omitted.
prelude : str, optional
Source injected before the generated kernel, e.g. ``#include`` lines or
helper functions.
Examples
--------
.. code-block:: python
with T.Kernel(T.ceildiv(N, 128), threads=128) as bx:
...
"""
return launch_kernel(blocks, threads=threads, prelude=prelude)
+3 -1
View File
@@ -17,7 +17,9 @@ from tilelang.metal.transform import MetalFragmentToSimdgroup
def MetalPassPipelineBody(mod: IRModule, target: Target) -> IRModule:
mod = tirx.transform.BindTarget(target)(mod)
mod = tilelang.transform.MaterializeKernelLaunch()(mod)
mod = tilelang.transform.MaterializeKernelLaunch(
lower_grid_binding=True, lower_thread_binding=True, default_threads=128, unsupported_annotations=["cluster_dims"]
)(mod)
pass_ctx = tilelang.transform.get_pass_context()
if should_force_let_inline():
+4 -2
View File
@@ -7,8 +7,10 @@ from tilelang.language.common import __all__ as _COMMON_ALL
from .intrinsics import * # noqa: F401,F403
from .intrinsics import __all__ as _ROCM_ALL
from .kernel import * # noqa: F401,F403
from .kernel import __all__ as _KERNEL_ALL
__tilelang_dialect__ = "rocm"
__all__ = tuple(dict.fromkeys((*_COMMON_ALL, *_ROCM_ALL)))
__all__ = tuple(dict.fromkeys((*_COMMON_ALL, *_ROCM_ALL, *_KERNEL_ALL)))
del _COMMON_ALL, _ROCM_ALL
del _COMMON_ALL, _ROCM_ALL, _KERNEL_ALL
+45
View File
@@ -0,0 +1,45 @@
"""ROCm dialect of ``T.Kernel``: the common launch plus ROCm launch annotations."""
from __future__ import annotations
from tvm import tirx
from tilelang.language.kernel import KernelLaunchFrame, kernel_launch_factory, launch_kernel
__all__ = ["Kernel"]
@kernel_launch_factory
def Kernel(
*blocks: int | tirx.PrimExpr,
threads: int | list[int] | tuple[int, ...] | None = None,
prelude: str | None = None,
) -> KernelLaunchFrame:
"""Construct a kernel launch frame for ROCm: a grid of workgroups.
Code inside the launch operates at the workgroup level: ``T.Parallel``,
``T.copy`` and friends are mapped onto threads by the compiler.
``T.get_thread_binding()`` exposes the thread index for thread-level code.
The keyword arguments are recorded at trace time and materialized by the
ROCm pipeline once the target is known.
Parameters
----------
*blocks : int | PrimExpr
Grid extent along each axis (1-3 dimensions). The launch yields one
workgroup index per axis.
threads : int | list[int] | tuple[int, ...], optional
Threads per workgroup: a count or up to three per-dimension extents.
Defaults to 128 when omitted.
prelude : str, optional
Source injected before the generated kernel, e.g. ``#include`` lines or
helper functions.
Examples
--------
.. code-block:: python
with T.Kernel(T.ceildiv(N, 128), threads=128) as bx:
...
"""
return launch_kernel(blocks, threads=threads, prelude=prelude)
+3 -1
View File
@@ -16,7 +16,9 @@ from tilelang.backend.pass_pipeline.pipeline_utils import (
def ROCMPassPipelineBody(mod: IRModule, target: Target) -> IRModule:
mod = tirx.transform.BindTarget(target)(mod)
mod = tilelang.transform.MaterializeKernelLaunch()(mod)
mod = tilelang.transform.MaterializeKernelLaunch(
lower_grid_binding=True, lower_thread_binding=True, default_threads=128, unsupported_annotations=["cluster_dims"]
)(mod)
pass_ctx = tilelang.transform.get_pass_context()
if should_force_let_inline():
+4 -3
View File
@@ -73,9 +73,10 @@ This writes `gemm_relu_passes.html` (the interactive browser) and a sibling
from tilelang.tools.pass_visualizer.viewer import build_pass_data, emit_html
name, stages = build_pass_data(
"path/to/kernel.py", factory=None, target="auto",
kwargs={"M": 1024, "N": 1024, "K": 1024,
"block_M": 128, "block_N": 128, "block_K": 32},
"path/to/kernel.py",
factory=None,
target="auto",
kwargs={"M": 1024, "N": 1024, "K": 1024, "block_M": 128, "block_N": 128, "block_K": 32},
source=open("path/to/kernel.py").read(),
)
html = emit_html(name, stages)
+60 -10
View File
@@ -198,26 +198,76 @@ def MakePackedAPI():
return _ffi_api.MakePackedAPI() # type: ignore
def MaterializeKernelLaunch(lower_thread_binding: bool = True):
"""Materialize the target-neutral kernel launch nest (thread_binding
For loops emitted by T.Kernel) into a backend-specific form. Each
backend pipeline decides the mode for itself:
DEFAULT_SIMT_THREADS = 128
def MaterializeKernelLaunch(
lower_grid_binding: bool | None = None,
lower_thread_binding: bool = True,
default_threads: int | list[int] | tuple | None = DEFAULT_SIMT_THREADS,
unsupported_annotations: list[str] | tuple[str, ...] | None = None,
launch_dim_tags: list[str] | tuple[str, ...] | None = None,
):
"""Materialize the target-neutral kernel launch nest emitted by T.Kernel
into a backend-specific form. Each backend pipeline decides the mode for
itself; this is where the target-dependent parts of a launch (whether a
program-index space exists, whether threads exist and how many run by
default) are decided.
Parameters
----------
lower_grid_binding : bool | None
If True (targets with a real block/core-level launch, e.g. CUDA,
Ascend), lower the launch loops (blockIdx.* grid axes and any tag in
``launch_dim_tags``) into thread_extent AttrStmts carrying each loop's
own thread tag.
If False (targets with no program-index space, e.g. CPU), lower those
loops into plain serial For loops.
If None (the default), follow ``lower_thread_binding``: a backend with
SIMT threads has a program-index space too, and a backend without them
has none. Pass this explicitly to decouple the two, as Ascend does.
lower_thread_binding : bool
If True (SIMT backends, e.g. CUDA/ROCm/Metal), lower the
blockIdx.*/threadIdx.* loops into thread_extent AttrStmts.
If False (backends without SIMT, e.g. CPU), lower blockIdx.*
loops into plain serial For loops and ignore threadIdx.* loops
(their extents are dropped; the loop vars are pinned to 0).
If True (SIMT backends, e.g. CUDA/ROCm/Metal), bind the thread
placeholders as threadIdx.* thread_extent scopes.
If False (backends without SIMT, e.g. CPU and Ascend), drop the thread
placeholders. A body that references a thread index is rejected on such
targets. Ascend pairs this with ``lower_grid_binding=True``: its NPU
launch is a real 1-D core grid, while thread domains only exist inside
``T.SimtVF``, which emits its own thread scopes below the launch nest.
default_threads : int | list[int] | tuple | None
Thread-block extents used by SIMT backends when T.Kernel was called
without ``threads=``. Ignored when ``lower_thread_binding`` is False.
None means the backend has no default and ``threads=`` is required.
unsupported_annotations : list[str] | None
Launch annotations (keys on the ``tilelang_root`` block, e.g.
``cluster_dims``) that have no meaning on this backend. A launch
carrying one is rejected here instead of being silently ignored by
later passes.
launch_dim_tags : list[str] | None
Extra thread_binding tags that belong to the launch nest rather than to
the thread domain, so a backend can extend the launch vocabulary
without this pass knowing about it. Ascend passes ``["cthread"]`` for
``T.MixedKernel``'s sub-block-id dimension.
Returns
-------
fpass : tvm.transform.Pass
The result pass
"""
return _ffi_api.MaterializeKernelLaunch(lower_thread_binding) # type: ignore
if lower_grid_binding is None:
lower_grid_binding = lower_thread_binding
if default_threads is not None:
if isinstance(default_threads, int):
default_threads = [default_threads, 1, 1]
else:
default_threads = list(default_threads) + [1] * (3 - len(default_threads))
if unsupported_annotations is not None:
unsupported_annotations = list(unsupported_annotations)
if launch_dim_tags is not None:
launch_dim_tags = list(launch_dim_tags)
return _ffi_api.MaterializeKernelLaunch( # type: ignore
lower_grid_binding, lower_thread_binding, default_threads, unsupported_annotations, launch_dim_tags
)
def AnnotateDeviceRegions():
+3 -5
View File
@@ -76,11 +76,9 @@ class PassConfigKey(str, Enum):
TL_ENABLE_AUTO_SCHEDULE = "tl.enable_auto_schedule"
"""Enable Ascend auto scheduling. Default: True."""
TL_CUDA_AUTO_SCHEDULE = "tl.cuda_auto_schedule"
"""Name of the CUDA automatic warp-specialization scheduler (e.g.
"role_based"). Default: unset (disabled). Separate from Ascend's boolean
TL_ENABLE_AUTO_SCHEDULE because pass-config types are process-global.
"""
TL_ENABLE_AUTO_WARP_SPECIALIZATION = "tl.enable_auto_warp_specialization"
"""Name of the automatic warp-specialization scheduler to run (e.g.
"role_based"). Default: unset (disabled)."""
TL_ENABLE_FAST_MATH = "tl.enable_fast_math"
"""
+6 -3
View File
@@ -5,7 +5,10 @@ from __future__ import annotations
from tilelang.language.common import * # noqa: F401,F403
from tilelang.language.common import __all__ as _COMMON_ALL
__tilelang_dialect__ = "webgpu"
__all__ = tuple(_COMMON_ALL)
from .kernel import * # noqa: F401,F403
from .kernel import __all__ as _KERNEL_ALL
del _COMMON_ALL
__tilelang_dialect__ = "webgpu"
__all__ = tuple(dict.fromkeys((*_COMMON_ALL, *_KERNEL_ALL)))
del _COMMON_ALL, _KERNEL_ALL
+41
View File
@@ -0,0 +1,41 @@
"""WebGPU dialect of ``T.Kernel``: the common launch plus WebGPU launch annotations."""
from __future__ import annotations
from tvm import tirx
from tilelang.language.kernel import KernelLaunchFrame, kernel_launch_factory, launch_kernel
__all__ = ["Kernel"]
@kernel_launch_factory
def Kernel(
*blocks: int | tirx.PrimExpr,
threads: int | list[int] | tuple[int, ...] | None = None,
) -> KernelLaunchFrame:
"""Construct a kernel launch frame for WebGPU: a grid of workgroups.
Code inside the launch operates at the workgroup level: ``T.Parallel``,
``T.copy`` and friends are mapped onto invocations by the compiler.
``T.get_thread_binding()`` exposes the invocation index for thread-level
code. ``threads`` is recorded at trace time and materialized by the WebGPU
pipeline once the target is known.
Parameters
----------
*blocks : int | PrimExpr
Grid extent along each axis (1-3 dimensions). The launch yields one
workgroup index per axis.
threads : int | list[int] | tuple[int, ...], optional
Invocations per workgroup: a count or up to three per-dimension
extents. Defaults to 128 when omitted.
Examples
--------
.. code-block:: python
with T.Kernel(T.ceildiv(N, 128), threads=128) as bx:
...
"""
return launch_kernel(blocks, threads=threads)
+3 -1
View File
@@ -18,7 +18,9 @@ def WebGPUPassPipelineBody(mod: IRModule, target: Target) -> IRModule:
mod = tirx.transform.BindTarget(target)(mod)
# WebGPU is a SIMT backend: lower the launch nest to thread_extent
# bindings for codegen.
mod = tilelang.transform.MaterializeKernelLaunch()(mod)
mod = tilelang.transform.MaterializeKernelLaunch(
lower_grid_binding=True, lower_thread_binding=True, default_threads=128, unsupported_annotations=["cluster_dims"]
)(mod)
pass_ctx = tilelang.transform.get_pass_context()
if should_force_let_inline():