Skip to content

Commit ad17615

Browse files
committed
Fix roll empty tensor dim handling
1 parent c13751c commit ad17615

2 files changed

Lines changed: 34 additions & 3 deletions

File tree

src/flag_gems/ops/roll.py

Lines changed: 13 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -30,7 +30,11 @@ def roll(inp: torch.Tensor, shifts, dims=None) -> torch.Tensor:
3030
return _candidate_triton_single_dim(
3131
inp,
3232
_as_tuple(shifts)[0],
33-
_canonicalize_dim(dim_values[0], inp.dim()),
33+
_canonicalize_dim(
34+
dim_values[0],
35+
inp.dim(),
36+
allow_empty_wrap=inp.numel() == 0,
37+
),
3438
)
3539
return _candidate_triton(inp, shifts, dims)
3640
return _candidate_fallback(inp, shifts, dims)
@@ -145,7 +149,11 @@ def _candidate_fallback(
145149

146150
result = inp
147151
for shift, dim in zip(shift_values, _as_tuple(dims)):
148-
result = _roll_along_dim(result, shift, _canonicalize_dim(dim, inp.dim()))
152+
result = _roll_along_dim(
153+
result,
154+
shift,
155+
_canonicalize_dim(dim, inp.dim(), allow_empty_wrap=inp.numel() == 0),
156+
)
149157
return result
150158

151159

@@ -189,9 +197,11 @@ def _roll_along_dim(inp: torch.Tensor, shift: int, dim: int) -> torch.Tensor:
189197
)
190198

191199

192-
def _canonicalize_dim(dim: int, ndim: int) -> int:
200+
def _canonicalize_dim(dim: int, ndim: int, allow_empty_wrap: bool = False) -> int:
193201
if ndim == 0:
194202
raise IndexError(f"Dimension specified as {dim} but tensor has no dimensions")
203+
if allow_empty_wrap:
204+
return dim % ndim
195205
if dim < -ndim or dim >= ndim:
196206
raise IndexError(
197207
f"Dimension out of range (expected to be in range of [{-ndim}, {ndim - 1}], but got {dim})"

tests/test_unary_pointwise_ops.py

Lines changed: 21 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2780,6 +2780,27 @@ def test_roll_empty_input(shape, dtype):
27802780
gems_assert_equal(res_out, ref_out)
27812781

27822782

2783+
@pytest.mark.roll
2784+
@pytest.mark.parametrize(
2785+
"shape, shifts, dims",
2786+
[
2787+
((0,), 1, 5),
2788+
((0,), 1, -5),
2789+
((2, 0, 3), 1, 5),
2790+
((2, 0, 3), (1, 2), (5, 9)),
2791+
],
2792+
)
2793+
def test_roll_empty_input_with_out_of_range_dims(shape, shifts, dims):
2794+
inp = torch.empty(shape, dtype=torch.float32, device=flag_gems.device)
2795+
ref_inp = to_reference(inp, False)
2796+
2797+
ref_out = torch.roll(ref_inp, shifts, dims)
2798+
with flag_gems.use_gems():
2799+
res_out = torch.roll(inp, shifts, dims)
2800+
2801+
gems_assert_equal(res_out, ref_out)
2802+
2803+
27832804
@pytest.mark.roll
27842805
def test_roll_scalar_flatten():
27852806
inp = torch.tensor(7.0, dtype=torch.float32, device=flag_gems.device)

0 commit comments

Comments
 (0)