Skip to content

Commit 612c428

Browse files
committed
feat: add LLVM backend support
Enable LLVM as a TileLang backend in CMake and route llvm targets through TVM LLVM codegen during lowering. Fix tvm_ffi output tensor device inference by deriving the output device from tensor inputs, and add LLVM codegen/compile regression tests for matmul, T.gemm, T.copy, and T.While.
1 parent 446e8da commit 612c428

6 files changed

Lines changed: 329 additions & 2 deletions

File tree

‎CMakeLists.txt‎

Lines changed: 12 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -250,11 +250,12 @@ else()
250250
endif()
251251

252252
# Configs
253-
set(TILELANG_BACKENDS CUDA ROCM METAL)
253+
set(TILELANG_BACKENDS CUDA ROCM METAL LLVM)
254254

255255
set(TILELANG_BACKEND_DOC_CUDA "Enable CUDA backend (ON/OFF/or CUDA SDK path)")
256256
set(TILELANG_BACKEND_DOC_ROCM "Enable ROCm backend (ON/OFF/or ROCm SDK path)")
257257
set(TILELANG_BACKEND_DOC_METAL "Enable Metal backend")
258+
set(TILELANG_BACKEND_DOC_LLVM "Enable LLVM backend")
258259

259260
# TVM's config.cmake redefines USE_* options later, so we cache the user's choice
260261
# (including explicit -DUSE_XXX arguments) before we include TVM and restore it
@@ -418,6 +419,16 @@ if(NOT TILELANG_BACKEND_USER_SELECTED)
418419
endif()
419420
endif()
420421

