{"record":{"id":"22ac2d6e42040e92","repo":"jax-ml/jax","slug":"pallas-call-does-not-support-hijax-for-index-map","errorCode":null,"errorMessage":"pallas_call does not support hijax for index_map","messagePattern":"pallas_call does not support hijax for index_map","errorType":"exception","errorClass":"NotImplementedError","httpStatus":null,"severity":"error","filePath":"jax/_src/pallas/pallas_call.py","lineNumber":182,"sourceCode":"    grid_mapping: GridMapping,\n    mesh: pallas_core.Mesh | None,\n    debug: bool,\n    interpret: Any,\n    compiler_params: Any,\n    cost_estimate: CostEstimate | None,\n    out_avals: tuple[jax_core.AbstractValue, ...],\n    metadata: FrozenDict[str, str] | None,\n    name: str | None,\n):\n  closed_jaxpr = jaxpr\n  with grid_mapping.trace_env():\n    closed_lo_jaxpr = pe.lower_jaxpr2(closed_jaxpr)\n  assert not closed_lo_jaxpr.consts\n  lo_jaxpr = closed_lo_jaxpr\n  for block_mapping in grid_mapping.block_mappings:\n    index_map_jaxpr = block_mapping.index_map_jaxpr\n    if index_map_jaxpr.is_high:\n      raise NotImplementedError(\n          \"pallas_call does not support hijax for index_map\"\n      )\n  avals = [jax_core.typeof(a) for a in hi_args]\n  lo_args = [lo_val for aval, x in zip(avals, hi_args)\n             for lo_val in aval.lower_val(x)]\n  lo_out_avals = [\n      lo_aval\n      for aval in out_avals\n      for lo_aval in (aval.lo_ty() if aval.is_high else [aval])\n  ]\n  lo_grid_mapping = grid_mapping.to_lojax()\n  in_avals = [v.aval for v in lo_jaxpr.invars]\n  scalar_prefetch_avals = in_avals[lo_grid_mapping.slice_index_ops]\n  operand_avals = in_avals[lo_grid_mapping.slice_block_ops]\n  scratch_avals = in_avals[lo_grid_mapping.slice_scratch_ops]\n  # Some basic checks\n  assert len(scalar_prefetch_avals) + len(operand_avals) + len(\n      scratch_avals","sourceCodeStart":164,"sourceCodeEnd":200,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/pallas/pallas_call.py#L164-L200","documentation":"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.","triggerScenarios":"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.","commonSituations":"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.","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"],"exampleFix":null,"handlingStrategy":"fallback","validationCode":null,"typeGuard":null,"tryCatchPattern":"try:\n    compiled = jax.jit(f).lower(x)\nexcept NotImplementedError as e:\n    if 'hijax' in str(e):\n        raise RuntimeError('index_map must use only lowerable ops') from e\n    raise","preventionTips":["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"],"tags":["jax","pallas","pallas-call","index-map","lowering","notimplementederror"],"backgroundTag":"unsupported-operation-in-lowering","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}