{"record":{"id":"72838ed561e36095","repo":"jax-ml/jax","slug":"collective-axes-must-be-specified-when-leader-t","errorCode":null,"errorMessage":"`collective_axes` must be specified when `leader_tracked` is set","messagePattern":"`collective_axes` must be specified when `leader_tracked` is set","errorType":"validation","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/_src/pallas/mosaic_gpu/primitives.py","lineNumber":1303,"sourceCode":"  flat_dst_transforms, dst_transforms_treedef = tree_util.tree_flatten(\n      dst_transforms\n  )\n  has_barrier = barrier is not None\n  if has_barrier:\n    barrier, barrier_transforms = state_primitives.get_ref_and_transforms(\n        barrier, None, \"copy_gmem_to_smem\"\n    )\n    barrier_operands = [barrier]\n  else:\n    barrier_transforms = []\n    barrier_operands = []\n  flat_barrier_transforms, barrier_transforms_treedef = tree_util.tree_flatten(\n      barrier_transforms\n  )\n  if isinstance(collective_axes, str):\n    collective_axes = (collective_axes,)\n  if leader_tracked is not None and collective_axes is None:\n    raise ValueError(\n        \"`collective_axes` must be specified when `leader_tracked` is set\"\n    )\n  copy_gmem_to_smem_p.bind(\n      src,\n      dst,\n      *barrier_operands,\n      *flat_src_transforms,\n      *flat_dst_transforms,\n      *flat_barrier_transforms,\n      *[] if predicate is None else [predicate],\n      src_transforms_treedef=src_transforms_treedef,\n      dst_transforms_treedef=dst_transforms_treedef,\n      barrier_transforms_treedef=barrier_transforms_treedef,\n      collective_axes=collective_axes,\n      leader_tracked=leader_tracked,\n      oob_mode=oob_mode,\n      has_barrier=has_barrier,\n      has_user_predicate=predicate is not None,","sourceCodeStart":1285,"sourceCodeEnd":1321,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/pallas/mosaic_gpu/primitives.py#L1285-L1321","documentation":"copy_gmem_to_smem's leader_tracked parameter marks a copy partitioned across a collective (e.g. a TP or multi-GPU collective axis). To track which leader performs the copy, the implementation must know the collective axes, so leader_tracked without collective_axes is a user API error.","triggerScenarios":"Calling copy_gmem_to_smem(..., leader_tracked=CopyPartition.PARTITIONED(axis)) (or similar) without also passing collective_axes.","commonSituations":"Writing multi-GPU/collective Pallas kernels and copy-pasting a leader_tracked example while forgetting the collective_axes argument added in the same API revision.","solutions":["Pass collective_axes (tuple of axis names or a single string) whenever leader_tracked is set.","If you did not intend collective behavior, drop leader_tracked."],"exampleFix":"# before\ncopy_gmem_to_smem(src, dst, leader_tracked=CopyPartition.PARTITIONED(0))\n# after\ncopy_gmem_to_smem(src, dst,\n    leader_tracked=CopyPartition.PARTITIONED(0), collective_axes=('tp',))","handlingStrategy":"validation","validationCode":"if leader_tracked is not None:\n    assert collective_axes is not None, 'collective_axes required with leader_tracked'","typeGuard":null,"tryCatchPattern":null,"preventionTips":["Pass collective_axes and leader_tracked together by convention; wrap them in a small helper function."],"tags":["mosaic-gpu","pallas","api-misuse","collectives"],"backgroundTag":"missing-required-argument","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}