422+
423+
if(DEFINED ENV{USE_LLVM})
424+
set(_tilelang_backend_env_selected ON)
425+
if($ENV{USE_LLVM})
426+
set(USE_LLVM ON)
427+
else()
428+
set(USE_LLVM OFF)
429+
endif()
430+
endif()
431+
421432
if(NOT _tilelang_backend_env_selected)
422433
if(APPLE)
423434
message(STATUS "Enable Metal support by default.")
Lines changed: 155 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,155 @@
1+
import tilelang
2+
import tilelang.testing
3+
from tilelang import tvm as tvm
4+
import tilelang.language as T
5+
import torch
6+
7+
8+
def matmul(M, N, K, block_M, block_N, block_K, dtype=T.float16, accum_dtype=T.float32):
9+
num_stages = 0
10+
11+
@T.prim_func
12+
def matmul(
13+
A: T.Tensor((M, K), dtype),
14+
B: T.Tensor((K, N), dtype),
15+
C: T.Tensor((M, N), dtype),
16+
):
17+
with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M)) as (bx, by):
18+
A_local = T.alloc_local((block_M, block_K), dtype)
19+
B_local = T.alloc_local((block_K, block_N), dtype)
20+
C_local = T.alloc_local((block_M, block_N), accum_dtype)
21+
22+
T.clear(C_local)
23+
24+
# Apply layout optimizations or define your own layout
25+
# (Optional).
26+
# T.annotate_layout(
27+
# {
28+
# A_local: make_swizzle_layout(A_local),
29+
# B_local: make_swizzle_layout(B_local),
30+
# }
31+
# )
32+
33+
for ko in T.Pipelined(K // block_K, num_stages=num_stages):
34+
T.copy(A[by * block_M, ko * block_K], A_local)
35+
36+
# Or Copy with Parallel
37+
for k, j in T.Parallel(block_K, block_N):
38+
B_local[k, j] = B[ko * block_K + k, by * block_N + j]
39+
40+
for i, j, k in T.grid(block_M, block_N, block_K):
41+
C_local[i, j] += A_local[i, k] * B_local[k, j]
42+
43+
T.copy(C_local, C[by * block_M, bx * block_N])
44+
45+
return matmul
46+
47+
48+
def assert_matmul_codegen(M=1024, N=1024, K=1024, block_M=128, block_N=128, block_K=32):
49+
func = matmul(M, N, K, block_M, block_N, block_K)
50+
51+
with tvm.target.Target("llvm"):
52+
artifact = tilelang.lower(func)
53+
54+
code = artifact.kernel_source
55+
56+
assert code is not None, "Code generation failed"
57+
58+
59+
def test_matmul_codegen():
60+
assert_matmul_codegen(M=1024, N=1024, K=1024, block_M=128, block_N=128, block_K=32)
61+
62+
63+
def test_matmul_compile():
64+
def matmul_jit_test(M, N, K, block_M, block_N, block_K, dtype=T.float16, accum_dtype=T.float32):
65+
# a simple kernel just for jit test
66+
@T.prim_func
67+
def matmul(
68+
A: T.Tensor((M, K), dtype),
69+
B: T.Tensor((K, N), dtype),
70+
C: T.Tensor((M, N), dtype),
71+
):
72+
with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M)) as (bx, by):
73+
A_local = T.alloc_local((block_M, block_K), dtype)
74+
B_local = T.alloc_local((block_K, block_N), dtype)
75+
C_local = T.alloc_local((block_M, block_N), accum_dtype)
76+
77+
for p in T.serial(block_M):
78+
for w in T.serial(block_N):
79+
C_local[p, w] = 0
80+
for ko in T.serial(K // block_K):
81+
for i in T.serial(block_M):
82+
for k in T.serial(block_K):
83+
A_local[i, k] = A[by * block_M + i, ko * block_K + k]
84+
85+
for k in T.serial(block_K):
86+
for j in T.serial(block_N):
87+
B_local[k, j] = B[ko * block_K + k, bx * block_N + j]
88+
89+
for i in T.serial(block_M):
90+
for j in T.serial(block_N):
91+
for k in T.serial(block_K):
92+
C_local[i, j] += A_local[i, k] * B_local[k, j]
93+
94+
for i in T.serial(block_M):
95+
for j in T.serial(block_N):
96+
C[by * block_M + i, bx * block_N + j] = C_local[i, j]
97+
98+
return matmul
99+
100+
M, N, K = 1024, 512, 512
101+
block_M, block_N, block_K = M // 4, N // 4, K // 4
102+
llvm_func = matmul_jit_test(M, N, K, block_M, block_N, block_K)
103+
with tvm.target.Target("llvm"):
104+
complied_fun = tilelang.compile(llvm_func, -1, execution_backend="tvm_ffi")
105+
106+
in_dtype = T.float16
107+
A = torch.randn(M, K, dtype=torch.__getattribute__(in_dtype))
108+
B = torch.randn(K, N, dtype=torch.__getattribute__(in_dtype))
109+
110+
C = complied_fun(A, B)
111+
C_torch = torch.matmul(A, B)
112+
113+
tilelang.testing.torch_assert_close(C, C_torch, atol=1e-2, rtol=1e-2, max_mismatched_ratio=0.05)
114+
115+
116+
def test_matmul_with_copy_tvm_ffi():
117+
"""LLVM kernel using T.copy with tvm_ffi backend.
118+
119+
Verifies that T.copy works end-to-end on LLVM backend: the vectorized copy
120+
uses vector types (e.g. float4) defined in common.h, and the
121+
wrapper correctly skips redundant re-lowering.
122+
"""
123+
M, N, K = 128, 128, 128
124+
block_M, block_N, block_K = 32, 32, 32
125+
126+
@T.prim_func
127+
def matmul(
128+
A: T.Tensor((M, K), "float32"),
129+
B: T.Tensor((K, N), "float32"),
130+
C: T.Tensor((M, N), "float32"),
131+
):
132+
with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M)) as (bx, by):
133+
A_local = T.alloc_local((block_M, block_K), "float32")
134+
B_local = T.alloc_local((block_K, block_N), "float32")
135+
C_local = T.alloc_local((block_M, block_N), "float32")
136+
137+
T.clear(C_local)
138+
for ko in T.serial(K // block_K):
139+
T.copy(A[by * block_M, ko * block_K], A_local)
140+
T.copy(B[ko * block_K, bx * block_N], B_local)
141+
for i, j, k in T.grid(block_M, block_N, block_K):
142+
C_local[i, j] += A_local[i, k] * B_local[k, j]
143+
T.copy(C_local, C[by * block_M, bx * block_N])
144+
145+
compiled = tilelang.compile(matmul, target="llvm", out_idx=-1, execution_backend="tvm_ffi")
146+
147+
a = torch.randn(M, K, dtype=torch.float32)
148+
b = torch.randn(K, N, dtype=torch.float32)
149+
c = compiled(a, b)
150+
ref = a @ b
151+
torch.testing.assert_close(c, ref, rtol=1e-5, atol=1e-5)
152+
153+
154+
if __name__ == "__main__":
155+
tilelang.testing.main()
Lines changed: 128 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,128 @@
1+
"""Tests for T.gemm on LLVM target (GemmScalar path).
2+
3+
Verifies that T.gemm works correctly with for various
4+
matrix sizes, block sizes, and transpose combinations.
5+
"""
6+
7+
import pytest
8+
import torch
9+
import tilelang
10+
import tilelang.testing
11+
from tilelang import tvm as tvm
12+
import tilelang.language as T
13+
14+
15+
def matmul(M, N, K, block_M, block_N, block_K, trans_A=False, trans_B=False, dtype=T.float32, accum_dtype=T.float32):
16+
A_shape = (K, M) if trans_A else (M, K)
17+
B_shape = (N, K) if trans_B else (K, N)
18+
A_local_shape = (block_K, block_M) if trans_A else (block_M, block_K)
19+
B_local_shape = (block_N, block_K) if trans_B else (block_K, block_N)
20+
21+
@T.prim_func
22+
def main(
23+
A: T.Tensor(A_shape, dtype),
24+
B: T.Tensor(B_shape, dtype),
25+
C: T.Tensor((M, N), dtype),
26+
):
27+
with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M)) as (bx, by):
28+
A_local = T.alloc_local(A_local_shape, dtype)
29+
B_local = T.alloc_local(B_local_shape, dtype)
30+
C_local = T.alloc_local((block_M, block_N), accum_dtype)
31+
T.clear(C_local)
32+
for ko in T.serial(T.ceildiv(K, block_K)):
33+
if trans_A:
34+
T.copy(A[ko * block_K, by * block_M], A_local)
35+
else:
36+
T.copy(A[by * block_M, ko * block_K], A_local)
37+
if trans_B:
38+
T.copy(B[bx * block_N, ko * block_K], B_local)
39+
else:
40+
T.copy(B[ko * block_K, bx * block_N], B_local)
41+
T.gemm(A_local, B_local, C_local, trans_A, trans_B)
42+
T.copy(C_local, C[by * block_M, bx * block_N])
43+
44+
return main
45+
46+
47+
def ref_matmul(A, B, trans_A, trans_B):
48+
if trans_A:
49+
A = A.T
50+
if trans_B:
51+
B = B.T
52+
return torch.matmul(A.float(), B.float()).to(A.dtype)
53+
54+
55+
def run_gemm_codegen(M, N, K, block_M, block_N, block_K, trans_A=False, trans_B=False):
56+
func = matmul(M, N, K, block_M, block_N, block_K, trans_A, trans_B)
57+
with tvm.target.Target("llvm"):
58+
artifact = tilelang.lower(func)
59+
code = artifact.kernel_source
60+
assert code is not None, "Code generation failed"
61+
return code
62+
63+
64+
def run_gemm_compile(M, N, K, block_M, block_N, block_K, trans_A=False, trans_B=False, dtype=T.float32):
65+
func = matmul(M, N, K, block_M, block_N, block_K, trans_A, trans_B, dtype=dtype)
66+
kernel = tilelang.compile(func, target="llvm", out_idx=[2], execution_backend="tvm_ffi")
67+
68+
torch_dtype = torch.__getattribute__(dtype)
69+
A_shape = (K, M) if trans_A else (M, K)
70+
B_shape = (N, K) if trans_B else (K, N)
71+
A = torch.randn(A_shape, dtype=torch_dtype)
72+
B = torch.randn(B_shape, dtype=torch_dtype)
73+
74+
C = kernel(A, B)
75+
C_ref = ref_matmul(A, B, trans_A, trans_B)
76+
77+
tilelang.testing.torch_assert_close(C, C_ref, atol=1e-2, rtol=1e-2)
78+
79+
80+
# --- Codegen tests ---
81+
82+
83+
def test_codegen_basic():
84+
run_gemm_codegen(128, 128, 128, 64, 64, 64)
85+
86+
87+
def test_codegen_rectangular():
88+
run_gemm_codegen(256, 512, 128, 64, 64, 64)
89+
90+
91+
def test_codegen_trans_A():
92+
run_gemm_codegen(128, 128, 128, 64, 64, 64, trans_A=True)
93+
94+
95+
def test_codegen_trans_B():
96+
run_gemm_codegen(128, 128, 128, 64, 64, 64, trans_B=True)
97+
98+
99+
# --- Compile + correctness tests ---
100+
101+
102+
@pytest.mark.parametrize(
103+
"M,N,K,block_M,block_N,block_K",
104+
[
105+
(128, 128, 128, 64, 64, 64),
106+
(256, 256, 256, 64, 64, 64),
107+
(256, 512, 128, 64, 64, 64),
108+
(512, 512, 512, 128, 128, 128),
109+
],
110+
)
111+
def test_gemm_f32_nn(M, N, K, block_M, block_N, block_K):
112+
run_gemm_compile(M, N, K, block_M, block_N, block_K)
113+
114+
115+
def test_gemm_f32_tn():
116+
run_gemm_compile(128, 128, 128, 64, 64, 64, trans_A=True)
117+
118+
119+
def test_gemm_f32_nt():
120+
run_gemm_compile(128, 128, 128, 64, 64, 64, trans_B=True)
121+
122+
123+
def test_gemm_f32_tt():
124+
run_gemm_compile(128, 128, 128, 64, 64, 64, trans_A=True, trans_B=True)
125+
126+
127+
if __name__ == "__main__":
128+
tilelang.testing.main()
Lines changed: 28 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,28 @@
1+
import tilelang
2+
import tilelang.language as T
3+
4+
5+
def test_llvm_kernel_source_generation_with_while() -> None:
6+
"""Regression test: LLVM `T.While` kernels should compile without errors.
7+
8+
See: https://github.com/tile-ai/tilelang/issues/2202
9+
10+
Historically, a CPU kernel containing `T.While(...)` inside
11+
`T.Kernel(...)` could leak the synthetic fallback thread variable
12+
`v_thread` into host/device splitting. LLVM uses the same scalar lowering
13+
path for this case, so this keeps equivalent coverage for>
14+
"""
15+
16+
@T.prim_func
17+
def main(flag: T.Tensor((1,), "int32"), out: T.Tensor((1,), "int32")):
18+
with T.Kernel(1):
19+
state = T.alloc_fragment((1,), "int32")
20+
state[0] = 0
21+
22+
with T.While(state[0] == 0):
23+
state[0] = flag[0]
24+
25+
out[0] = state[0]
26+
27+
compiled = tilelang.compile(main, target="llvm", execution_backend="tvm_ffi")
28+
assert compiled is not None

‎tilelang/engine/lower.py‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -241,6 +241,8 @@ def device_codegen(device_mod: tvm.IRModule, target: Target) -> tvm.IRModule:
241241
device_mod = tvm.ffi.get_global_func("target.build.tilelang_hip")(device_mod, target)
242242
elif target.kind.name == "metal":
243243
device_mod = tvm.ffi.get_global_func("target.build.tilelang_metal")(device_mod, target)
244+
elif target.kind.name == "llvm":
245+
device_mod = tvm.ffi.get_global_func("target.build.llvm")(device_mod, target)
244246
else:
245247
raise ValueError(f"Target {target.kind.name} is not supported")
246248

‎tilelang/jit/adapter/tvm_ffi.py‎

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -209,7 +209,10 @@ def func(*inputs: torch.Tensor | Any):
209209

210210
# Resolve the device used for outputs. Prefer the first tensor input's device
211211
# if available, otherwise use PyTorch's current device.
212-
out_device: torch.device | None = None
212+
out_device: torch.device | None = next(
213+
(input.device for input in inputs if isinstance(input, torch.Tensor)),
214+
None,
215+
)
213216

214217
# Stitch the full positional argument list expected by the TVM executable
215218
ins_idx: int = 0

0 commit comments

Comments
 (0)