Skip to content
Merged
Changes from 1 commit
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
Prev Previous commit
Next Next commit
Fix type annotation and improve code structure in tuner.py
  • Loading branch information
ColmaLiu committed Mar 10, 2026
commit 4a0783aa63a9737344e1b0ba2ac1790ea75c55a6
48 changes: 21 additions & 27 deletions tilelang/autotuner/tuner.py
Original file line number Diff line number Diff line change
Expand Up @@ -99,6 +99,21 @@ def get_available_cpu_count() -> int:
return cpu_count or 1


def _normalize_value(value, sort_dict_items: bool = False):
if isinstance(value, torch.Tensor):
return ("tensor", str(value.dtype), tuple(value.shape), value.stride())
if isinstance(value, Var):
return str(value)
if isinstance(value, (list, tuple)):
return tuple(_normalize_value(v, sort_dict_items=sort_dict_items) for v in value)
if isinstance(value, dict):
items = ((str(k), _normalize_value(v, sort_dict_items=sort_dict_items)) for k, v in value.items())
if sort_dict_items:
return tuple(sorted(items))
return {k: v for k, v in items}
return value


class AutoTuner:
"""Auto-tuner for tilelang programs.

Expand All @@ -113,7 +128,7 @@ class AutoTuner:
compile_args = CompileArgs()
profile_args = ProfileArgs()

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

return self

def set_kernel_parameters(self, k_parameters: tuple[str, ...], f_parameters: dict[str, Any]):
def set_kernel_parameters(self, k_parameters: tuple[tuple[Any, ...], tuple[tuple[str, Any], ...]], f_parameters: dict[str, Any]):
# for cache key generation
self._kernel_parameters = k_parameters
self._function_parameters = f_parameters

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

def _normalize_param(value):
if isinstance(value, Var):
return str(value)
if isinstance(value, (list, tuple)):
return [_normalize_param(v) for v in value]
if isinstance(value, dict):
return {str(k): _normalize_param(v) for k, v in value.items()}
return value

# extract parameters from the function signature
op_parameters = []
for _, default_value in parameters.items():
if default_value.default is not inspect.Parameter.empty:
op_parameters.append(default_value.default)

if self._kernel_parameters is not None:
op_parameters += _normalize_param(self._kernel_parameters)
op_parameters += _normalize_value(self._kernel_parameters)

func_source = inspect.getsource(self.fn)
key_data = {
Expand Down Expand Up @@ -687,19 +693,8 @@ def __call__(self, *args: _P.args, **kwargs: _P.kwargs) -> JITKernel | _T:

mode = self.jit_impl.initialize_jit_mode(*args, **kwargs)

def _normalize_for_key(x):
import torch

if isinstance(x, torch.Tensor):
return ("tensor", str(x.dtype), tuple(x.shape))
if isinstance(x, (list, tuple)):
return tuple(_normalize_for_key(v) for v in x)
if isinstance(x, dict):
return tuple(sorted((k, _normalize_for_key(v)) for k, v in x.items()))
return x

norm_args = _normalize_for_key(args)
norm_kwargs = tuple(sorted((k, _normalize_for_key(v)) for k, v in kwargs.items()))
norm_args = _normalize_value(args, sort_dict_items=True)
norm_kwargs = _normalize_value(kwargs, sort_dict_items=True)
key = (norm_args, norm_kwargs)
Comment thread
ColmaLiu marked this conversation as resolved.
if key not in self._tuner_cache:
if mode == "lazy":
Expand All @@ -709,8 +704,7 @@ def jit_compile(**config_arg):

autotuner = self.get_tunner()
autotuner.jit_compile = jit_compile
raw_key = (args, tuple(sorted(kwargs.items())))
autotuner.set_kernel_parameters(raw_key, self.jit_impl.signature.parameters)
autotuner.set_kernel_parameters(key, self.jit_impl.signature.parameters)
else:

def jit_compile(**config_arg):
Expand Down