jax-ml/jax · error · ValueError
'shift' and 'axis' arguments to roll must be scalars or 1D a
Error message
'shift' and 'axis' arguments to roll must be scalars or 1D arrays
What it means
Raised inside jnp.roll's dynamic path when broadcasting shift against axis yields more than one dimension — i.e. shift and/or axis are arrays whose broadcast shape is not 1-D (scalars or 1-D sequences are required so each shift pairs with one axis).
Source
Thrown at jax/_src/numpy/lax_numpy.py:8487
return _nanargmin(a, None if axis is None else operator.index(axis), keepdims=bool(keepdims))
@api.jit(static_argnames=('axis', 'keepdims'))
def _nanargmin(a: Array, axis: int | None = None, keepdims : bool = False):
if not issubdtype(a.dtype, np.inexact):
return argmin(a, axis=axis, keepdims=keepdims)
nan_mask = ufuncs.isnan(a)
a = where(nan_mask, np.inf, a)
res = argmin(a, axis=axis, keepdims=keepdims)
return where(reductions.all(nan_mask, axis=axis, keepdims=keepdims), -1, res)
@api.jit(static_argnums=(2,))
def _roll_dynamic(a: Array, shift: Array, axis: Sequence[int]) -> Array:
b_shape = lax.broadcast_shapes(shift.shape, np.shape(axis))
if len(b_shape) != 1:
msg = "'shift' and 'axis' arguments to roll must be scalars or 1D arrays"
raise ValueError(msg)
for x, i in zip(broadcast_to(shift, b_shape),
np.broadcast_to(axis, b_shape)): # pyrefly: ignore[no-matching-overload]
a_shape_i = array(a.shape[i], dtype=np.int32)
x = ufuncs.remainder(lax.convert_element_type(x, np.int32),
lax.max(a_shape_i, np.int32(1)))
a_concat = lax.concatenate((a, a), i)
a = lax_slicing.dynamic_slice_in_dim(a_concat, a_shape_i - x, a.shape[i], axis=i)
return a
@api.jit(static_argnums=(1, 2))
def _roll_static(a: Array, shift: Sequence[int], axis: Sequence[int]) -> Array:
for ax, s in zip(*np.broadcast_arrays(axis, shift)):
if a.shape[ax] == 0:
continue
i = (-s) % a.shape[ax]
a = lax.concatenate([lax_slicing.slice_in_dim(a, i, a.shape[ax], axis=ax),
lax_slicing.slice_in_dim(a, 0, i, axis=ax)],View on GitHub (pinned to 1e1c6a8fc0)
Solutions
- Flatten shift and axis to 1-D (or scalars) before calling roll
- Loop or vmap over rows for per-row shifts
- Ensure len(shift) == len(axis) as flat sequences for paired rolling
Example fix
// before jnp.roll(a, shifts_matrix, axis=1) # shifts_matrix is 2-D // after jax.vmap(lambda row, s: jnp.roll(row, s))(a, shifts_matrix.ravel())
Defensive patterns
Strategy: validation
Validate before calling
shift = jnp.ravel(jnp.asarray(shift)) axis = jnp.ravel(jnp.asarray(axis)) assert shift.ndim <= 1 and axis.ndim <= 1
Type guard
def valid_roll_args(shift, axis):
return jnp.asarray(shift).ndim <= 1 and jnp.asarray(axis).ndim <= 1 Prevention
- Pass scalars or flat 1-D shift/axis
- Ravel programmatically-built shift arrays
- Use vmap for per-row circular shifts
When it happens
Trigger: jnp.roll(a, shift_2d, axis=0); passing axis as a 2-D array; shift shape (n,1) with axis shape (m,) broadcasting to (n,m).
Common situations: Programmatic roll code building shift/axis from meshgrid or outer products; batching rolls where a 2-D shift matrix seemed natural.
Related errors
- axis is out of range.
- stride_axis is out of range
- scan got `length` argument of {} which disagrees with leadin
- conv_general_dilated batch_group_count must divide lhs batch
- conv_general_dilated rhs output feature dimension size must
AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27).
Data as JSON: /api/errors/3a67ae0d613bb027.
Report an issue: GitHub.