jax-ml/jax · error · NotImplementedError

pallas_call does not support hijax for index_map

Error message

pallas_call does not support hijax for index_map

What it means

pallas_call lowers high-level JAXprs to 'lo' (lower) JAXprs. Block mapping index_map functions must already be lowered; if an index_map jaxpr is still 'high' (hijax — containing high-level primitives that were not lowered), lowering cannot proceed and NotImplementedError is raised.

Source

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

    grid_mapping: GridMapping,
    mesh: pallas_core.Mesh | None,
    debug: bool,
    interpret: Any,
    compiler_params: Any,
    cost_estimate: CostEstimate | None,
    out_avals: tuple[jax_core.AbstractValue, ...],
    metadata: FrozenDict[str, str] | None,
    name: str | None,
):
  closed_jaxpr = jaxpr
  with grid_mapping.trace_env():
    closed_lo_jaxpr = pe.lower_jaxpr2(closed_jaxpr)
  assert not closed_lo_jaxpr.consts
  lo_jaxpr = closed_lo_jaxpr
  for block_mapping in grid_mapping.block_mappings:
    index_map_jaxpr = block_mapping.index_map_jaxpr
    if index_map_jaxpr.is_high:
      raise NotImplementedError(
          "pallas_call does not support hijax for index_map"
      )
  avals = [jax_core.typeof(a) for a in hi_args]
  lo_args = [lo_val for aval, x in zip(avals, hi_args)
             for lo_val in aval.lower_val(x)]
  lo_out_avals = [
      lo_aval
      for aval in out_avals
      for lo_aval in (aval.lo_ty() if aval.is_high else [aval])
  ]
  lo_grid_mapping = grid_mapping.to_lojax()
  in_avals = [v.aval for v in lo_jaxpr.invars]
  scalar_prefetch_avals = in_avals[lo_grid_mapping.slice_index_ops]
  operand_avals = in_avals[lo_grid_mapping.slice_block_ops]
  scratch_avals = in_avals[lo_grid_mapping.slice_scratch_ops]
  # Some basic checks
  assert len(scalar_prefetch_avals) + len(operand_avals) + len(
      scratch_avals

View on GitHub (pinned to 1e1c6a8fc0)

Solutions

  1. Rewrite the index_map to use only basic/lowerable operations (index arithmetic on program IDs) so it lowers cleanly
  2. Rebuild the GridMapping through public APIs (pallas_call/grid helpers) that lower index maps automatically instead of hand-constructing BlockMapping
  3. Update JAX — hi/lo ('hijax') support in index maps is actively evolving; a newer version may lower your case
Defensive patterns

Strategy: fallback

Try / catch

try:
    compiled = jax.jit(f).lower(x)
except NotImplementedError as e:
    if 'hijax' in str(e):
        raise RuntimeError('index_map must use only lowerable ops') from e
    raise

Prevention

When it happens

Trigger: Calling pallas_call (or an API that internally lowers it, e.g. _pallas_call_to_lojax during compilation/export) where a BlockMapping's index_map was built with high-level JAX operations that were never lowered to the lo representation.

Common situations: Using higher-level or transform-based operations (jit, vmap, autodiff) inside the index_map of a GridMapping; constructing GridMapping/BlockMapping manually without lowering index_map; version changes in the hi/lo JAXpr pipeline.

Related errors


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