Skip to content

Commit 7252103

Browse files
authored
fix index and index_put_ oom bug (#1217)
1 parent f56417e commit 7252103

2 files changed

Lines changed: 3 additions & 13 deletions

File tree

src/flag_gems/ops/index.py

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -57,7 +57,6 @@ def generate_index_kernel(
5757
with code.indent():
5858
code.writeline('configs=runtime.get_tuned_config("index"),')
5959
code.writeline('key=["M", "N"],')
60-
code.writeline('restore_value=["input_ptr"],')
6160
code.writeline('strategy=["align32", "align32"],')
6261
code.writeline("warmup=5,")
6362
code.writeline("rep=10,")

src/flag_gems/ops/index_put.py

Lines changed: 3 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)