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

  1. Index the barrier down to a single element, e.g. barrier[0] or barrier[i], before passing it
  2. Allocate barriers as scalar-shaped SMEM (pl.SMEM((), plint.barrier_dtype))
  3. 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

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


AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27). Data as JSON: /api/errors/79e0d391bd1e95c0. Report an issue: GitHub.