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_memory

View on GitHub (pinned to 1e1c6a8fc0)

Solutions

  1. Alias outputs only to non-scalar (array) inputs
  2. Return updated scalars as separate kernel outputs instead of aliasing the scalar input
  3. 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

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


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