jax-ml/jax · error · NotImplementedError
Unsupported transforms: {b_scale_transforms}
Error message
Unsupported transforms: {b_scale_transforms} What it means
Same restriction as the A-scale case but for the B operand's scale factor reference: after lowering._handle_transforms runs, any residual transforms on b_scale_ref are unsupported and raise NotImplementedError.
Source
Thrown at jax/_src/pallas/mosaic_gpu/primitives.py:2820
)
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(
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}"View on GitHub (pinned to 1e1c6a8fc0)
Solutions
- Eliminate transforms on the b_scale reference; load the scale already in the required layout
- Materialize any transpose of the scale tensor manually before the MMA
Example fix
// before tcgen05_mma(a, b, acc, b_scale=scale_ref, b_scale_transforms=transforms) // after tcgen05_mma(a, b, acc, b_scale=scale_ref)
Defensive patterns
Strategy: validation
Validate before calling
assert not b_scale_transforms or all(t is None for t in b_scale_transforms)
Prevention
- Materialize scale layouts manually in SMEM
- Keep scale refs transform-free
When it happens
Trigger: Passing b_scale_ref with remaining transforms (transposes etc.) in b_scale_transforms_tree to tcgen05_mma.
Common situations: Block-scaled (fp8) kernels where the B-scale tensor is accessed through a transformed reference.
Related errors
- Unsupported transforms: {a_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/d68f4e27d6117820.
Report an issue: GitHub.