{"record":{"id":"843676e74052dcc9","repo":"jax-ml/jax","slug":"input-pinned-buffers-without-input-output-aliases","errorCode":null,"errorMessage":"input pinned buffers without input_output_aliases:{missing}","messagePattern":"input pinned buffers without input_output_aliases:(.+?)","errorType":"exception","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/_src/pallas/pallas_call.py","lineNumber":110,"sourceCode":"    compiler_params: CompilerParams | None,\n    input_output_aliases,\n    grid_mapping,\n    **params,\n):\n  del params  # Unused.\n\n  effs: Set[jax_core.Effect] = {*pallas_core.get_interpret_effects(interpret)}\n\n  # closed-over refs and dynamic grid bounds aren't reflected in\n  # input_output_aliases, though they are present in `avals`, so split them off\n  num_refs = sum(isinstance(a, state.AbstractRef) for a in avals)\n  _, _, avals = split_list(avals, [num_refs, grid_mapping.num_dynamic_grid_bounds])\n\n  inout_aliases = dict(input_output_aliases)\n  lin_avals = {i for i, a in enumerate(avals)\n               if isinstance(a, state_types.AbstractLinVal)}\n  if (missing := lin_avals - set(inout_aliases)):\n    raise ValueError(f\"input pinned buffers without input_output_aliases:\"\n                     f\"{missing}\")\n  outin_aliases = {out_idx: in_idx for in_idx, out_idx in inout_aliases.items()}\n  out_avals = tuple(\n      avals[outin_aliases[out_idx]] if out_idx in outin_aliases else a\n      for out_idx, a in enumerate(out_avals)\n  )\n  # Make sure we don't return ShapedArray with pallas memory space to the\n  # outside world.\n  out_avals = tuple(a.update(memory_space=jax_core.MemorySpace.Device)\n                    if isinstance(a, jax_core.ShapedArray) else a\n                    for a in out_avals)\n\n  # TODO(mattjj,yashkatariya): if we hide vmapped away mesh axes, use this:\n  # if not (all(a.sharding.mesh.are_all_axes_manual for a in avals) and\n  #         all(a.sharding.mesh.are_all_axes_manual for a in out_avals) and\n  #         get_abstract_mesh().are_all_axes_manual):\n  #   raise ValueError(\"pallas_call requires all mesh axes to be Manual, \"\n  #                    f\"got {get_abstract_mesh().axis_types}\")","sourceCodeStart":92,"sourceCodeEnd":128,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/pallas/pallas_call.py#L92-L128","documentation":"In pallas_call, inputs whose abstract values are 'linear' (pinned) buffers (AbstractLinVal) must be aliased to an output via input_output_aliases, because a pinned input is mutated in place and must be surfaced as an output. If any pinned input index is missing from input_output_aliases, a ValueError listing the missing indices is raised.","triggerScenarios":"Calling pallas_call (or the public pallas_call wrapper) with an input that is a pinned/linear buffer while input_output_aliases is None, empty, or does not map that input index to an output index.","commonSituations":"Upgrading JAX/Pallas versions where buffer donation/pinning semantics were introduced; passing donated buffers without updating the input_output_aliases parameter; miscounting indices when adding a new input argument shifts positions.","solutions":["Add an entry (input_index, output_index) to input_output_aliases for each pinned buffer input listed in the error message","Re-check input ordering after adding/removing kernel arguments and fix stale alias indices","If the input should not be mutated, pass a regular (non-pinned) array instead of a pinned buffer"],"exampleFix":"# before\npallas_call(kernel, grid=grid, out_shape=out_shape)(pinned_buf, x)\n# after\npallas_call(kernel, grid=grid, out_shape=out_shape,\n            input_output_aliases=[(0, 0)])(pinned_buf, x)","handlingStrategy":"validation","validationCode":"lin_idx = {i for i, a in enumerate(flat_avals)\n           if type(a).__name__ == 'AbstractLinVal'}\nmissing = lin_idx - {i for i, _ in input_output_aliases or []}\nassert not missing, f'pinned inputs missing aliases: {missing}'","typeGuard":null,"tryCatchPattern":null,"preventionTips":["Whenever passing pinned/donated buffers to pallas_call, immediately pair them in input_output_aliases","Define aliases in one place next to the kernel signature so argument reordering updates them together"],"tags":["jax","pallas","pallas-call","input-output-alias","pinned-buffer"],"backgroundTag":"missing-input-output-alias","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}