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

  1. Remove the remaining transforms on the named ref (message shows exactly which ref and transforms)
  2. 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

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


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