jax-ml/jax · error · TypeError
Expected a single barrier, got a barrier reference with shap
Error message
Expected a single barrier, got a barrier reference with shape {transformed_barrier.shape} What it means
async_store_smem takes a barrier reference used to signal transaction completion; the abstract eval requires the transformed barrier to hold exactly one element (transformed_barrier.size == 1). A barrier with a non-scalar shape cannot be passed to the arrive_expect_tx operation.
Source
Thrown at jax/_src/pallas/mosaic_gpu/primitives.py:586
flat_ref_transforms_avals
)
barrier_transform_avals = barrier_transforms_treedef.unflatten(
flat_barrier_transforms_avals
)
transformed_ref = pallas_core.TransformedRef(ref, ref_transform_avals)
if src.shape != transformed_ref.shape:
raise TypeError(
f"The stored value has shape {src.shape}, but the target reference has"
f" shape {transformed_ref.shape}"
)
if src.dtype != transformed_ref.dtype:
raise TypeError(
f"The stored value has dtype {src.dtype}, but the target reference has"
f" dtype {transformed_ref.dtype}"
)
transformed_barrier = pallas_core.TransformedRef(barrier, barrier_transform_avals)
if transformed_barrier.size != 1:
raise TypeError(
"Expected a single barrier, got a barrier reference with shape"
f" {transformed_barrier.shape}"
)
effs = {gpu_core._memory_effect, state.WriteEffect(1)}
return (), effs
@lowering.register_lowering_rule(async_store_smem_p, mgpu.LoweringSemantics.Lane)
@lowering.register_lowering_rule(async_store_smem_p, mgpu.LoweringSemantics.Warpgroup)
def _async_store_smem_lowering(
ctx: lowering.LoweringRuleContext,
src,
ref,
barrier,
cluster_idx,
*flat_transforms,
ref_transforms_treedef,View on GitHub (pinned to 1e1c6a8fc0)
Solutions
- Index the barrier down to a single element, e.g. barrier[0] or barrier[i], before passing it
- Allocate barriers as scalar-shaped SMEM (pl.SMEM((), plint.barrier_dtype))
- Use the dedicated async_barrier/barrier APIs if per-warp barriers are needed
Example fix
# before async_store_smem(smem, x, barriers) # after async_store_smem(smem, x, barriers[0])
Defensive patterns
Strategy: validation
Validate before calling
assert barrier_ref.size == 1, 'async_store_smem needs a scalar barrier'
Prevention
- Allocate scalar barriers; index barrier arrays to a single element before passing
When it happens
Trigger: Passing an array-shaped barrier (e.g. a (4,) barrier vector) or slicing the barrier so more than one element remains, to async_store_smem.
Common situations: Allocating one barrier per warp/iteration as a vector and passing the whole array instead of a single element; mis-indexing barrier buffers with block indices meant for the value.
Related errors
- The stored value has shape {src.shape}, but the target refer
- The stored value has dtype {src.dtype}, but the target refer
- Scalars are not supported in async_store_smem
- Unexpected unhandled transforms: {remaining_ref_transforms}
- async_store_smem requires a tiled and swizzled ref
AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27).
Data as JSON: /api/errors/79e0d391bd1e95c0.
Report an issue: GitHub.