{"record":{"id":"d3d1e47d655b2125","repo":"jax-ml/jax","slug":"transpose-with-multiref-is-not-supported","errorCode":null,"errorMessage":"Transpose with multiref is not supported.","messagePattern":"Transpose with multiref is not supported\\.","errorType":"exception","errorClass":"NotImplementedError","httpStatus":null,"severity":"error","filePath":"jax/_src/state/types.py","lineNumber":377,"sourceCode":"      return TransformedRef(self, (BitcastTransform(dtype),))\n    return TransformedRef(self.ref, (*self.transforms, BitcastTransform(dtype)))\n\n  def reshape(self, *shape):\n    if self.is_dynamic_size:\n      raise NotImplementedError(\n          \"Reshape ref with dynamic size is not supported.\"\n      )\n    if len(shape) == 1 and isinstance(shape[0], tuple):\n      shape = shape[0]\n    input_shape = tuple(operator.index(s) for s in self.shape)\n    shape = _canonicalize_reshape(input_shape, shape)\n    if self.multiref:\n      return TransformedRef(self, (ReshapeTransform(shape),))\n    return TransformedRef(self.ref, (*self.transforms, ReshapeTransform(shape)))\n\n  def transpose(self, permutation: Sequence[int]):\n    if self.multiref:\n      raise NotImplementedError(\"Transpose with multiref is not supported.\")\n    transposer = TransposeTransform(tuple(permutation))\n    if self.multiref:\n      return TransformedRef(self, (transposer,))\n    return TransformedRef(self.ref, (*self.transforms, transposer))\n\n  def set(self, value, idx=()):\n    from jax._src.state.primitives import ref_set  # pyrefly: ignore[missing-import]\n    return ref_set(self, idx, value)\n\n  def swap(self, value, idx=()):\n    from jax._src.state.primitives import ref_swap  # pyrefly: ignore[missing-import]\n    return ref_swap(self, idx, value)\n\n  def get(self, idx=()):\n    from jax._src.state.primitives import ref_get  # pyrefly: ignore[missing-import]\n    return ref_get(self, idx)\n\n  @property","sourceCodeStart":359,"sourceCodeEnd":395,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/state/types.py#L359-L395","documentation":"TransformedRef.transpose explicitly rejects multirefs with NotImplementedError, because a permutation over grouped refs is ambiguous (which ref's axes does it permute?).","triggerScenarios":"Calling .transpose(...) or .T on a TransformedRef whose .ref is a tuple of refs.","commonSituations":"Using .T on a group of refs; generic code applying transforms to both single refs and multirefs.","solutions":["Transpose each individual ref in the group instead","Split the multiref before transforming","Use .reshape only if semantically equivalent (reshape supports multiref)"],"exampleFix":"# before\nout = multiref.transpose(perm)\n# after\nout = tuple(r.transpose(perm) for r in multiref.ref)","handlingStrategy":"type-guard","validationCode":"if isinstance(ref.ref, tuple):\n    raise TypeError(\"transpose on multiref not supported; transpose each ref\")","typeGuard":"def is_multiref(ref):\n    return isinstance(getattr(ref, \"ref\", None), tuple)","tryCatchPattern":null,"preventionTips":["Dispatch transforms based on single-ref vs multiref","Use .T only on single refs"],"tags":["jax","transpose","multiref"],"backgroundTag":"unsupported-operation-on-group","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}