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)

  @property

View on GitHub (pinned to 1e1c6a8fc0)

Solutions

  1. Transpose each individual ref in the group instead
  2. Split the multiref before transforming
  3. 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

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


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