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

  1. Flatten shift and axis to 1-D (or scalars) before calling roll
  2. Loop or vmap over rows for per-row shifts
  3. 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

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


AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27). Data as JSON: /api/errors/3a67ae0d613bb027. Report an issue: GitHub.