{"record":{"id":"3a67ae0d613bb027","repo":"jax-ml/jax","slug":"shift-and-axis-arguments-to-roll-must-be-scala","errorCode":null,"errorMessage":"'shift' and 'axis' arguments to roll must be scalars or 1D arrays","messagePattern":"'shift' and 'axis' arguments to roll must be scalars or 1D arrays","errorType":"exception","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/_src/numpy/lax_numpy.py","lineNumber":8487,"sourceCode":"  return _nanargmin(a, None if axis is None else operator.index(axis), keepdims=bool(keepdims))\n\n\n@api.jit(static_argnames=('axis', 'keepdims'))\ndef _nanargmin(a: Array, axis: int | None = None, keepdims : bool = False):\n  if not issubdtype(a.dtype, np.inexact):\n    return argmin(a, axis=axis, keepdims=keepdims)\n  nan_mask = ufuncs.isnan(a)\n  a = where(nan_mask, np.inf, a)\n  res = argmin(a, axis=axis, keepdims=keepdims)\n  return where(reductions.all(nan_mask, axis=axis, keepdims=keepdims), -1, res)\n\n\n@api.jit(static_argnums=(2,))\ndef _roll_dynamic(a: Array, shift: Array, axis: Sequence[int]) -> Array:\n  b_shape = lax.broadcast_shapes(shift.shape, np.shape(axis))\n  if len(b_shape) != 1:\n    msg = \"'shift' and 'axis' arguments to roll must be scalars or 1D arrays\"\n    raise ValueError(msg)\n\n  for x, i in zip(broadcast_to(shift, b_shape),\n                  np.broadcast_to(axis, b_shape)):  # pyrefly: ignore[no-matching-overload]\n    a_shape_i = array(a.shape[i], dtype=np.int32)\n    x = ufuncs.remainder(lax.convert_element_type(x, np.int32),\n                         lax.max(a_shape_i, np.int32(1)))\n    a_concat = lax.concatenate((a, a), i)\n    a = lax_slicing.dynamic_slice_in_dim(a_concat, a_shape_i - x, a.shape[i], axis=i)\n  return a\n\n@api.jit(static_argnums=(1, 2))\ndef _roll_static(a: Array, shift: Sequence[int], axis: Sequence[int]) -> Array:\n  for ax, s in zip(*np.broadcast_arrays(axis, shift)):\n    if a.shape[ax] == 0:\n      continue\n    i = (-s) % a.shape[ax]\n    a = lax.concatenate([lax_slicing.slice_in_dim(a, i, a.shape[ax], axis=ax),\n                         lax_slicing.slice_in_dim(a, 0, i, axis=ax)],","sourceCodeStart":8469,"sourceCodeEnd":8505,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/numpy/lax_numpy.py#L8469-L8505","documentation":"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).","triggerScenarios":"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).","commonSituations":"Programmatic roll code building shift/axis from meshgrid or outer products; batching rolls where a 2-D shift matrix seemed natural.","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"],"exampleFix":"// before\njnp.roll(a, shifts_matrix, axis=1)  # shifts_matrix is 2-D\n// after\njax.vmap(lambda row, s: jnp.roll(row, s))(a, shifts_matrix.ravel())\n","handlingStrategy":"validation","validationCode":"shift = jnp.ravel(jnp.asarray(shift))\naxis = jnp.ravel(jnp.asarray(axis))\nassert shift.ndim <= 1 and axis.ndim <= 1","typeGuard":"def valid_roll_args(shift, axis):\n    return jnp.asarray(shift).ndim <= 1 and jnp.asarray(axis).ndim <= 1","tryCatchPattern":null,"preventionTips":["Pass scalars or flat 1-D shift/axis","Ravel programmatically-built shift arrays","Use vmap for per-row circular shifts"],"tags":["jax","roll","shape-validation"],"backgroundTag":"array-dimension-validation","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}