{"record":{"id":"fb4012dd0a2fffa4","repo":"jax-ml/jax","slug":"can-only-store-scalars-or-vectors-got-value","errorCode":null,"errorMessage":"Can only store scalars or vectors (got {value}).","messagePattern":"Can only store scalars or vectors \\(got (.+?)\\)\\.","errorType":"exception","errorClass":"TypeError","httpStatus":null,"severity":"error","filePath":"jax/_src/pallas/mosaic_gpu/lowering.py","lineNumber":2378,"sourceCode":"          old_value = mgpu.FragmentedArray.load_strided(\n              x_smem, is_signed=mgpu_utils.is_signed(v_aval.dtype)\n          )\n          value.store_untiled(x_smem)\n    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)","sourceCodeStart":2360,"sourceCodeEnd":2396,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/pallas/mosaic_gpu/lowering.py#L2360-L2396","documentation":"Warpgroup swap lowering requires the stored value to be a scalar (no shape) or an MLIR vector. A shaped value that is not an ir.VectorType (e.g. a FragmentedArray or a non-vector lowered value) is a TypeError.","triggerScenarios":"Calling x_smem[...] = value under warpgroup semantics where value has shape but its lowered type is not ir.VectorType — usually an internal lowering inconsistency or a custom primitive feeding the wrong IR type.","commonSituations":"Writing custom Pallas primitives/lowering rules; mixing lane-semantics FragmentedArray values into warpgroup stores.","solutions":["Ensure values stored under warpgroup semantics are lowered as vectors or scalars","Use the standard library ops (plgpu ops) rather than hand-building values for the store","File/check against recent JAX if using only public Pallas APIs"],"exampleFix":null,"handlingStrategy":"type-guard","validationCode":null,"typeGuard":"import jax._src.interpreters.mlir as mlir\ndef is_scalar_or_vector(v) -> bool:\n    return not getattr(v, 'type', None) and True or isinstance(getattr(v, 'type', None), mlir.ir.VectorType)","tryCatchPattern":null,"preventionTips":["Use only public Pallas ops under warpgroup semantics","Ensure stored values lower to vectors/scalars"],"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"}