Skip to content
Closed
Show file tree
Hide file tree
Changes from 7 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion 3rdparty/tvm
Submodule tvm updated from e47e76 to 001022
22 changes: 11 additions & 11 deletions examples/lazy_jit/lazyjit.en.ipynb
Original file line number Diff line number Diff line change
Expand Up @@ -53,7 +53,7 @@
"metadata": {},
"outputs": [],
"source": [
"@tilelang.lazy_jit\n",
"@tilelang.jit\n",
"def gemm(\n",
" A,\n",
" B,\n",
Expand Down Expand Up @@ -209,7 +209,7 @@
"metadata": {},
"outputs": [],
"source": [
"@tilelang.lazy_jit\n",
"@tilelang.jit\n",
"def gemm_dyn_K(A, B):\n",
" M, N, K = T.dynamic(\"M, N, K\")\n",
" A: T.Tensor[[M, K], T.float16]\n",
Expand Down Expand Up @@ -248,7 +248,7 @@
"metadata": {},
"outputs": [],
"source": [
"@tilelang.lazy_jit\n",
"@tilelang.jit\n",
"def as_contingious(A):\n",
" M, N, dM, dN = T.dynamic(\"M, N, dM, dN\")\n",
" A: T.StridedTensor[[M, N], [dM, dN], T.float32]\n",
Expand Down Expand Up @@ -307,7 +307,7 @@
"metadata": {},
"outputs": [],
"source": [
"@tilelang.lazy_jit\n",
"@tilelang.jit\n",
"def gemm_ptr(\n",
" A,\n",
" B,\n",
Expand Down Expand Up @@ -359,7 +359,7 @@
"metadata": {},
"outputs": [],
"source": [
"@tilelang.lazy_jit\n",
"@tilelang.jit\n",
"def gemm_ptr_dyn(A, B, M, N, K):\n",
" M: T.int32\n",
" N: T.int32\n",
Expand Down Expand Up @@ -421,7 +421,7 @@
}
],
"source": [
"@tilelang.lazy_jit\n",
"@tilelang.jit\n",
"def example_wrong_kernel(A):\n",
" M = T.const(\"M\")\n",
" A: T.Tensor[[M * 2, M * 3], T.float32]\n",
Expand Down Expand Up @@ -470,7 +470,7 @@
}
],
"source": [
"@tilelang.lazy_jit\n",
"@tilelang.jit\n",
"def dyn_annot(\n",
" A: T.ptr, # 1. T.ptr type annotation\n",
" is_2d=False,\n",
Expand Down Expand Up @@ -515,7 +515,7 @@
"metadata": {},
"outputs": [],
"source": [
"@tilelang.lazy_jit\n",
"@tilelang.jit\n",
"def add_one(X, data: T.float32 = 1):\n",
" M, N = T.const(\"M, N\")\n",
" X: T.Tensor[[M, N], T.float32]\n",
Expand Down Expand Up @@ -577,7 +577,7 @@
"B = torch.randn(128, 128, dtype=torch.float16, device=\"cuda\")\n",
"\n",
"\n",
"@tilelang.lazy_jit\n",
"@tilelang.jit\n",
"def dummy_kernel(A, B):\n",
" M, N = T.const(\"M, N\")\n",
" A: T.Tensor[[M, N], T.float16]\n",
Expand Down Expand Up @@ -797,7 +797,7 @@
"metadata": {},
"outputs": [],
"source": [
"@tilelang.lazy_jit\n",
"@tilelang.jit\n",
"def element_wise(A, fn):\n",
" N = T.dynamic(\"N\")\n",
" A: T.Tensor[[N], T.float32]\n",
Expand Down Expand Up @@ -857,7 +857,7 @@
" n31(x * 3 + 1, var)\n",
"\n",
"\n",
"@tilelang.lazy_jit\n",
"@tilelang.jit\n",
"def foo(A: T.Tensor[[1], T.int32], n: int):\n",
" with T.Kernel(1) as _:\n",
" n31(n, A[0])"
Expand Down
22 changes: 11 additions & 11 deletions examples/lazy_jit/lazyjit.zh.ipynb
Original file line number Diff line number Diff line change
Expand Up @@ -53,7 +53,7 @@
"metadata": {},
"outputs": [],
"source": [
"@tilelang.lazy_jit\n",
"@tilelang.jit\n",
"def gemm(\n",
" A,\n",
" B,\n",
Expand Down Expand Up @@ -209,7 +209,7 @@
"metadata": {},
"outputs": [],
"source": [
"@tilelang.lazy_jit\n",
"@tilelang.jit\n",
"def gemm_dyn_K(A, B):\n",
" M, N, K = T.dynamic(\"M, N, K\")\n",
" A: T.Tensor[[M, K], T.float16]\n",
Expand Down Expand Up @@ -248,7 +248,7 @@
"metadata": {},
"outputs": [],
"source": [
"@tilelang.lazy_jit\n",
"@tilelang.jit\n",
"def as_contingious(A):\n",
" M, N, dM, dN = T.dynamic(\"M, N, dM, dN\")\n",
" A: T.StridedTensor[[M, N], [dM, dN], T.float32]\n",
Expand Down Expand Up @@ -307,7 +307,7 @@
"metadata": {},
"outputs": [],
"source": [
"@tilelang.lazy_jit\n",
"@tilelang.jit\n",
"def gemm_ptr(\n",
" A,\n",
" B,\n",
Expand Down Expand Up @@ -359,7 +359,7 @@
"metadata": {},
"outputs": [],
"source": [
"@tilelang.lazy_jit\n",
"@tilelang.jit\n",
"def gemm_ptr_dyn(A, B, M, N, K):\n",
" M: T.int32\n",
" N: T.int32\n",
Expand Down Expand Up @@ -421,7 +421,7 @@
}
],
"source": [
"@tilelang.lazy_jit\n",
"@tilelang.jit\n",
"def example_wrong_kernel(A):\n",
" M = T.const(\"M\")\n",
" A: T.Tensor[[M * 2, M * 3], T.float32]\n",
Expand Down Expand Up @@ -470,7 +470,7 @@
}
],
"source": [
"@tilelang.lazy_jit\n",
"@tilelang.jit\n",
"def dyn_annot(\n",
" A: T.ptr, # 1. T.ptr type annotation\n",
" is_2d=False,\n",
Expand Down Expand Up @@ -515,7 +515,7 @@
"metadata": {},
"outputs": [],
"source": [
"@tilelang.lazy_jit\n",
"@tilelang.jit\n",
"def add_one(X, data: T.float32 = 1):\n",
" M, N = T.const(\"M, N\")\n",
" X: T.Tensor[[M, N], T.float32]\n",
Expand Down Expand Up @@ -577,7 +577,7 @@
"B = torch.randn(128, 128, dtype=torch.float16, device=\"cuda\")\n",
"\n",
"\n",
"@tilelang.lazy_jit\n",
"@tilelang.jit\n",
"def dummy_kernel(A, B):\n",
" M, N = T.const(\"M, N\")\n",
" A: T.Tensor[[M, N], T.float16]\n",
Expand Down Expand Up @@ -797,7 +797,7 @@
"metadata": {},
"outputs": [],
"source": [
"@tilelang.lazy_jit\n",
"@tilelang.jit\n",
"def element_wise(A, fn):\n",
" N = T.dynamic(\"N\")\n",
" A: T.Tensor[[N], T.float32]\n",
Expand Down Expand Up @@ -857,7 +857,7 @@
" n31(x * 3 + 1, var)\n",
"\n",
"\n",
"@tilelang.lazy_jit\n",
"@tilelang.jit\n",
"def foo(A: T.Tensor[[1], T.int32], n: int):\n",
" with T.Kernel(1) as _:\n",
" n31(n, A[0])"
Expand Down
28 changes: 14 additions & 14 deletions testing/python/language/test_tilelang_language_lazy_jit.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@


def test_jit2_gemm():
@tilelang.lazy_jit(verbose=True)
@tilelang.jit(verbose=True)
def gemm(
A,
B,
Expand Down Expand Up @@ -45,7 +45,7 @@ def gemm(


def test_jit2_gemm_ptr():
@tilelang.lazy_jit
@tilelang.jit
def gemm_ptr(
A: T.ptr,
B: T.ptr,
Expand Down Expand Up @@ -102,43 +102,43 @@ def copy_impl(A, B):
with T.Kernel(T.ceildiv(M, 128), T.ceildiv(N, 128), threads=128) as (bx, by):
T.copy(A[bx * 128 : bx * 128 + 128, by * 128 : by * 128 + 128], B[bx * 128 : bx * 128 + 128, by * 128 : by * 128 + 128])

@tilelang.lazy_jit
@tilelang.jit
def copy1(A, B):
N, M = T.const("N, M")
A: T.Tensor[[N, M], T.float32]
B: T.Tensor[[N, M], T.float32]
copy_impl(A, B)

@tilelang.lazy_jit
@tilelang.jit
def copy2(
A: T.Tensor[[128, 128], T.float32],
B: T.Tensor[[128, 128], T.float32],
):
copy_impl(A, B)

@tilelang.lazy_jit
@tilelang.jit
def copy3(A, B):
N = T.const("N")
A: T.Tensor[[N, 128], T.float32]
B: T.Tensor[[N, 128], T.float32]
copy_impl(A, B)

@tilelang.lazy_jit
@tilelang.jit
def copy4(A, B):
N = T.dynamic("N")
M = T.const("M")
A: T.Tensor[[N, M], T.float32]
B: T.Tensor[[N, M], T.float32]
copy_impl(A, B)

@tilelang.lazy_jit
@tilelang.jit
def copy5(A, B):
N, M, N_, M_ = T.const("N, M, N_, M_")
A: T.StridedTensor[[N, M], [N_, M_], T.float32]
B: T.StridedTensor[[N, M], [N_, M_], T.float32]
copy_impl(A, B)

@tilelang.lazy_jit
@tilelang.jit
def copy6(A, B):
N = T.dynamic("N")
M, N_, M_ = T.const("M, N_, M_")
Expand Down Expand Up @@ -175,37 +175,37 @@ def copy_impl(A):
T.copy(A[bx * 128 : bx * 128 + 128, by * 128 : by * 128 + 128], B[bx * 128 : bx * 128 + 128, by * 128 : by * 128 + 128])
return B

@tilelang.lazy_jit
@tilelang.jit
def copy1(A):
M, N = T.const("M, N")
A: T.Tensor[[M, N], T.float32]
return copy_impl(A)

@tilelang.lazy_jit
@tilelang.jit
def copy2(A):
A: T.Tensor[[128, 128], T.float32]
return copy_impl(A)

@tilelang.lazy_jit
@tilelang.jit
def copy3(A):
N = T.const("N")
A: T.Tensor[[N, 128], T.float32]
return copy_impl(A)

@tilelang.lazy_jit
@tilelang.jit
def copy4(A):
N = T.dynamic("N")
M = T.const("M")
A: T.Tensor[[N, M], T.float32]
return copy_impl(A)

@tilelang.lazy_jit
@tilelang.jit
def copy5(A):
N, M, N_, M_ = T.const("N, M, N_, M_")
A: T.StridedTensor[[N, M], [N_, M_], T.float32]
return copy_impl(A)

@tilelang.lazy_jit
@tilelang.jit
def copy6(A):
N = T.dynamic("N")
M, N_, M_ = T.const("M, N_, M_")
Expand Down
6 changes: 3 additions & 3 deletions testing/python/layout/test_tilelang_annotate_loop_layout.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@


# TODO(lei): replicate loop layout and more complicated layout cases
@tilelang.lazy_jit
@tilelang.jit
def loop_layout_kernel(A, B, loop_layout):
M, N = T.const("M, N")
A: T.Tensor[(M, N), T.float32]
Expand Down Expand Up @@ -51,7 +51,7 @@ def loop_layout_fn(i, j):
assert "*(float4*)(B + ((((int)threadIdx.x) * 32) + (i * 4))) = *(float4*)(A + ((((int)threadIdx.x) * 32) + (i * 4)));" in code


@tilelang.lazy_jit
@tilelang.jit
def copy_with_layout_kernel(A, B, loop_layout):
M, N = T.const("M, N")
A: T.Tensor[(M, N), T.float32]
Expand Down Expand Up @@ -79,7 +79,7 @@ def loop_layout_fn(i, j, rep):
assert "*(float4*)(B + ((i * 512) + (((int)threadIdx.x) * 4))) = *(float4*)(A + ((i * 512) + (((int)threadIdx.x) * 4)));" in code


@tilelang.lazy_jit
@tilelang.jit
def replicate_loop_layout_kernel(A, B, loop_layout):
M, N = T.const("M, N")
A: T.Tensor[(M, N), T.float32]
Expand Down
2 changes: 1 addition & 1 deletion tilelang/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -142,7 +142,7 @@ def _load_tile_lang_lib():
if env.SKIP_LOADING_TILELANG_SO == "0":
_LIB, _LIB_PATH = _load_tile_lang_lib()

from .jit import jit, lazy_jit, JITKernel, compile, par_compile # noqa: F401
from .jit import jit, JITKernel, compile, par_compile # noqa: F401
from .profiler import Profiler # noqa: F401
from .cache import clear_cache # noqa: F401
from .utils import (
Expand Down
Loading
Loading