@@ -38,7 +38,7 @@ def generate_imports(code: IndentedBuffer) -> IndentedBuffer:
3838 code .writeline ("import triton" )
3939 code .writeline ("import triton.language as tl" )
4040 code .newline ()
41- code .writeline ("from flag_gems.utils import libentry, libtuner " )
41+ code .writeline ("from flag_gems.utils import libentry" )
4242 code .writeline ("from flag_gems import runtime" )
4343 code .writeline ("from flag_gems.utils.shape_utils import volume" )
4444 code .writeline ("from flag_gems.utils import triton_lang_extension as tle" )
@@ -52,15 +52,6 @@ def generate_index_put_kernel(
5252 inp_rank , indices_len , index_rank , kernel_name : str , code : IndentedBuffer
5353):
5454 code .writeline ("@libentry()" )
55- code .writeline ("@libtuner(" )
56- with code .indent ():
57- code .writeline ('configs=runtime.get_tuned_config("index_put"),' )
58- code .writeline ('key=["M", "N"],' )
59- code .writeline ('restore_value=["input_ptr"],' )
60- code .writeline ('strategy=["align32", "align32"],' )
61- code .writeline ("warmup=5," )
62- code .writeline ("rep=10," )
63- code .writeline (")" )
6455 code .writeline ("@triton.jit" )
6556 code .writeline (f"def { kernel_name } (" )
6657 with code .indent ():
@@ -80,8 +71,8 @@ def generate_index_put_kernel(
8071 "M," ,
8172 "N," ,
8273 "IS_ACCUMULATE: tl.constexpr," ,
83- "BLOCK_SIZE0: tl.constexpr," ,
84- "BLOCK_SIZE1: tl.constexpr," ,
74+ "BLOCK_SIZE0: tl.constexpr = 2 ," ,
75+ "BLOCK_SIZE1: tl.constexpr = 2048 ," ,
8576 ]
8677 code .writelines (args )
8778 code .writeline ("):" )
0 commit comments