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
- Remove the unsupported transform from the stack before entering the mosaic pipeline (materialize/undo it)
- Implement a matching case in batch_transform for your custom transform (if vendoring/patching)
- 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
- Keep custom transforms out of ref.transforms entering the GPU pipeline
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
- Non-indexing transforms on GMEM refs are not implemented.
- Not all transforms could be handled. Remaining transforms: {
- Unsupported dtype: {ref.dtype}
- Only SMEM and TMEM refs are supported.
- Transpose cannot be moved before a tiling transform when it
AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27).
Data as JSON: /api/errors/b5adbd6cecd797d3.
Report an issue: GitHub.