jax-ml/jax · error · NotImplementedError
Unsupported transforms for {ref}. Transforms {transforms}.
Error message
Unsupported transforms for {ref}. Transforms {transforms}. What it means
The warp-group variant of the tcgen05 MMA lowering routes the accumulator and operand references through _handle_transforms; if any transforms remain unhandled (notably transposes when handle_transposes cannot apply), it raises NotImplementedError naming the offending ref and transforms.
Source
Thrown at jax/_src/pallas/mosaic_gpu/primitives.py:2969
acc_transforms_leaves_avals,
a_transforms_leaves_avals,
b_transforms_leaves_avals,
_,
a_scale_transforms_leaves_avals,
b_scale_transforms_leaves_avals,
a_sparse_metadata_transforms_leaves_avals,
) = transforms_avals_lists
def handle_transforms_and_get_ref(tree, leaves, leaves_avals, ref, ref_aval, handle_transposes=True):
if tree is None:
return ref
transforms = tree.unflatten(leaves)
transform_avals = tree.unflatten(leaves_avals)
ref, _, transforms = lowering._handle_transforms(
ctx, ref_aval, ref, transform_avals, transforms, handle_transposes=handle_transposes
)
if transforms:
raise NotImplementedError(
f"Unsupported transforms for {ref}. Transforms {transforms}."
)
return ref
acc_ref = handle_transforms_and_get_ref(
acc_transforms_tree,
acc_transforms_leaves,
acc_transforms_leaves_avals,
acc_ref,
acc_aval,
handle_transposes=False,
)
a_ref = handle_transforms_and_get_ref(
a_transforms_tree,
a_transforms_leaves,
a_transforms_leaves_avals,
a_ref,View on GitHub (pinned to 1e1c6a8fc0)
Solutions
- Remove the remaining transforms on the named ref (message shows exactly which ref and transforms)
- Apply the transpose/layout change manually to the buffer contents instead of via a ref transform
Example fix
// before acc_ref_t = plgpu.transpose_ref(acc_ref) tcgen05_mma_wg(a, b, acc_ref_t) // after tcgen05_mma_wg(a, b, acc_ref)
Defensive patterns
Strategy: validation
Validate before calling
def check_ref_clean(ref, transforms):
if transforms:
raise NotImplementedError(f'clean {transforms} on {ref} first') Prevention
- Audit transform trees on all MMA refs before calling warp-group MMA
- Use handle_transposes-compatible layouts
When it happens
Trigger: Using the warp-group tcgen05 MMA lowering with a reference whose transform tree leaves residual transforms (e.g. transpose on a TMEM accumulator ref).
Common situations: Writing warp-group MMA kernels with transposed accumulators or exotic layouts on refs fed to the MMA.
Related errors
- Unsupported transform: {type(transform)}
- Sparse metadata format not implemented for {operand_dtype=}
- Unsupported TMEM ref {ref}.
- Non-indexing transforms on GMEM refs are not implemented.
- Unsupported transforms for ACC: {acc_transforms}.
AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27).
Data as JSON: /api/errors/d728e3b3c82fdf86.
Report an issue: GitHub.