{"record":{"id":"5095806d16f0e851","repo":"jax-ml/jax","slug":"jax-numpy-put-along-axis-cannot-modify-arrays-in-p","errorCode":null,"errorMessage":"jax.numpy.put_along_axis cannot modify arrays in-place, because JAX arraysare immutable. Pass inplace=False to instead return an updated array.","messagePattern":"jax\\.numpy\\.put_along_axis cannot modify arrays in-place, because JAX arraysare immutable\\. Pass inplace=False to instead return an updated array\\.","errorType":"exception","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/_src/numpy/indexing.py","lineNumber":1035,"sourceCode":"    - :func:`jax.numpy.place`: place elements into an array via boolean mask.\n    - :func:`jax.numpy.ndarray.at`: array updates using NumPy-style indexing.\n    - :func:`jax.numpy.take`: extract values from an array at given indices.\n    - :func:`jax.numpy.take_along_axis`: extract values from an array along an axis.\n\n  Examples:\n    >>> from jax import numpy as jnp\n    >>> a = jnp.array([[10, 30, 20], [60, 40, 50]])\n    >>> i = jnp.argmax(a, axis=1, keepdims=True)\n    >>> print(i)\n    [[1]\n     [0]]\n    >>> b = jnp.put_along_axis(a, i, 99, axis=1, inplace=False)\n    >>> print(b)\n    [[10 99 20]\n     [99 40 50]]\n  \"\"\"\n  if inplace:\n    raise ValueError(\n      \"jax.numpy.put_along_axis cannot modify arrays in-place, because JAX arrays\"\n      \"are immutable. Pass inplace=False to instead return an updated array.\")\n\n  arr, indices, values = util.ensure_arraylike(\"put_along_axis\", arr, indices, values)\n\n  original_axis = axis\n  original_arr_shape = arr.shape\n\n  if axis is None:\n    arr = arr.ravel()\n    axis = 0\n\n  if not arr.ndim == indices.ndim:\n    raise ValueError(\n      \"put_along_axis arguments 'arr' and 'indices' must have same ndim. Got \"\n      f\"{arr.ndim=} and {indices.ndim=}.\"\n    )\n","sourceCodeStart":1017,"sourceCodeEnd":1053,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/numpy/indexing.py#L1017-L1053","documentation":"JAX arrays are immutable, so put_along_axis cannot write in-place despite exposing an inplace flag for NumPy API compatibility. With inplace=True it raises ValueError telling you to use inplace=False.","triggerScenarios":"Calling jnp.put_along_axis(arr, indices, values, axis, inplace=True) (the NumPy default).","commonSituations":"Porting np.put_along_axis calls verbatim; users assuming the NumPy signature works identically in JAX.","solutions":["Pass inplace=False and use the returned array: arr = jnp.put_along_axis(arr, idx, vals, axis, inplace=False)","Remember all JAX updates are functional — always rebind the result"],"exampleFix":"// before\njnp.put_along_axis(a, i, 99, axis=1, inplace=True)\n// after\na = jnp.put_along_axis(a, i, 99, axis=1, inplace=False)","handlingStrategy":"validation","validationCode":"assert not inplace, 'JAX is immutable; use inplace=False'","typeGuard":null,"tryCatchPattern":null,"preventionTips":["Always pass inplace=False and rebind: a = jnp.put_along_axis(a, ...)","Remember JAX transformations require functional updates"],"tags":["jax","put-along-axis","immutability","inplace"],"backgroundTag":"immutable-array-inplace","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}