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
- 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
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
- 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
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
- pallas_call requires all mesh axes to be Manual, got {get_ab
- pallas_call does not support hijax for index_map
- JVP with aliasing not supported.
- This gmm kernel only supports either (m, k) x (g, k, n) -> (
- Group sizes {group_sizes.shape=} must match first dimension
AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27).
Data as JSON: /api/errors/843676e74052dcc9.
Report an issue: GitHub.