jax-ml/jax · error · NotImplementedError

Unsupported transforms: {a_sparse_metadata_transforms}

Error message

Unsupported transforms: {a_sparse_metadata_transforms}

What it means

For sparse MMA on Blackwell, the sparse-metadata reference (index/end pointers) must be passed without residual layout transforms; _handle_transforms must consume all of them or lowering fails with NotImplementedError.

Source

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

    if b_scale_transforms:
      raise NotImplementedError(f"Unsupported transforms: {b_scale_transforms}")
  if a_sparse_metadata_transforms_tree is not None:
    a_sparse_metadata_transforms = a_sparse_metadata_transforms_tree.unflatten(
        a_sparse_metadata_transforms_leaves
    )
    a_sparse_metadata_transform_avals = (
        a_sparse_metadata_transforms_tree.unflatten(
            a_sparse_metadata_transforms_leaves_avals
        )
    )
    assert isinstance(a_sparse_metadata_ref_aval, state_types.AbstractRef)
    a_sparse_metadata_ref, _, a_sparse_metadata_transforms = (
        lowering._handle_transforms(  # pyrefly: ignore[bad-specialization]
            ctx, a_sparse_metadata_ref_aval, a_sparse_metadata_ref,
            a_sparse_metadata_transform_avals, a_sparse_metadata_transforms)
    )
    if a_sparse_metadata_transforms:
      raise NotImplementedError(
          f"Unsupported transforms: {a_sparse_metadata_transforms}"
      )

  predicate = ctx.module_ctx.single_lane_predicate
  if collective_axis is not None:
    assert predicate is not None
    is_leader_block = _collective_mma_predicate(ctx, collective_axis)
    predicate = arith_dialect.andi(predicate, is_leader_block)
    collective = True
  else:
    collective = False

  with mgpu.when(predicate):
    tcgen05.mma(
        acc,
        a_ref,
        b_ref,
        a_swizzle=int(lhs_swizzle),

View on GitHub (pinned to 1e1c6a8fc0)

Solutions

  1. Pass the sparse metadata reference without transforms
  2. Rearrange metadata layout manually before the MMA call

Example fix

// before
tcgen05_mma(a_sparse, b, acc, a_sparse_metadata=meta_ref, a_sparse_metadata_transforms=transforms)
// after
tcgen05_mma(a_sparse, b, acc, a_sparse_metadata=meta_ref)
Defensive patterns

Strategy: validation

Validate before calling

assert not a_sparse_metadata_transforms

Prevention

When it happens

Trigger: Supplying a_sparse_metadata_ref with transform trees (transposes etc.) to tcgen05_mma such that transforms remain after handling.

Common situations: Writing 2:4 sparse GEMM kernels and routing metadata through transformed references.

Related errors


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