{"record":{"id":"9008df1105d034fb","repo":"jax-ml/jax","slug":"axis-mixes-jax-mesh-and-pallas-mesh-grid-axes","errorCode":null,"errorMessage":"{axis} mixes JAX mesh and Pallas mesh grid axes","messagePattern":"(.+?) mixes JAX mesh and Pallas mesh grid axes","errorType":"exception","errorClass":"NotImplementedError","httpStatus":null,"severity":"error","filePath":"jax/_src/pallas/mosaic/interpret/utils.py","lineNumber":160,"sourceCode":"  def __init__(self, initial_value: int):\n    self.value = initial_value\n    self.lock = threading.Lock()\n\n  def get_next(self):\n    with self.lock:\n      result = self.value\n      self.value += 1\n    return result\n\n\n# TODO(sharadmv): De-dup this w/ the impl in primitives.py.\ndef _device_id_dict_to_mesh(device_id_dict, axis_sizes, axis_indices):\n  physical_axis_dict = {}\n  axis_names = axis_sizes.keys()\n  for axis, idx in device_id_dict.items():\n    if isinstance(axis, tuple) and any(a in axis_names for a in axis):\n      if not all(a in axis_names for a in axis):\n        raise NotImplementedError(\n            f\"{axis} mixes JAX mesh and Pallas mesh grid axes\"\n        )\n      axes_dimensions = [axis_sizes[name] for name in axis]\n      for axis_index, axis_name in enumerate(axis):\n        axis_size = axis_sizes[axis_name]\n        inner_mesh_size = math.prod(axes_dimensions[axis_index + 1 :])\n        minor_divisor = inner_mesh_size\n\n        # Fast path for power of 2s\n        if inner_mesh_size & (inner_mesh_size - 1) == 0:\n          shift_len = (inner_mesh_size & -inner_mesh_size).bit_length() - 1\n          partial_device_idx = idx >> shift_len\n        else:\n          partial_device_idx = idx // minor_divisor\n\n        if axis_size & (axis_size - 1) == 0:\n          device_idx = partial_device_idx & (axis_size - 1)\n        else:","sourceCodeStart":142,"sourceCodeEnd":178,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/pallas/mosaic/interpret/utils.py#L142-L178","documentation":"A device-coordinate dict contained a tuple axis mixing JAX mesh axis names with Pallas mesh (grid) axis names. Mixed groupings are not supported when mapping device coordinates to a logical program ID.","triggerScenarios":"Passing device_id as a dict whose keys are tuples where some names come from the JAX mesh (named_sharding axes) and others from the Pallas grid/mesh, when DeviceIdType.MESH is used in interpret mode.","commonSituations":"Using Mesh/JIT sharding with names that partially overlap Pallas grid axis names; constructing manual device coordinate dicts for multi-mesh interpret runs.","solutions":["Split the tuple axis so each tuple contains only JAX mesh axes or only Pallas grid axes","Rename axes so the two mesh namespaces don't overlap in one tuple key","Pass device_id as a plain coordinate tuple instead of a dict when possible"],"exampleFix":null,"handlingStrategy":"validation","validationCode":"mesh_axes = set(axis_sizes)\nfor axis in device_id_dict:\n    if isinstance(axis, tuple):\n        assert all(a in mesh_axes for a in axis) or not any(a in mesh_axes for a in axis), 'mixed mesh/grid axis tuple'","typeGuard":null,"tryCatchPattern":null,"preventionTips":["Keep JAX mesh axis names and Pallas grid axis names in separate namespaces","Pass plain coordinate tuples rather than dict device IDs when possible"],"tags":["jax","pallas","mesh","device-id","interpret-mode","not-implemented"],"backgroundTag":"unsupported-axis-mapping","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}