{"record":{"id":"b641309ac432b30e","repo":"jax-ml/jax","slug":"only-single-indexer-is-supported","errorCode":null,"errorMessage":"Only single indexer is supported.","messagePattern":"Only single indexer is supported\\.","errorType":"exception","errorClass":"NotImplementedError","httpStatus":null,"severity":"error","filePath":"jax/_src/pallas/triton/primitives.py","lineNumber":461,"sourceCode":"\n@lowering.register_lowering(atomic_rmw_p)\ndef _atomic_lowering_rule(\n    ctx: lowering.LoweringRuleContext,\n    *args_flat,\n    args_tree,\n    atomic_type: AtomicOpType,\n):\n  block_info, *_ = ctx.block_infos\n  assert block_info is not None\n  ptr, indexers, val, mask = args_tree.unflatten(args_flat)\n  *_, value_aval, mask_aval = args_tree.unflatten(ctx.avals_in)\n  indexers = list(indexers)\n  if not indexers or not isinstance(indexers[-1], indexing.NDIndexer):\n    ref_aval = state.transform_type(indexers, ctx.avals_in[0])\n    assert isinstance(ref_aval, state.AbstractRef)\n    indexers.append(indexing.NDIndexer.make_trivial_indexer(ref_aval.shape))\n  if len(indexers) != 1:\n    raise NotImplementedError(\"Only single indexer is supported.\")\n  idx = indexers[0]\n  ptr = lowering._compute_pointers_from_indices(ptr, block_info, idx)\n  val = lowering._ensure_ir_value(val, value_aval)\n  if mask is not None:\n    mask = lowering._ensure_ir_value(mask, mask_aval)\n  if atomic_type == AtomicOpType.XCHG:\n    op = tt_dialect.RMWOp.XCHG\n  elif atomic_type == AtomicOpType.ADD:\n    if isinstance(val.type, ir.IntegerType):\n      op = tt_dialect.RMWOp.ADD\n    else:\n      op = tt_dialect.RMWOp.FADD\n  elif atomic_type == AtomicOpType.MIN:\n    if isinstance(val.type, ir.IntegerType):\n      op = (\n        tt_dialect.RMWOp.MIN\n        if jnp.issubdtype(value_aval.dtype, jnp.signedinteger)\n        else tt_dialect.RMWOp.UMIN","sourceCodeStart":443,"sourceCodeEnd":479,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/pallas/triton/primitives.py#L443-L479","documentation":"Raised by the Pallas Triton lowering rule for atomic operations when the reference being atomically updated is indexed by more than one indexer (i.e., a chained/multi-dimensional indexing that produces multiple NDIndexer elements). The Triton backend can only compute a single flat pointer per atomic op, so compound indexing cannot be lowered.","triggerScenarios":"Calling pallas.triton primitives like atomic_add/atomic_or/atomic_xchg with an idx argument that is a tuple of multiple indexers, or applying an atomic op to a Ref produced via chained indexing (e.g. ref[i][j] or a view chain) inside a pallas triton kernel.","commonSituations":"Porting a GPU kernel that does atomic updates on a 2D location using two separate bracket operations; using jnp-style fancy indexing inside pallas kernels which silently creates multiple indexers.","solutions":["Flatten the array (or use a single NDIndexer) so the atomic access uses one indexer, e.g. ref[i * cols + j] instead of ref[i, j] via chained views","If indexing a block, index into a temporary non-Ref array first, then do a scalar atomic store","File a feature request / check newer JAX versions if multi-indexer atomics were added"],"exampleFix":"# before\np.atomic_add(ref[i], j, val)  # chained -> multiple indexers\n# after\np.atomic_add(ref, i * ref.shape[1] + j, val)  # single flat indexer","handlingStrategy":"validation","validationCode":"idxers = idx if isinstance(idx, tuple) else (idx,)\nassert len(idxers) == 1, 'use a single (flat) indexer for atomics'","typeGuard":null,"tryCatchPattern":null,"preventionTips":["Always index Refs with a single indexer expression for atomic ops","Avoid chained bracket indexing on Refs inside pallas kernels"],"tags":["jax","pallas","triton","atomics","indexing"],"backgroundTag":"unsupported-operation-lowering","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}