jax-ml/jax · error · NotImplementedError

Unimplemented transforms for TMEM refs. {transforms=}

Error message

Unimplemented transforms for TMEM refs. {transforms=}

What it means

When lowering a TMEM operation, residual transforms remained after _handle_transforms with transposes and reshapes disabled; TMEM refs do not support those transforms. Any leftover transform raises NotImplementedError.

Source

Thrown at jax/_src/pallas/mosaic_gpu/primitives.py:4107

def _async_load_tmem_lowering_rule(
    ctx: lowering.LoweringRuleContext,
    x_ref,
    *leaves,
    tree,
    reduce: Literal["max", "min", "absmax", "absmin"] | None = None,
):
  assert isinstance(x_ref, tcgen05.TMEMRef)
  x_aval = ctx.avals_in[0]
  assert isinstance(x_aval, state_types.AbstractRef)
  transforms = jax.tree.unflatten(tree, leaves)
  transform_avals = tree.unflatten(
      ctx.avals_in[1 : 1 + tree.num_leaves]
  )
  x_tmem, _, transforms = lowering._handle_transforms(
      ctx, x_aval, x_ref, transform_avals, transforms, handle_transposes=False,
      handle_reshapes=False)
  if transforms:
    raise NotImplementedError(
        f"Unimplemented transforms for TMEM refs. {transforms=}"
    )
  layout_hint = None
  if isinstance(ctx.out_layout_hint, mgpu.TiledLayout):
    layout_hint = ctx.out_layout_hint
  is_signed = mgpu_utils.is_signed(ctx.avals_out[0].dtype)
  res = x_tmem.load(layout=layout_hint, is_signed=is_signed, reduce=reduce)
  return (res,) if reduce is None else res


@lowering.register_lowering_rule(
    async_load_tmem_p, mgpu.LoweringSemantics.Warpgroup
)
def _async_load_tmem_lowering_rule_wg(
    ctx: lowering.LoweringRuleContext,
    x_ref: ir.Value,
    *leaves,
    tree,

View on GitHub (pinned to 1e1c6a8fc0)

Solutions

  1. Remove the transpose/reshape on the TMEM ref; do the transpose after loading into registers
  2. Allocate the TMEM ref with the final desired shape/layout instead of transforming it

Example fix

// before
x = load(tmem_ref.T)
// after
x = load(tmem_ref).T
Defensive patterns

Strategy: validation

Validate before calling

assert not transforms, f'leftover TMEM transforms: {transforms}'

Prevention

When it happens

Trigger: Applying a transpose or reshape (or any transform) to a TMEM ref before an operation lowered with handle_transposes=False and handle_reshapes=False.

Common situations: Writing tmem_ref.T or reshaping a TMEM ref; assuming TMEM behaves like SMEM regarding transforms.

Related errors


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