jax-ml/jax · error · NotImplementedError
Unsupported transforms: {a_scale_transforms}
Error message
Unsupported transforms: {a_scale_transforms} What it means
When using scaled MMA (fp8/fp4 block scaling on the A operand), any layout transforms (e.g. transposes) left over on the a_scale reference after _handle_transforms cannot be lowered, so the lowering raises NotImplementedError.
Source
Thrown at jax/_src/pallas/mosaic_gpu/primitives.py:2804
accumulate = mgpu.c(accumulate, ir.IntegerType.get_signless(1))
elif isinstance(accumulate, mgpu.FragmentedArray):
accumulate = accumulate.registers.item()
assert isinstance(accumulate, ir.Value)
if a_scale_ref is not None and a_scale_transforms_tree is not None:
assert isinstance(a_scale_ref_aval, state.AbstractRef)
a_scale_transforms = a_scale_transforms_tree.unflatten(
a_scale_transforms_leaves
)
a_scale_transform_avals = a_scale_transforms_tree.unflatten(
a_scale_transforms_leaves_avals
)
a_scale_ref, _, a_scale_transforms = lowering._handle_transforms(
ctx, a_scale_ref_aval, a_scale_ref, a_scale_transform_avals,
a_scale_transforms
)
if a_scale_transforms:
raise NotImplementedError(
f"Unsupported transforms: {a_scale_transforms}"
)
if b_scale_ref is not None and b_scale_transforms_tree is not None:
assert isinstance(b_scale_ref_aval, state.AbstractRef)
b_scale_transforms = b_scale_transforms_tree.unflatten(
b_scale_transforms_leaves
)
b_scale_transform_avals = b_scale_transforms_tree.unflatten(
b_scale_transforms_leaves_avals
)
b_scale_ref, _, b_scale_transforms = lowering._handle_transforms(
ctx, b_scale_ref_aval, b_scale_ref, b_scale_transform_avals,
b_scale_transforms
)
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(View on GitHub (pinned to 1e1c6a8fc0)
Solutions
- Remove the transforms on the a_scale reference (load it un-transformed; materialize the transpose manually)
- Pre-transpose/rearrange the scale tensor in SMEM before passing it to the MMA
Example fix
// before tcgen05_mma(a, b, acc, a_scale=scale_ref, a_scale_transforms=transforms_with_transpose) // after tcgen05_mma(a, b, acc, a_scale=scale_ref) # scale loaded already in the right layout
Defensive patterns
Strategy: validation
Validate before calling
assert not a_scale_transforms or all(t is None for t in a_scale_transforms)
Prevention
- Load scale tensors with the exact final layout
- Avoid applying BlockSpec transforms to scale refs
When it happens
Trigger: Passing a_scale_ref to tcgen05_mma together with transform trees (e.g. a transpose transform) that apply to the A-scale reference and cannot be handled during lowering.
Common situations: Building fp8 block-scaled GEMM kernels where the scale tensors are loaded from transposed or transformed references.
Related errors
- Unsupported transforms: {b_scale_transforms}
- Sparse metadata format not implemented for {operand_dtype=}
- Unsupported TMEM ref {ref}.
- Unsupported transforms for ACC: {acc_transforms}.
- Unsupported transforms for LHS: {a_transforms}.
AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27).
Data as JSON: /api/errors/1eeda67d60ef78ae.
Report an issue: GitHub.