Skip to content

Commit 7b6b64f

Browse files
authored
[Enhancement] Add eager-mode support for tilelang.autotune (#1906)
* [Enhancement] Add eager-mode support for tilelang.autotune * Fix type annotation and improve code structure in tuner.py * Fix misleading docstring in autotune tests
1 parent ccfc127 commit 7b6b64f

5 files changed

Lines changed: 233 additions & 32 deletions

File tree

‎testing/python/autotune/test_tilelang_autotune.py‎

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -11,14 +11,14 @@ def ref_program(A, B):
1111
1212
Parameters
1313
----------
14-
A : numpy.ndarray
14+
A : torch.Tensor
1515
The matrix with shape (M, K).
16-
B : numpy.ndarray
16+
B : torch.Tensor
1717
The matrix with shape (N, K).
1818
1919
Returns
2020
-------
21-
np.ndarray
21+
torch.Tensor
2222
The result of A @ B.T, shape (M, N).
2323
"""
2424
return A @ B.T
Lines changed: 141 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,141 @@
1+
import itertools
2+
import logging
3+
import tilelang
4+
import tilelang.testing
5+
from tilelang.autotuner import set_autotune_inputs
6+
import tilelang.language as T
7+
8+
# Configure logger
9+
logger = logging.getLogger(__name__)
10+
logger.setLevel(logging.DEBUG)
11+
12+
13+
def ref_program(A, B):
14+
"""
15+
A reference matrix multiplication program, used to compare performance.
16+
17+
Parameters
18+
----------
19+
A : torch.Tensor
20+
The matrix with shape (M, K).
21+
B : torch.Tensor
22+
The matrix with shape (N, K).
23+
24+
Returns
25+
-------
26+
torch.Tensor
27+
The result of A @ B.T, shape (M, N).
28+
"""
29+
return A @ B.T
30+
31+
32+
def get_configs():
33+
iter_params = dict(block_M=[64], block_N=[64], block_K=[32], num_stages=[0, 1], thread_num=[128], enable_rasterization=[False])
34+
return [{k: v for k, v in zip(iter_params, values)} for values in itertools.product(*iter_params.values())]
35+
36+
37+
@tilelang.autotune(
38+
configs=get_configs(),
39+
)
40+
@tilelang.jit
41+
def matmul(A, B, block_M=128, block_N=128, block_K=32, num_stages=0, thread_num=128, enable_rasterization=False):
42+
M, N, K = T.const("M, N, K")
43+
44+
dtype = T.float16
45+
accum_dtype = T.float32
46+
47+
A: T.Tensor((M, K), dtype)
48+
B: T.Tensor((N, K), dtype)
49+
C = T.empty((M, N), dtype)
50+
51+
# Bind x-dimension to block index in N,
52+
# y-dimension to block index in M.
53+
with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M), threads=thread_num) as (bx, by):
54+
# Allocate shared memory for A sub-block of shape (block_M, block_K)
55+
A_shared = T.alloc_shared((block_M, block_K), dtype)
56+
# Allocate shared memory for B sub-block of shape (block_N, block_K)
57+
B_shared = T.alloc_shared((block_N, block_K), dtype)
58+
# Allocate a local fragment for intermediate accumulation
59+
C_local = T.alloc_fragment((block_M, block_N), accum_dtype)
60+
61+
# Enable (or disable) swizzling optimization
62+
T.use_swizzle(panel_size=10, enable=enable_rasterization)
63+
64+
# Clear out the accumulation buffer
65+
T.clear(C_local)
66+
67+
# Loop over sub-blocks in K dimension, pipelined by num_stages
68+
for k in T.Pipelined(T.ceildiv(K, block_K), num_stages=num_stages):
69+
# Load a sub-block of A from global memory into A_shared
70+
T.copy(
71+
A[by * block_M, k * block_K],
72+
A_shared,
73+
)
74+
# Load a sub-block of B from global memory into B_shared
75+
T.copy(
76+
B[bx * block_N, k * block_K],
77+
B_shared,
78+
)
79+
# Perform a partial matrix multiplication:
80+
# C_local += A_shared @ B_shared^T
81+
T.gemm(
82+
A_shared,
83+
B_shared,
84+
C_local,
85+
transpose_B=True,
86+
)
87+
# Write back the results from C_local to the global memory C
88+
T.copy(C_local, C[by * block_M, bx * block_N])
89+
90+
return C
91+
92+
93+
def run_autotune(M, N, K, M_value=None, N_value=None, K_value=None, return_kernel=False):
94+
import torch
95+
96+
def _resolve(dim, provided, name):
97+
if isinstance(dim, T.Var):
98+
if provided is None:
99+
raise ValueError(f"Dynamic dimension {name} requires a concrete value.")
100+
return provided
101+
return dim
102+
103+
actual_M = _resolve(M, M_value, "M")
104+
actual_N = _resolve(N, N_value, "N")
105+
actual_K = _resolve(K, K_value, "K")
106+
107+
a = torch.randn(actual_M, actual_K, dtype=torch.float16).cuda()
108+
b = torch.randn(actual_N, actual_K, dtype=torch.float16).cuda()
109+
110+
if return_kernel:
111+
with set_autotune_inputs([a, b]):
112+
kernel = matmul.compile(M=M, N=N, K=K)
113+
c = kernel(a, b)
114+
else:
115+
with set_autotune_inputs([a, b]):
116+
c = matmul(a, b)
117+
118+
ref_c = ref_program(a, b)
119+
torch.testing.assert_close(c, ref_c, rtol=1e-2, atol=1e-2)
120+
121+
122+
def test_autotune_matmul():
123+
"""
124+
Run the autotuning validation for the matmul kernel on a 1024x1024x1024 problem.
125+
126+
This test constructs random CUDA tensors, autotunes the JIT-compiled block-level matrix-multiplication kernel,
127+
executes it, and asserts the result matches a reference PyTorch implementation within tolerances.
128+
"""
129+
run_autotune(1024, 1024, 1024)
130+
131+
132+
def test_autotune_matmul_compile():
133+
run_autotune(1024, 1024, 1024, return_kernel=True)
134+
135+
136+
def test_autotune_matmul_symbolic_m():
137+
run_autotune(T.symbolic("m"), 1024, 1024, M_value=1024)
138+
139+
140+
if __name__ == "__main__":
141+
tilelang.testing.main()

‎testing/python/autotune/test_tilelang_autotune_with_inputs.py‎

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -16,14 +16,14 @@ def ref_program(A, B):
1616
1717
Parameters
1818
----------
19-
A : numpy.ndarray
19+
A : torch.Tensor
2020
The matrix with shape (M, K).
21-
B : numpy.ndarray
21+
B : torch.Tensor
2222
The matrix with shape (N, K).
2323
2424
Returns
2525
-------
26-
np.ndarray
26+
torch.Tensor
2727
The result of A @ B.T, shape (M, N).
2828
"""
2929
return A @ B.T
@@ -131,7 +131,7 @@ def test_autotune_matmul():
131131
Run the autotuning validation for the matmul kernel on a 1024x1024x1024 problem.
132132
133133
This test constructs random CUDA tensors, autotunes the JIT-compiled block-level matrix-multiplication kernel,
134-
executes it, and asserts the result matches a reference CPU implementation within tolerances.
134+
executes it, and asserts the result matches a reference PyTorch implementation within tolerances.
135135
"""
136136
run_autotune(1024, 1024, 1024)
137137

‎tilelang/autotuner/param.py‎

Lines changed: 22 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,7 @@
2323

2424
BEST_CONFIG_PATH = "best_config.json"
2525
FUNCTION_PATH = "function.pkl"
26+
OUT_IDX_PATH = "out_idx.json"
2627
LATENCY_PATH = "latency.json"
2728

2829
# Align file names with cache/kernel_cache.py
@@ -378,6 +379,20 @@ def save_to_disk(self, path: Path, verbose: bool = False):
378379
logger.debug(f"Saving function to file: {path / FUNCTION_PATH}")
379380
self._safe_write_file(str(path / FUNCTION_PATH), "wb", lambda f: cloudpickle.dump(self.func, f))
380381

382+
# save out idx (atomic)
383+
if verbose:
384+
logger.debug(f"Saving out idx to file: {path / OUT_IDX_PATH}")
385+
self._safe_write_file(
386+
str(path / OUT_IDX_PATH),
387+
"w",
388+
lambda f: json.dump(
389+
{
390+
"out_idx": getattr(self.func, "out_idx_override", None),
391+
},
392+
f,
393+
),
394+
)
395+
381396
# save ref latency (atomic)
382397
if verbose:
383398
logger.debug(f"Saving latency to file: {path / LATENCY_PATH}")
@@ -421,6 +436,12 @@ def load_from_disk(cls, path: Path, compile_args: CompileArgs) -> AutotuneResult
421436
with open(path / FUNCTION_PATH, "rb") as f:
422437
func = cloudpickle.load(f)
423438

439+
# load out idx
440+
if verbose:
441+
logger.debug(f"Loading out idx from file: {path / OUT_IDX_PATH}")
442+
with open(path / OUT_IDX_PATH) as f:
443+
out_idx_override = json.load(f)["out_idx"]
444+
424445
# load latency
425446
if verbose:
426447
logger.debug(f"Loading latency from file: {path / LATENCY_PATH}")
@@ -433,7 +454,7 @@ def load_from_disk(cls, path: Path, compile_args: CompileArgs) -> AutotuneResult
433454
path,
434455
norm_target,
435456
compile_args.target_host,
436-
compile_args.out_idx,
457+
out_idx_override if out_idx_override is not None else compile_args.out_idx,
437458
resolved_backend,
438459
compile_args.pass_configs,
439460
None, # compile_flags not tracked here

‎tilelang/autotuner/tuner.py‎

Lines changed: 63 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -99,6 +99,21 @@ def get_available_cpu_count() -> int:
9999
return cpu_count or 1
100100

101101

102+
def _normalize_value(value, sort_dict_items: bool = False):
103+
if isinstance(value, torch.Tensor):
104+
return ("tensor", str(value.dtype), tuple(value.shape), value.stride())
105+
if isinstance(value, Var):
106+
return str(value)
107+
if isinstance(value, (list, tuple)):
108+
return tuple(_normalize_value(v, sort_dict_items=sort_dict_items) for v in value)
109+
if isinstance(value, dict):
110+
items = ((str(k), _normalize_value(v, sort_dict_items=sort_dict_items)) for k, v in value.items())
111+
if sort_dict_items:
112+
return tuple(sorted(items))
113+
return {k: v for k, v in items}
114+
return value
115+
116+
102117
class AutoTuner:
103118
"""Auto-tuner for tilelang programs.
104119
@@ -113,7 +128,7 @@ class AutoTuner:
113128
compile_args = CompileArgs()
114129
profile_args = ProfileArgs()
115130

116-
_kernel_parameters: tuple[str, ...] | None = None
131+
_kernel_parameters: tuple[tuple[Any, ...], tuple[tuple[str, Any], ...]] | None = None
117132
_function_parameters: dict[str, Any] | None = None
118133
_lock = threading.Lock() # For thread safety
119134
_memory_cache = {} # In-memory cache dictionary
@@ -259,31 +274,22 @@ def set_profile_args(
259274

260275
return self
261276

262-
def set_kernel_parameters(self, k_parameters: tuple[str, ...], f_parameters: dict[str, Any]):
277+
def set_kernel_parameters(self, k_parameters: tuple[tuple[Any, ...], tuple[tuple[str, Any], ...]], f_parameters: dict[str, Any]):
263278
# for cache key generation
264279
self._kernel_parameters = k_parameters
265280
self._function_parameters = f_parameters
266281

267282
def generate_cache_key(self, parameters: dict[str, Any], extra_parameters: dict[str, Any]) -> AutotuneResult | None:
268283
"""Generate a cache key for the auto-tuning process."""
269284

270-
def _normalize_param(value):
271-
if isinstance(value, Var):
272-
return str(value)
273-
if isinstance(value, (list, tuple)):
274-
return [_normalize_param(v) for v in value]
275-
if isinstance(value, dict):
276-
return {str(k): _normalize_param(v) for k, v in value.items()}
277-
return value
278-
279285
# extract parameters from the function signature
280286
op_parameters = []
281287
for _, default_value in parameters.items():
282288
if default_value.default is not inspect.Parameter.empty:
283289
op_parameters.append(default_value.default)
284290

285291
if self._kernel_parameters is not None:
286-
op_parameters += _normalize_param(self._kernel_parameters)
292+
op_parameters += _normalize_value(self._kernel_parameters)
287293

288294
func_source = inspect.getsource(self.fn)
289295
key_data = {
@@ -342,7 +348,9 @@ def run(self, warmup: int = 25, rep: int = 100, timeout: int = 30):
342348
extra_parameters[var_name] = cell.cell_contents
343349

344350
if isinstance(self.configs, Callable):
345-
self.configs = self.configs(*self._kernel_parameters)
351+
kernel_args, kernel_kwargs = self._kernel_parameters
352+
kernel_kwargs = dict(kernel_kwargs)
353+
self.configs = self.configs(*kernel_args, **kernel_kwargs)
346354

347355
key = self.generate_cache_key(parameters, extra_parameters)
348356

@@ -680,21 +688,52 @@ def get_tunner(self):
680688
autotuner.run = partial(autotuner.run, self.warmup, self.rep, self.timeout)
681689
return autotuner
682690

683-
def __call__(self, *args: _P.args, **kwargs: _P.kwargs) -> JITKernel:
684-
key_args_tuple = args
685-
key_kwargs_tuple = tuple(sorted(kwargs.items()))
686-
key = (key_args_tuple, key_kwargs_tuple)
691+
def __call__(self, *args: _P.args, **kwargs: _P.kwargs) -> JITKernel | _T:
692+
return_kernel = kwargs.pop("__return_kernel", False)
693+
694+
mode = self.jit_impl.initialize_jit_mode(*args, **kwargs)
695+
696+
norm_args = _normalize_value(args, sort_dict_items=True)
697+
norm_kwargs = _normalize_value(kwargs, sort_dict_items=True)
698+
key = (norm_args, norm_kwargs)
687699
if key not in self._tuner_cache:
700+
if mode == "lazy":
701+
702+
def jit_compile(**config_arg):
703+
return self.jit_impl(*args, **kwargs, __tune_params=config_arg)
688704

689-
def jit_compile(**config_arg):
690-
return self.jit_impl(*args, **kwargs, __tune_params=config_arg)
705+
autotuner = self.get_tunner()
706+
autotuner.jit_compile = jit_compile
707+
autotuner.set_kernel_parameters(key, self.jit_impl.signature.parameters)
708+
else:
709+
710+
def jit_compile(**config_arg):
711+
merged = dict(kwargs)
712+
merged.update(config_arg)
713+
return self.jit_impl.compile(*args, **merged)
714+
715+
autotuner = self.get_tunner()
716+
autotuner.jit_compile = jit_compile
717+
autotuner.set_kernel_parameters(key, self.jit_impl.signature.parameters)
691718

692-
autotuner = self.get_tunner()
693-
autotuner.jit_compile = jit_compile
694-
autotuner.set_kernel_parameters(key, self.jit_impl.signature.parameters)
695719
artifact = autotuner.run()
696-
self._tuner_cache[key] = artifact.kernel
697-
return self._tuner_cache[key]
720+
self._tuner_cache[key] = artifact.kernel, artifact.config
721+
722+
best_kernel, best_config = self._tuner_cache[key]
723+
724+
if mode == "lazy":
725+
return best_kernel
726+
else:
727+
if return_kernel:
728+
return best_kernel
729+
exec_kwargs = dict(kwargs)
730+
if best_config is not None:
731+
exec_kwargs.update(best_config)
732+
_, kernel_args = self.jit_impl.func.parse_args(*args, **exec_kwargs)
733+
return best_kernel(*kernel_args.values())
734+
735+
def compile(self, *args: _P.args, **kwargs: _P.kwargs) -> JITKernel:
736+
return self(*args, **kwargs, __return_kernel=True)
698737

699738

700739
def autotune( # This is the new public interface

0 commit comments

Comments
 (0)