{"record":{"id":"070aec008d86b681","repo":"jax-ml/jax","slug":"swap-does-not-support-masked-scalar-stores","errorCode":null,"errorMessage":"Swap does not support masked scalar stores","messagePattern":"Swap does not support masked scalar stores","errorType":"exception","errorClass":"NotImplementedError","httpStatus":null,"severity":"error","filePath":"jax/_src/pallas/mosaic/sc_lowering.py","lineNumber":233,"sourceCode":"        \"Integer indexing of refs that follows a non-trivial slice is not\"\n        \" supported on SC\"\n    )\n  if not all(s == 1 for s in strides):\n    raise NotImplementedError(\n        \"Swap only supports slices with stride 1, got {strides}\"\n    )\n\n  if (out_aval.ndim == 0) != (ref_memory_space is MemorySpace.SMEM):\n    message = \"Swap only supports scalars in SMEM.\"\n    if ref_memory_space is MemorySpace.SMEM:\n      message += \" Trying to swap an array of shape {out_aval.shape}.\"\n    else:\n      message += f\" Trying to swap a scalar in {ref_memory_space!r}.\"\n    raise NotImplementedError(message)\n\n  if out_aval.ndim == 0:\n    if mask is not None:\n      raise NotImplementedError(\"Swap does not support masked scalar stores\")\n    if add:\n      # TODO(slebedev): We can use memref.atomic_rmw here, but the SC compiler\n      # doesn't support it yet.\n      raise NotImplementedError(\"Swap does not support atomic scalar adds\")\n    old_val = memref.load(ref, starts)\n    memref.store(val, ref, starts)\n    return old_val\n\n  if not ctx.lowering_context.needs_layout_passes:\n    _check_aval_is_supported(\"Swap\", out_aval)\n  out_vec_type = ir.VectorType.get(\n      out_aval.shape, _dtype_to_ir_type(out_aval.dtype)\n  )\n  if not ctx.lowering_context.needs_layout_passes:\n    old_val = tpu.vector_load(out_vec_type, ref, starts, strides=[], mask=mask)\n    _ = tpu.vector_store(\n        val, ref, indices=starts, strides=[], mask=mask, add=add\n    )","sourceCodeStart":215,"sourceCodeEnd":251,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/pallas/mosaic/sc_lowering.py#L215-L251","documentation":"Atomic scalar add via swap semantics is not implemented in the SC lowering (the TODO notes memref.atomic_rmw is unsupported by the SC compiler). swap_add on scalars therefore raises.","triggerScenarios":"Using swap-add (e.g. ref[idx] += v via swap with add=True) where the value is a scalar in a SparseCore kernel.","commonSituations":"Counter increments or scalar accumulations written with atomic add semantics on SC.","solutions":["Accumulate into a vector element instead: load a 1-element vector, add, store back","Serialize the update (only one lane performs it) and use a normal load-add-store","Accumulate in VMEM arrays and reduce at the end"],"exampleFix":"# before\nref[i] += delta  # scalar atomic add on SC\n# after\nv = ref[pl.ds(i, 1)]\nv = v.at[0].add(delta)\nref[pl.ds(i, 1)] = v","handlingStrategy":"fallback","validationCode":"# accumulate via 1-element vector instead of scalar atomic add\nv = ref[pl.ds(i, 1)]\nref[pl.ds(i, 1)] = v + delta","typeGuard":"null","tryCatchPattern":"null","preventionTips":["Avoid scalar atomics on SC; accumulate in vectors or reduce at the end"],"tags":["jax","pallas","tpu","sparsecore","atomic","swap"],"backgroundTag":"unsupported-atomic-operation","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}