jax-ml/jax · error · NotImplementedError

Unsupported transform: {type(transform)}

Error message

Unsupported transform: {type(transform)}

What it means

batch_transform converts state transforms to mosaic-GPU transforms and hit a transform type not in its supported set (TransposeTransform via permutation, TilingTransform, SwizzleTransform). Anything else — including new or GPU-specific transforms like PeerMemRef — raises NotImplementedError.

Source

Thrown at jax/_src/pallas/mosaic_gpu/core.py:956

    return state_types.ReshapeTransform(new_shape), UntilingTransform(new_tiling)

  def pretty_print(self, context: jax_core.JaxprPpContext) -> pp.Doc:
    return pp.text(f"{{untile({list(self.tiling)})}}")


def batch_transform(
    transform: state_types.Transform, leading_rank: int
) -> state_types.Transform:
  match transform:
    case TransposeTransform() as t:
      return TransposeTransform(
          (*range(leading_rank), *(d + leading_rank for d in t.permutation))
      )
    case TilingTransform() | SwizzleTransform() as t:
      return t
    case _:
      raise NotImplementedError(f"Unsupported transform: {type(transform)}")


def to_gpu_transform(
    transform: state_types.Transform,
) -> mgpu.MemRefTransform:
  match transform:
    case TransposeTransform(permutation):
      return mgpu.TransposeTransform(permutation)
    case TilingTransform(tiling):
      return mgpu.TileTransform(tiling)
    case _:
      raise TypeError(f"Unsupported transform: {type(transform)}")


def to_transform_attr(
    transform: state_types.Transform,
) -> ir.Attribute:
  match transform:

View on GitHub (pinned to 1e1c6a8fc0)

Solutions

  1. Remove the unsupported transform from the stack before entering the mosaic pipeline (materialize/undo it)
  2. Implement a matching case in batch_transform for your custom transform (if vendoring/patching)
  3. Check for version mismatch: upgrade or downgrade jax so the transform set is consistent
Defensive patterns

Strategy: type-guard

Validate before calling

from jax._src.pallas.mosaic_gpu import core
SUPPORTED = (state_types.TilingTransform, state_types.SwizzleTransform, state_types.TransposeTransform)

Type guard

def is_batch_supported(t): return isinstance(t, (TilingTransform, SwizzleTransform)) or hasattr(t, 'permutation')

Try / catch

catch NotImplementedError and inspect type(transform) in the message

Prevention

When it happens

Trigger: Passing a transform object that is not TilingTransform/SwizzleTransform/a transposable transform into the mosaic GPU pipeline, e.g. via a custom Transform subclass added to ref.transforms.

Common situations: Extending pallas with custom transforms; version skew where a transform exists in jax but isn't handled in mosaic_gpu core; user code manipulating TransformedRef.transforms directly.

Related errors


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