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
- Let transform inference compute the output tiling instead of annotating it manually
- Or set out tiling = tuple(dims[p[i]] ... ) i.e. the input tiling permuted by the permutation
- 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
- Let inference compute output tilings
- Derive out tiling by permuting the in tiling
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
- Transpose cannot be moved before a tiling transform when it
- Cannot tile a transpose ({permutation}). Tiling dims ({permu
- Input/output tiling mismatch when attempting to collapse a s
- Commuting a `UntilingTransform` with a `ReshapeTransform` is
- Folding tiled dimensions into untiled dimensions is not supp
AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27).
Data as JSON: /api/errors/f8a4bf6f80915e49.
Report an issue: GitHub.