{"record":{"id":"64cd036e70bb7b60","repo":"jax-ml/jax","slug":"accumulate-does-not-allow-multiple-axes","errorCode":null,"errorMessage":"accumulate does not allow multiple axes","messagePattern":"accumulate does not allow multiple axes","errorType":"exception","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/_src/numpy/ufunc_api.py","lineNumber":390,"sourceCode":"      raise ValueError(\"accumulate only supported for binary ufuncs\")\n    if self.nout != 1:\n      raise ValueError(\"accumulate only supported for functions returning a single value\")\n    if out is not None:\n      raise NotImplementedError(f\"out argument of {self.__name__}.accumulate()\")\n    accumulate = self.__static_props['accumulate'] or self._accumulate_via_scan\n    return accumulate(a, axis=axis, dtype=dtype)\n\n  def _accumulate_via_scan(self, arr: ArrayLike, axis: int = 0,\n                           dtype: DTypeLike | None = None) -> Array:\n    assert self.nin == 2 and self.nout == 1\n    check_arraylike(f\"{self.__name__}.accumulate\", arr)\n    arr = lax.asarray(arr)\n\n    if dtype is None:\n      dtype = api.eval_shape(self._func, lax._one(arr), lax._one(arr)).dtype\n\n    if axis is None or isinstance(axis, tuple):\n      raise ValueError(\"accumulate does not allow multiple axes\")\n    axis = canonicalize_axis(axis, np.ndim(arr))\n\n    if arr.size == 0:\n      return lax.full(arr.shape, 0, dtype)\n    arr = _moveaxis(arr, axis, 0)\n    def scan_fun(carry, _):\n      i, x = carry\n      y = _where(i == 0, arr[0].astype(dtype), self(x.astype(dtype), arr[i].astype(dtype)))\n      return (i + 1, y), y\n    _, result = control_flow.scan(scan_fun, (0, arr[0].astype(dtype)), None, length=arr.shape[0])\n    return _moveaxis(result, 0, axis)\n\n  @api.jit(static_argnums=[0], static_argnames=['inplace'])\n  def at(self, a: ArrayLike, indices: Any, b: ArrayLike | None = None, /, *,\n         inplace: bool = True) -> Array:\n    \"\"\"Update elements of an array via the specified unary or binary ufunc.\n\n    JAX implementation of :func:`numpy.ufunc.at`.","sourceCodeStart":372,"sourceCodeEnd":408,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/numpy/ufunc_api.py#L372-L408","documentation":"ufunc.accumulate operates along exactly one axis; numpy's semantics of axis=None (flatten) or a tuple of axes are rejected by JAX's scan-based accumulate with this ValueError.","triggerScenarios":"jnp.add.accumulate(x, axis=None) or jnp.add.accumulate(x, axis=(0,1)).","commonSituations":"Porting numpy code using axis=None to accumulate over a flattened array; generic code that forwards tuple axes.","solutions":["Pass a single integer axis","For axis=None behavior, ravel first: jnp.add.accumulate(x.ravel()) then reshape back"],"exampleFix":"// before\njnp.add.accumulate(x, axis=None)\n// after\njnp.add.accumulate(x.ravel()).reshape(x.shape)","handlingStrategy":"validation","validationCode":"assert isinstance(axis, int), 'accumulate requires a single int axis'","typeGuard":null,"tryCatchPattern":null,"preventionTips":["Handle axis=None by ravelling before .accumulate"],"tags":["jax","ufunc","accumulate","axis"],"backgroundTag":"accumulate-single-axis-only","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}