{"record":{"id":"cf78a9e1b57e7f3c","repo":"jax-ml/jax","slug":"can-only-store-scalars-in-warp-level-lowering","errorCode":null,"errorMessage":"Can only store scalars in warp-level lowering.","messagePattern":"Can only store scalars in warp-level lowering\\.","errorType":"exception","errorClass":"NotImplementedError","httpStatus":null,"severity":"error","filePath":"jax/_src/pallas/mosaic_gpu/lowering.py","lineNumber":2264,"sourceCode":"    return mgpu.dialect.vector_load(x_ref, optimized=optimized)\n  else:\n    return memref_dialect.load(x_ref, [])\n\n\n@register_lowering_rule(sp.swap_p, mgpu.LoweringSemantics.Lane)\n@register_lowering_rule(sp.swap_p, *gpu_core.LANExWARP_SEMANTICS)\ndef _swap_lowering_rule(\n    ctx: LoweringRuleContext, x_ref, value, *leaves, tree\n):\n  if isinstance(x_ref, tcgen05.TMEMRef):\n    raise RuntimeError(\n        \"Stores to TMEM are asynchronous operations and cannot be performed\"\n        \" using the usual syntax. Please use plgpu.async_store_tmem instead.\"\n    )\n  barrier = mgpu.warpgroup_barrier\n  if ctx.module_ctx.primitive_semantics == gpu_core.PrimitiveSemantics.Warp:\n    if ctx.avals_out[0].shape:\n      raise NotImplementedError(\"Can only store scalars in warp-level lowering.\")\n    i32 = ir.IntegerType.get_signless(32)\n    barrier = functools.partial(\n        nvvm_dialect.bar_warp_sync, arith_dialect.constant(i32, -1)\n    )\n  value = _ensure_fa(value, ctx.avals_in[1].dtype)\n\n  if not isinstance(x_ref, ir.Value) and isinstance(x_ref, ir.MemRefType):\n    raise TypeError(f\"Can only store to references (got {x_ref}).\")\n  v_aval = ctx.avals_in[1]\n  transforms = jax.tree.unflatten(tree, leaves)\n  transform_avals = jax.tree.unflatten(tree, ctx.avals_in[2:])\n\n  if ctx.module_ctx.auto_barriers:\n    barrier()  # Make sure reads have completed before we write.\n\n  if transforms and isinstance(transforms[0], gpu_core.UnswizzleRef):\n    swizzle = transforms[0].swizzle\n    transforms = transforms[1:]","sourceCodeStart":2246,"sourceCodeEnd":2282,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/pallas/mosaic_gpu/lowering.py#L2246-L2282","documentation":"In warp-level lowering of swap (store), only scalar values can be stored. Storing a shaped value under PrimitiveSemantics.Warp has no defined layout, so it raises NotImplementedError.","triggerScenarios":"Assigning a tensor (non-scalar) to a ref inside a plgpu.warp_semantics region: x_ref[...] = vec where avals_out[0].shape is non-empty under Warp semantics.","commonSituations":"Converting a lane-level store loop into warp-level code without making stores per-scalar; using vector ops inside warp semantics.","solutions":["Store scalar elements (loop or broadcast per-lane scalars) inside warp semantics","Hoist the tensor store out to warpgroup/lane semantics","Verify the BlockSpec shape vs the semantics decorator so stores remain scalar"],"exampleFix":"# before\nwith plgpu.warp_semantics():\n  x_ref[...] = vec  # vec has shape\n# after\nfor i in range(vec.shape[0]):\n  x_ref[i] = vec[i]  # scalar stores, or move store out of warp region","handlingStrategy":"validation","validationCode":"assert np.ndim(value_to_store) == 0 or semantics != 'warp'","typeGuard":null,"tryCatchPattern":null,"preventionTips":["Make stores scalar inside warp semantics","Hoist tensor stores out of warp-level regions"],"tags":["jax","pallas","mosaic-gpu","warp-semantics","store"],"backgroundTag":"gpu-kernel-semantics-mismatch","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}