mirror of
https://github.com/tile-ai/tilelang.git
synced 2026-10-03 14:59:54 +08:00
[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:
@@ -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
|
||||
|
||||
@@ -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 }}
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
```
|
||||
|
||||
@@ -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)`.
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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]
|
||||
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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])
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:")
|
||||
```
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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])
|
||||
```
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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
|
||||
|
||||
```
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
@@ -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
|
||||
```
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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])
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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";
|
||||
|
||||
Vendored
+59
-59
@@ -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;
|
||||
|
||||
/**
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 {
|
||||
|
||||
|
||||
@@ -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>();
|
||||
|
||||
@@ -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));
|
||||
|
||||
|
||||
@@ -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 ®ion();
|
||||
|
||||
/*!
|
||||
* \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();
|
||||
|
||||
@@ -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";
|
||||
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
```
|
||||
|
||||
Vendored
+74
-41
@@ -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)
|
||||
|
||||
Vendored
+1
-2
@@ -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
|
||||
|
||||
@@ -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`
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
@@ -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():
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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")
|
||||
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
@@ -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():
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
@@ -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():
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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"
|
||||
"""
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
@@ -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():
|
||||
|
||||
Reference in New Issue
Block a user