jax-ml/jax · error · ValueError

Invalid in/out transforms. In transform: {in_transform}, out

Error message

Invalid in/out transforms. In transform: {in_transform}, out transform: {transform}

What it means

After applying the permutation to the input transform's tiling, the result must equal the output transform's declared tiling. This error means the annotated output tiling is not the permutation of the input tiling as the transpose requires.

Source

Thrown at jax/experimental/mosaic/gpu/dialect_lowering.py:2180

        in_transform, lc.TileTransform
    ):
      raise NotImplementedError(
          f"Unsupported in/out transforms. In transform: {in_transform}, out"
          f" transform: {transform}"
      )
    tiling_len = len(in_transform.tiling)
    tiling_offset = len(permutation) - tiling_len
    if any(dim < tiling_offset for dim in permutation[-tiling_len :]):
      raise ValueError(
          f"Cannot tile a transpose ({permutation}). Tiling dims"
          f" ({permutation[-tiling_len:]}) cannot contain non-tiled dims."
          f" All of them must be >= {tiling_offset}."
      )
    dims = [-1] * len(permutation)
    dims = dims[:-tiling_len] + list(in_transform.tiling)
    permuted_dims = tuple(dims[permutation[i]] for i in range(len(dims)))
    if permuted_dims[-tiling_len:] != transform.tiling:
      raise ValueError(
          f"Invalid in/out transforms. In transform: {in_transform}, out"
          f" transform: {transform}"
      )
    new_permutation = permutation + [
        x + tiling_len for x in permutation[-tiling_len:]
    ]
    new_permutation = _permutation_to_affine_map_attr(new_permutation)

  new_transpose_op = memref.TransposeOp(
      transform_type(ir.MemRefType(op.result.type), out_transforms),
      unwrapped_in_ref,
      new_permutation,
  )

  out_transforms = inference_utils.out_transforms(op)[0]
  wrapped_ref = wrap_transformed_memref(
      new_transpose_op.result, op.result.type, out_transforms
  )

View on GitHub (pinned to 1e1c6a8fc0)

Solutions

  1. Let transform inference compute the output tiling instead of annotating it manually
  2. Or set out tiling = tuple(dims[p[i]] ... ) i.e. the input tiling permuted by the permutation
  3. Verify by applying the permutation to in_transform.tiling and comparing to out_transform.tiling before emitting the op
Defensive patterns

Strategy: validation

Validate before calling

dims = [-1]*len(perm)
dims = dims[:-tl] + list(in_t.tiling)
expected = tuple(dims[perm[i]] for i in range(len(dims)))[-tl:]
assert expected == out_t.tiling, f'expected out tiling {expected}'

Prevention

When it happens

Trigger: memref.transpose with transforms where permuted input tiling != output transform's tiling, e.g. in tiling (2,4) and permutation that maps it to (4,2) but out tiling declared (2,4).

Common situations: Hand-specifying the output transform without deriving it from the permutation; inference usually computes it for you when you omit the annotation.

Related errors


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