Skip to content
Prev Previous commit
Next Next commit
add test
  • Loading branch information
Rachmanino committed Feb 26, 2026
commit e5c1ec70d99cbf2a29b7e232da3d486c4f14e47e
61 changes: 61 additions & 0 deletions testing/python/language/test_tilelang_language_cluster_launch.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,61 @@
import tilelang
import tilelang.language as T
import torch
import tilelang.testing


def matmul(M, N, K, block_M, block_N, block_K, dtype=T.float16, accum_dtype=T.float32):
@T.prim_func
def gemm(
A: T.Tensor((M, K), dtype),
B: T.Tensor((K, N), dtype),
C: T.Tensor((M, N), dtype),
):
with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M), threads=128, cluster_dims=(2, 1, 1)) 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)

T.clear(C_local)
for k in T.Pipelined(T.ceildiv(K, block_K), num_stages=3):
T.copy(A[by * block_M, k * block_K], A_shared)
T.copy(B[k * block_K, bx * block_N], B_shared)
T.gemm(A_shared, B_shared, C_local)

T.copy(C_local, C[by * block_M, bx * block_N])

return gemm


def run_cython_cluster_launch():
kernel = matmul(1024, 1024, 1024, 128, 128, 32)
mod = tilelang.compile(kernel, execution_backend="cython")
assert 'clusterDim = {2, 1, 1}' in mod.get_host_source()


def run_tvm_ffi_cluster_launch():
kernel = matmul(1024, 1024, 1024, 128, 128, 32)
mod = tilelang.compile(kernel, execution_backend="tvm_ffi")
check_str = r"""
(((TVMFFIAny*)stack_ffi_any)[3].type_index) = 1;
(((TVMFFIAny*)stack_ffi_any)[3].zero_padding) = 0;
(((TVMFFIAny*)stack_ffi_any)[3].v_int64) = ((int64_t)2);
(((TVMFFIAny*)stack_ffi_any)[4].type_index) = 1;
(((TVMFFIAny*)stack_ffi_any)[4].zero_padding) = 0;
(((TVMFFIAny*)stack_ffi_any)[4].v_int64) = ((int64_t)1);
(((TVMFFIAny*)stack_ffi_any)[5].type_index) = 1;
(((TVMFFIAny*)stack_ffi_any)[5].zero_padding) = 0;
(((TVMFFIAny*)stack_ffi_any)[5].v_int64) = ((int64_t)1);
"""
assert check_str in mod.get_host_source()

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

⚠️ Potential issue | 🟡 Minor

Make host-source assertions less formatting-fragile.

The current checks are tightly coupled to exact whitespace/layout, especially the multiline TVM FFI snippet, so harmless codegen formatting changes can fail the test.

Proposed refactor
+def _assert_tvm_ffi_cluster_dims(host_src: str) -> None:
+    required = (
+        "(((TVMFFIAny*)stack_ffi_any)[3].v_int64) = ((int64_t)2);",
+        "(((TVMFFIAny*)stack_ffi_any)[4].v_int64) = ((int64_t)1);",
+        "(((TVMFFIAny*)stack_ffi_any)[5].v_int64) = ((int64_t)1);",
+    )
+    for snippet in required:
+        assert snippet in host_src
+
 def run_cython_cluster_launch():
     kernel = matmul(1024, 1024, 1024, 128, 128, 32)
     mod = tilelang.compile(kernel, execution_backend="cython")
-    assert 'clusterDim = {2, 1, 1}' in mod.get_host_source()
+    host_src = mod.get_host_source()
+    assert "clusterDim" in host_src and "{2, 1, 1}" in host_src
@@
 def run_tvm_ffi_cluster_launch():
     kernel = matmul(1024, 1024, 1024, 128, 128, 32)
     mod = tilelang.compile(kernel, execution_backend="tvm_ffi")
-    check_str = r"""
-  (((TVMFFIAny*)stack_ffi_any)[3].type_index) = 1;
-  (((TVMFFIAny*)stack_ffi_any)[3].zero_padding) = 0;
-  (((TVMFFIAny*)stack_ffi_any)[3].v_int64) = ((int64_t)2);
-  (((TVMFFIAny*)stack_ffi_any)[4].type_index) = 1;
-  (((TVMFFIAny*)stack_ffi_any)[4].zero_padding) = 0;
-  (((TVMFFIAny*)stack_ffi_any)[4].v_int64) = ((int64_t)1);
-  (((TVMFFIAny*)stack_ffi_any)[5].type_index) = 1;
-  (((TVMFFIAny*)stack_ffi_any)[5].zero_padding) = 0;
-  (((TVMFFIAny*)stack_ffi_any)[5].v_int64) = ((int64_t)1);
-"""
-    assert check_str in mod.get_host_source()
+    _assert_tvm_ffi_cluster_dims(mod.get_host_source())
🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed.

In `@testing/python/language/test_tilelang_language_cluster_launch.py` around
lines 33 - 50, The test run_tvm_ffi_cluster_launch is brittle because it asserts
an exact multiline snippet (check_str) against mod.get_host_source(), which
breaks on harmless formatting changes; update the test to normalize or
pattern-match the host source instead: fetch the string via
mod.get_host_source(), collapse or normalize whitespace (e.g., replace
consecutive whitespace/newlines with a single space) or use regex to assert the
presence of the essential tokens like "stack_ffi_any", "[3].type_index",
"[3].v_int64", "[4].v_int64", "[5].v_int64" and the numeric values 2,1,1 rather
than comparing the exact multiline layout; apply this change inside
run_tvm_ffi_cluster_launch replacing the check_str exact-match assertion with
the whitespace-normalized or regex-based assertions so the test passes despite
formatting changes.



@tilelang.testing.requires_cuda
@tilelang.testing.requires_cuda_compute_version_ge(9, 0)
def test_cluster_launch():
run_cython_cluster_launch()
run_tvm_ffi_cluster_launch()


if __name__ == "__main__":
test_cluster_launch()