Skip to content

Commit c6b714e

Browse files
authored
fix repeat_interleave bug & update benchmark setting for inplace operators (flagos-ai#1227)
1 parent a0c2861 commit c6b714e

2 files changed

Lines changed: 3 additions & 8 deletions

File tree

benchmark/performance_utils.py

Lines changed: 1 addition & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -262,12 +262,7 @@ def set_gems(self, gems_op):
262262
self.gems_op = gems_op
263263

264264
def get_latency(self, op, *args, **kwargs):
265-
if self.is_inplace:
266-
fn = lambda: op(
267-
*[x.clone() if torch.is_tensor(x) else x for x in args], **kwargs
268-
)
269-
else:
270-
fn = lambda: op(*args, **kwargs)
265+
fn = lambda: op(*args, **kwargs)
271266
if self.is_backward:
272267
out = fn()
273268
dout = torch.randn_like(out)

tests/test_special_ops.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1004,7 +1004,7 @@ def test_accuracy_repeat_interleave_self_int(shape, dim, dtype):
10041004

10051005
ref_out = torch.repeat_interleave(ref_inp, repeats, dim)
10061006
with flag_gems.use_gems():
1007-
res_out = torch.repeat_interleave(ref_inp, repeats, dim)
1007+
res_out = torch.repeat_interleave(inp, repeats, dim)
10081008
gems_assert_equal(res_out, ref_out)
10091009

10101010

@@ -1019,7 +1019,7 @@ def test_accuracy_repeat_interleave_self_int_non_contiguous(shape, dim, dtype):
10191019

10201020
ref_out = torch.repeat_interleave(ref_inp, repeats, dim)
10211021
with flag_gems.use_gems():
1022-
res_out = torch.repeat_interleave(ref_inp, repeats, dim)
1022+
res_out = torch.repeat_interleave(inp, repeats, dim)
10231023
gems_assert_equal(res_out, ref_out)
10241024

10251025

0 commit comments

Comments
 (0)