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
- Pass the sparse metadata reference without transforms
- 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
- Pass sparse metadata without ref transforms
- Check transforms survive _handle_transforms in tests
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
- Sparse metadata format not implemented for {operand_dtype=}
- Expected metadata dtype to be uint2, got: {meta.dtype}
- Expected metadata to be 3-dimensional (M, K // 4, 2), but it
- Expected the trailing dimension of the metadata to be 2, got
- Unsupported TMEM ref {ref}.
AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27).
Data as JSON: /api/errors/659c330f470fa88b.
Report an issue: GitHub.