jax-ml/jax · error · ValueError
Aliasing of scalar prefetch arguments is not currently suppo
Error message
Aliasing of scalar prefetch arguments is not currently supported in TPU interpret mode.
What it means
In TPU interpret mode, pallas_call input_output_aliases that alias a scalar prefetch argument (an input index below the number of scalar inputs) are rejected, because the interpreter's buffer allocation scheme cannot map scalar prefetch inputs to output buffer ids.
Source
Thrown at jax/_src/pallas/mosaic/interpret/interpret_pallas_call.py:2055
input_buffer_ids = []
for i, var in enumerate(
jaxpr.invars[grid_mapping.num_index_operands:][:grid_mapping.num_inputs]):
assert var.aval.dtype == input_args[i].dtype # pyrefly: ignore[missing-attribute]
token, buffer_id = callback.io_callback(
_allocate_buffer,
(TOKEN_SHAPE_DTYPE, jax.ShapeDtypeStruct((), jnp.int16)),
token,
device_id,
None, # local_core_id
TPU_MEMORY_SPACE_IDXS[mosaic_core.MemorySpace.HBM],
input_args[i],
)
input_buffer_ids.append(buffer_id)
# Allocate buffers in HBM for pallas_call outputs.
oi_alias_map = {v: k - len(scalars) for k, v in input_output_aliases}
if any(i < 0 for i in oi_alias_map.keys()):
raise ValueError('Aliasing of scalar prefetch arguments is not currently '
'supported in TPU interpret mode.')
output_buffer_ids = []
output_buffer_shapes = []
output_vals = []
num_outputs = grid_mapping.num_outputs
output_block_shapes = block_shapes[num_inputs : num_inputs + num_outputs]
for i, bm in enumerate(grid_mapping.block_mappings_output):
if i in oi_alias_map:
# Reuse the HBM buffer for the aliased pallas_call input.
output_buffer_ids.append(input_buffer_ids[oi_alias_map[i]])
output_buffer_shapes.append(input_args[oi_alias_map[i]].shape)
output_vals.append(input_args[oi_alias_map[i]])
else:
out_val = interpret_utils.get_uninitialized_array(
bm.array_aval.shape, bm.array_aval.dtype,
interpret_params.uninitialized_memory)
padded_val = interpret_utils.pad_to_block_dimension(
out_val, output_block_shapes[i], interpret_params.uninitialized_memoryView on GitHub (pinned to 1e1c6a8fc0)
Solutions
- Alias outputs only to non-scalar (array) inputs
- Return updated scalars as separate kernel outputs instead of aliasing the scalar input
- Skip interpret mode for kernels requiring scalar aliasing
Example fix
# before pallas_call(kernel, out_shapes, inputs, input_output_aliases=(0, 0)) # input 0 is scalar # after pallas_call(kernel, out_shapes + scalar_shape, inputs) # return updated scalar as extra output
Defensive patterns
Strategy: validation
Validate before calling
num_scalars = sum(1 for x in inputs if getattr(x, 'shape', None) in ((), None))
for i, _ in input_output_aliases:
assert i >= num_scalars, 'cannot alias scalar prefetch inputs in TPU interpret mode' Try / catch
try:
interpret_run(kernel)
except ValueError as e:
if 'scalar prefetch' in str(e):
# move aliased scalar to an explicit output and rerun
raise Prevention
- Alias outputs only to array inputs
- Return updated scalars as extra outputs
- Check alias indices against scalar input count before interpret runs
When it happens
Trigger: Using input_output_aliases=(i, j) where i < len(scalars) — i.e., aliasing an output to a scalar (prefetch) input — while running with interpret mode on TPU.
Common situations: Kernels designed for GPU/compiled TPU that return an updated scalar argument (e.g., running token/seed) via output aliasing; porting such kernels to interpret-mode debugging.
Related errors
- Out-of-bounds read of ({device_id} {local_core_id} {memory_s
- masked load_p
- run_scoped_p with collective axes is not supported
- Non-decrementing wait is not supported.
- Kernel input {j} in HBM but does not have trivial BlockSpec.
AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27).
Data as JSON: /api/errors/823c5a9584b28651.
Report an issue: GitHub.