{"record":{"id":"24150c2d5773912e","repo":"jax-ml/jax","slug":"swap-only-supports-scalars-in-smem","errorCode":null,"errorMessage":"Swap only supports scalars in SMEM.","messagePattern":"Swap only supports scalars in SMEM\\.","errorType":"exception","errorClass":"NotImplementedError","httpStatus":null,"severity":"error","filePath":"jax/_src/pallas/mosaic/sc_lowering.py","lineNumber":229,"sourceCode":"  else:\n    first_nontrivial_dim = len(sizes)\n  if any(squeeze_dims[first_nontrivial_dim:]):\n    raise NotImplementedError(\n        \"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:","sourceCodeStart":211,"sourceCodeEnd":247,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/pallas/mosaic/sc_lowering.py#L211-L247","documentation":"When swapping a scalar on SC, a mask cannot be applied — the lowering raises immediately for masked scalar stores. Bounds must be handled by clamping or branching instead.","triggerScenarios":"swap=True store with mask set and 0-d value in a SparseCore kernel.","commonSituations":"Conditional scatter-style updates using masked swaps.","solutions":["Clamp or guard the index so no mask is needed","Swap a 1-element vector with a mask and extract element 0 afterwards"],"exampleFix":"# before\nold = pallas_swap(ref, idx, val, mask=m)  # scalar + mask\n# after\nidx = jnp.where(m, idx, 0)\nold = pallas_swap(ref, idx, val)","handlingStrategy":"fallback","validationCode":"if mask is not None and out.ndim == 0:\n    idx = jnp.where(mask, idx, 0)","typeGuard":"null","tryCatchPattern":"null","preventionTips":["Guard scalar swaps with index clamping instead of masks"],"tags":["jax","pallas","tpu","sparsecore","masked-store","swap"],"backgroundTag":"unsupported-masked-operation","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}