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
- Remove the transpose/reshape on the TMEM ref; do the transpose after loading into registers
- 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
- Never transpose or reshape TMEM refs; transform values after load
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
- Unsupported transform: {type(transform)}
- Unsupported TMEM ref {ref}.
- Non-indexing transforms on GMEM refs are not implemented.
- Unsupported transforms for {ref}. Transforms {transforms}.
- Not all transforms could be handled. Remaining transforms: {
AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27).
Data as JSON: /api/errors/d6653170a6337cac.
Report an issue: GitHub.