{"record":{"id":"b5adbd6cecd797d3","repo":"jax-ml/jax","slug":"unsupported-transform-type-transform","errorCode":null,"errorMessage":"Unsupported transform: {type(transform)}","messagePattern":"Unsupported transform: (.+?)","errorType":"exception","errorClass":"NotImplementedError","httpStatus":null,"severity":"error","filePath":"jax/_src/pallas/mosaic_gpu/core.py","lineNumber":956,"sourceCode":"\n    return state_types.ReshapeTransform(new_shape), UntilingTransform(new_tiling)\n\n  def pretty_print(self, context: jax_core.JaxprPpContext) -> pp.Doc:\n    return pp.text(f\"{{untile({list(self.tiling)})}}\")\n\n\ndef batch_transform(\n    transform: state_types.Transform, leading_rank: int\n) -> state_types.Transform:\n  match transform:\n    case TransposeTransform() as t:\n      return TransposeTransform(\n          (*range(leading_rank), *(d + leading_rank for d in t.permutation))\n      )\n    case TilingTransform() | SwizzleTransform() as t:\n      return t\n    case _:\n      raise NotImplementedError(f\"Unsupported transform: {type(transform)}\")\n\n\ndef to_gpu_transform(\n    transform: state_types.Transform,\n) -> mgpu.MemRefTransform:\n  match transform:\n    case TransposeTransform(permutation):\n      return mgpu.TransposeTransform(permutation)\n    case TilingTransform(tiling):\n      return mgpu.TileTransform(tiling)\n    case _:\n      raise TypeError(f\"Unsupported transform: {type(transform)}\")\n\n\ndef to_transform_attr(\n    transform: state_types.Transform,\n) -> ir.Attribute:\n  match transform:","sourceCodeStart":938,"sourceCodeEnd":974,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/pallas/mosaic_gpu/core.py#L938-L974","documentation":"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.","triggerScenarios":"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.","commonSituations":"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.","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"],"exampleFix":null,"handlingStrategy":"type-guard","validationCode":"from jax._src.pallas.mosaic_gpu import core\nSUPPORTED = (state_types.TilingTransform, state_types.SwizzleTransform, state_types.TransposeTransform)","typeGuard":"def is_batch_supported(t): return isinstance(t, (TilingTransform, SwizzleTransform)) or hasattr(t, 'permutation')","tryCatchPattern":"catch NotImplementedError and inspect type(transform) in the message","preventionTips":["Keep custom transforms out of ref.transforms entering the GPU pipeline"],"tags":["jax","pallas","mosaic-gpu","transforms","not-implemented"],"backgroundTag":"unsupported-transform-type","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}