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

  1. Eliminate transforms on the b_scale reference; load the scale already in the required layout
  2. 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

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


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