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_avalsView on GitHub (pinned to 1e1c6a8fc0)
Solutions
- Rewrite the index_map to use only basic/lowerable operations (index arithmetic on program IDs) so it lowers cleanly
- Rebuild the GridMapping through public APIs (pallas_call/grid helpers) that lower index maps automatically instead of hand-constructing BlockMapping
- 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
- Keep index_map functions limited to program-id arithmetic; avoid transforms like jit/vmap inside them
- Build GridMappings via public helpers rather than manual BlockMapping construction
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
- Index map function {debug_info.func_src_info} for {origin} m
- index_map returned a value of type {type(idx_aval)} at posit
- index_map returned a value of type {type(idx_aval)} at posit
- Index map function {debug_info.func_src_info} for {origin} m
- Index map function {debug_info.func_src_info} for {origin} m
AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27).
Data as JSON: /api/errors/22ac2d6e42040e92.
Report an issue: GitHub.