jax-ml/jax · error · NotImplementedError
Transpose with multiref is not supported.
Error message
Transpose with multiref is not supported.
What it means
TransformedRef.transpose explicitly rejects multirefs with NotImplementedError, because a permutation over grouped refs is ambiguous (which ref's axes does it permute?).
Source
Thrown at jax/_src/state/types.py:377
return TransformedRef(self, (BitcastTransform(dtype),))
return TransformedRef(self.ref, (*self.transforms, BitcastTransform(dtype)))
def reshape(self, *shape):
if self.is_dynamic_size:
raise NotImplementedError(
"Reshape ref with dynamic size is not supported."
)
if len(shape) == 1 and isinstance(shape[0], tuple):
shape = shape[0]
input_shape = tuple(operator.index(s) for s in self.shape)
shape = _canonicalize_reshape(input_shape, shape)
if self.multiref:
return TransformedRef(self, (ReshapeTransform(shape),))
return TransformedRef(self.ref, (*self.transforms, ReshapeTransform(shape)))
def transpose(self, permutation: Sequence[int]):
if self.multiref:
raise NotImplementedError("Transpose with multiref is not supported.")
transposer = TransposeTransform(tuple(permutation))
if self.multiref:
return TransformedRef(self, (transposer,))
return TransformedRef(self.ref, (*self.transforms, transposer))
def set(self, value, idx=()):
from jax._src.state.primitives import ref_set # pyrefly: ignore[missing-import]
return ref_set(self, idx, value)
def swap(self, value, idx=()):
from jax._src.state.primitives import ref_swap # pyrefly: ignore[missing-import]
return ref_swap(self, idx, value)
def get(self, idx=()):
from jax._src.state.primitives import ref_get # pyrefly: ignore[missing-import]
return ref_get(self, idx)
@propertyView on GitHub (pinned to 1e1c6a8fc0)
Solutions
- Transpose each individual ref in the group instead
- Split the multiref before transforming
- Use .reshape only if semantically equivalent (reshape supports multiref)
Example fix
# before out = multiref.transpose(perm) # after out = tuple(r.transpose(perm) for r in multiref.ref)
Defensive patterns
Strategy: type-guard
Validate before calling
if isinstance(ref.ref, tuple):
raise TypeError("transpose on multiref not supported; transpose each ref") Type guard
def is_multiref(ref):
return isinstance(getattr(ref, "ref", None), tuple) Prevention
- Dispatch transforms based on single-ref vs multiref
- Use .T only on single refs
When it happens
Trigger: Calling .transpose(...) or .T on a TransformedRef whose .ref is a tuple of refs.
Common situations: Using .T on a group of refs; generic code applying transforms to both single refs and multirefs.
Related errors
- Pure callbacks do not support transpose. Please use `jax.cus
- transpose output pytree structure must match that of linear
- for transpose support, subclass {type(self)} must implement
- {type(_prim).__name__}.transpose should return None or a dic
- transpose_solve required for backwards mode automatic differ
AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27).
Data as JSON: /api/errors/d3d1e47d655b2125.
Report an issue: GitHub.