From f1efd84c92abe2ea242ed4f5750f5c9e730efccb Mon Sep 17 00:00:00 2001 From: chenyu Date: Sat, 16 Nov 2024 10:15:29 -0500 Subject: [PATCH] fix repeat_interleave with negative dim (#7734) --- test/test_ops.py | 2 ++ tinygrad/tensor.py | 2 +- 2 files changed, 3 insertions(+), 1 deletion(-) diff --git a/test/test_ops.py b/test/test_ops.py index e121c44687..44f5ce6864 100644 --- a/test/test_ops.py +++ b/test/test_ops.py @@ -2047,6 +2047,8 @@ class TestOps(unittest.TestCase): helper_test_op([(3, 3)], lambda x: x.repeat_interleave(6)) helper_test_op([(3, 3)], lambda x: x.repeat_interleave(2, 1)) helper_test_op([(3, 3)], lambda x: x.repeat_interleave(2, 0)) + helper_test_op([(3, 3)], lambda x: x.repeat_interleave(2, -1)) + helper_test_op([(3, 3)], lambda x: x.repeat_interleave(2, -2)) def test_simple_repeat(self): repeats = [3, 3, 4] diff --git a/tinygrad/tensor.py b/tinygrad/tensor.py index ccbe0cae93..da6202ae0a 100644 --- a/tinygrad/tensor.py +++ b/tinygrad/tensor.py @@ -1279,7 +1279,7 @@ class Tensor(SimpleMathTrait): # pylint: disable=abstract-method print(t.repeat_interleave(2).numpy()) ``` """ - x, dim = (self.flatten(), 0) if dim is None else (self, dim) + x, dim = (self.flatten(), 0) if dim is None else (self, self._resolve_dim(dim)) shp = x.shape return x.reshape(*shp[:dim+1], 1, *shp[dim+1:]).expand(*shp[:dim+1], repeats, *shp[dim+1:]).reshape(*shp[:dim], shp[dim]*repeats, *shp[dim+1:])