fix roll on zero-sized tensors (#17603)

Signed-off-by: Bennett <[email protected]>
Co-authored-by: Bennett <[email protected]>
This commit is contained in:
Bennett
2026-08-19 15:31:21 -04:00
committed by GitHub
co-authored by Bennett
parent 0c5307b4f3
commit 7064e76bc8
2 changed files with 5 additions and 0 deletions
+4
View File
@@ -2164,6 +2164,10 @@ class TestOps(unittest.TestCase):
def test_roll(self):
helper_test_op([(2, 4)], lambda x: x.roll(1))
helper_test_op([(2, 4)], lambda x: x.roll((1,)))
helper_test_op([(0,)], lambda x: x.roll(1, 0))
helper_test_op([(2, 0, 3)], lambda x: x.roll(1, 0))
helper_test_op([(2, 0, 3)], lambda x: x.roll(1, 1))
helper_test_op([(2, 0, 3)], lambda x: x.roll(1))
self.helper_test_exception([(2, 4)], lambda x: x.roll((1, 2)), expected=RuntimeError)
helper_test_op([(2, 4)], lambda x: x.roll(1, 0))
helper_test_op([(2, 4)], lambda x: x.roll(-1, 0))
+1
View File
@@ -550,6 +550,7 @@ class MovementMixin:
if dims is None: return self.flatten().roll(shifts, 0).reshape(self.shape)
dims, shifts = tuple(self._resolve_dim(d) for d in make_tuple(dims, 1)), make_tuple(shifts, 1)
if len(dims) != len(shifts): raise RuntimeError(f"{len(dims)=} != {len(shifts)=}")
if 0 in self.shape: return self
shrink_arg: list[tuple[sint, sint]|None] = [None] * self.ndim
for d, s in zip(dims, shifts): shrink_arg[d] = (delta:=self.shape[d]-s%self.shape[d], delta+self.shape[d])
return self.repeat(*tuple(2 if i in dims else 1 for i in range(self.ndim))).shrink(tuple(shrink_arg))