jax-ml/jax · error · ValueError

input pinned buffers without input_output_aliases:{missing}

Error message

input pinned buffers without input_output_aliases:{missing}

What it means

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.

Source

Thrown at jax/_src/pallas/pallas_call.py:110

    compiler_params: CompilerParams | None,
    input_output_aliases,
    grid_mapping,
    **params,
):
  del params  # Unused.

  effs: Set[jax_core.Effect] = {*pallas_core.get_interpret_effects(interpret)}

  # closed-over refs and dynamic grid bounds aren't reflected in
  # input_output_aliases, though they are present in `avals`, so split them off
  num_refs = sum(isinstance(a, state.AbstractRef) for a in avals)
  _, _, avals = split_list(avals, [num_refs, grid_mapping.num_dynamic_grid_bounds])

  inout_aliases = dict(input_output_aliases)
  lin_avals = {i for i, a in enumerate(avals)
               if isinstance(a, state_types.AbstractLinVal)}
  if (missing := lin_avals - set(inout_aliases)):
    raise ValueError(f"input pinned buffers without input_output_aliases:"
                     f"{missing}")
  outin_aliases = {out_idx: in_idx for in_idx, out_idx in inout_aliases.items()}
  out_avals = tuple(
      avals[outin_aliases[out_idx]] if out_idx in outin_aliases else a
      for out_idx, a in enumerate(out_avals)
  )
  # Make sure we don't return ShapedArray with pallas memory space to the
  # outside world.
  out_avals = tuple(a.update(memory_space=jax_core.MemorySpace.Device)
                    if isinstance(a, jax_core.ShapedArray) else a
                    for a in out_avals)

  # TODO(mattjj,yashkatariya): if we hide vmapped away mesh axes, use this:
  # if not (all(a.sharding.mesh.are_all_axes_manual for a in avals) and
  #         all(a.sharding.mesh.are_all_axes_manual for a in out_avals) and
  #         get_abstract_mesh().are_all_axes_manual):
  #   raise ValueError("pallas_call requires all mesh axes to be Manual, "
  #                    f"got {get_abstract_mesh().axis_types}")

View on GitHub (pinned to 1e1c6a8fc0)

Solutions

  1. Add an entry (input_index, output_index) to input_output_aliases for each pinned buffer input listed in the error message
  2. Re-check input ordering after adding/removing kernel arguments and fix stale alias indices
  3. If the input should not be mutated, pass a regular (non-pinned) array instead of a pinned buffer

Example fix

# before
pallas_call(kernel, grid=grid, out_shape=out_shape)(pinned_buf, x)
# after
pallas_call(kernel, grid=grid, out_shape=out_shape,
            input_output_aliases=[(0, 0)])(pinned_buf, x)
Defensive patterns

Strategy: validation

Validate before calling

lin_idx = {i for i, a in enumerate(flat_avals)
           if type(a).__name__ == 'AbstractLinVal'}
missing = lin_idx - {i for i, _ in input_output_aliases or []}
assert not missing, f'pinned inputs missing aliases: {missing}'

Prevention

When it happens

Trigger: 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.

Common situations: 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.

Related errors


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