{"record":{"id":"cb8b46c42c8ddb52","repo":"jax-ml/jax","slug":"axis-name-mixes-jax-mesh-and-pallas-mesh-grid-ax","errorCode":null,"errorMessage":"{axis_name} 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/primitives.py","lineNumber":1215,"sourceCode":"    _semaphore_wait_discharge_rule\n)\n\n\ndef _device_id_dict_to_mesh(mesh_context: pallas_utils.MeshInfo | None, device_id_dict, get_axis_index):\n  if mesh_context is None:\n    mesh_axis_sizes = {}\n  else:\n    mesh_axis_sizes = dict(\n        zip(mesh_context.axis_names, mesh_context.mesh_shape)\n    )\n  physical_axis_dict = {}\n  # Handle joint axes (i.e., one logical axis over >1 physical axes)\n  for axis_name, idx in device_id_dict.items():\n    if isinstance(axis_name, tuple) and any(\n        a in mesh_axis_sizes for a in axis_name\n    ):\n      if not all(a in mesh_axis_sizes for a in axis_name):\n        raise NotImplementedError(\n            f\"{axis_name} mixes JAX mesh and Pallas mesh grid axes\"\n        )\n      axes_dimensions = [mesh_axis_sizes[name] for name in axis_name]\n      for axis_index, axis_name in enumerate(axis_name):\n        axis_size = mesh_axis_sizes[axis_name]\n        inner_mesh_size = math.prod(axes_dimensions[axis_index + 1 :])\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 // inner_mesh_size\n\n        if axis_size & (axis_size - 1) == 0:\n          device_idx = partial_device_idx & jnp.asarray(\n              axis_size - 1, dtype=partial_device_idx.dtype\n          )","sourceCodeStart":1197,"sourceCodeEnd":1233,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/pallas/primitives.py#L1197-L1233","documentation":"In device_id_to_logical, a joint axis name (tuple of axes) must be composed entirely of JAX mesh axes or entirely of Pallas mesh grid axes. Mixing both kinds inside one joint axis is not implemented.","triggerScenarios":"Declaring a collective whose axis name is a tuple containing both a jax.lax mesh axis and a Pallas grid axis name in the same tuple.","commonSituations":"Writing multi-axis collectives (e.g., ('data','rep') where 'data' is a JAX Mesh axis and 'rep' a Pallas grid axis) in distributed Pallas kernels.","solutions":["Split the joint axis so each tuple contains only JAX mesh axes or only Pallas grid axes","Use separate collectives per axis kind"],"exampleFix":"// before\naxis_name=('mesh_axis', 'pallas_grid_axis')\n// after\n# perform collectives on ('mesh_axis','mesh_axis2') and 'pallas_grid_axis' separately","handlingStrategy":"validation","validationCode":"for ax in axis_names:\n    if isinstance(ax, tuple):\n        kinds = {is_mesh_axis(a) for a in ax}  # your axis bookkeeping\n        assert len(kinds) == 1, f\"joint axis {ax} mixes JAX mesh and Pallas grid axes\"","typeGuard":null,"tryCatchPattern":null,"preventionTips":["Keep joint axis tuples homogeneous (all mesh or all grid)","Document which axis names belong to JAX Mesh vs Pallas grid"],"tags":["pallas","mesh","collectives","axis-name","not-implemented","jax"],"backgroundTag":"invalid-configuration","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}