{"record":{"id":"3d2a1eb671e452d1","repo":"jax-ml/jax","slug":"can-only-store-to-references-got-x-smem","errorCode":null,"errorMessage":"Can only store to references (got {x_smem}).","messagePattern":"Can only store to references \\(got (.+?)\\)\\.","errorType":"exception","errorClass":"TypeError","httpStatus":null,"severity":"error","filePath":"jax/_src/pallas/mosaic_gpu/lowering.py","lineNumber":2382,"sourceCode":"    case _:\n      raise NotImplementedError(f\"Unsupported transforms: {transforms}\")\n  if ctx.module_ctx.auto_barriers:\n    barrier()  # Make sure the writes have completed.\n  return old_value\n\n\n@register_lowering_rule(sp.swap_p, mgpu.LoweringSemantics.Warpgroup)\n@register_lowering_rule(sp.swap_p, *gpu_core.WGxWARP_SEMANTICS)\ndef _swap_lowering_rule_wg(\n    ctx: LoweringRuleContext, x_smem, value, *leaves, tree\n):\n  shape = ctx.avals_out[0].shape\n  if shape and not isinstance(value.type, ir.VectorType):\n    raise TypeError(f\"Can only store scalars or vectors (got {value}).\")\n  if not (\n      isinstance(x_smem, ir.Value) and isinstance(x_smem.type, ir.MemRefType)\n  ):\n    raise TypeError(f\"Can only store to references (got {x_smem}).\")\n  if shape and ctx.module_ctx.primitive_semantics == gpu_core.PrimitiveSemantics.Warp:\n    raise NotImplementedError(\"Can only store scalars in warp-level lowering.\")\n  transforms = jax.tree.unflatten(tree, leaves)\n  transform_avals = jax.tree.unflatten(tree, ctx.avals_in[2:])\n  assert isinstance(ctx.avals_in[0], state_types.AbstractRef)\n  x_smem, _, transforms = _handle_transforms(\n      ctx, ctx.avals_in[0], x_smem, transform_avals, transforms,\n      allow_peer_refs=True\n  )\n  if transforms:\n    raise NotImplementedError(\n        \"Transforms are not yet implemented for warpgroup semantics\"\n    )\n  assert isinstance(x_smem, ir.Value)\n  value = _ensure_ir_value(value, ctx.avals_in[1].dtype)\n  if shape:\n    old_value = mgpu.dialect.vector_load(x_smem)\n    mgpu.dialect.vector_store(value, x_smem)","sourceCodeStart":2364,"sourceCodeEnd":2400,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/pallas/mosaic_gpu/lowering.py#L2364-L2400","documentation":"Warpgroup swap requires the destination to be an MLIR value of MemRefType (shared/global memory ref). A raw MemRefType or anything else is rejected, mirroring the lane-level type check for stores.","triggerScenarios":"Storing via swap_p under warpgroup semantics where x_smem is not a materialized memref value (internal misuse, peer refs, or custom lowering).","commonSituations":"Custom primitives or internal code paths passing unmaterialized refs; ref handling bugs after _handle_transforms.","solutions":["Pass a proper materialized shared-memory reference","If writing custom lowering code, lower the ref through the standard path so it becomes ir.Value with MemRefType"],"exampleFix":null,"handlingStrategy":"type-guard","validationCode":null,"typeGuard":"import jax._src.interpreters.mlir as mlir\ndef is_materialized_memref_value(x) -> bool:\n    return isinstance(x, mlir.ir.Value) and isinstance(x.type, mlir.ir.MemRefType)","tryCatchPattern":null,"preventionTips":["Pass standard refs to stores","Avoid custom lowering that leaks raw MemRefType"],"tags":["jax","pallas","mosaic-gpu","type-error","warpgroup"],"backgroundTag":"invalid-argument-type","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}