# Source code for tensorplay._transforms.einops.rearrange
Source: https://www.tensorplay.cn/docs/_modules/tensorplay/_transforms/einops/rearrange.html
```
"""Rearranging a tensor by naming its axes.
``rearrange(x, "b h w c -> b c h w")`` says what the axes *are*, so the reader
never has to decode a permutation of integers. Parentheses group axes:
``"(b h) w -> b h w"`` splits an axis whose factors are given as keyword
arguments, and ``"b h w -> b (h w)"`` merges two. ``...`` stands for any
number of leading or interior axes that pass straight through.
"""
from __future__ import annotations
import functools
from collections.abc import Sequence
from typing import Any, Union
__all__ = ["rearrange"]
_ELLIPSIS = "…"
class _AnonymousAxis:
"""A literal axis length written in the pattern, such as the ``1`` in
``"b h w -> b h w 1"``. Distinct occurrences stay distinct."""
__slots__ = ("value",)
def __init__(self, value: str) -> None:
self.value = int(value)
if self.value < 1:
raise ValueError(
f"Anonymous axis should be a positive integer, not {self.value}"
)
def __repr__(self) -> str:
return f"{self.value}-axis"
def _is_valid_identifier(name: str) -> bool:
if name == _ELLIPSIS:
return True
if not str.isidentifier(name):
return False
if name[0] == "_" or name[-1] == "_":
return False
return True
class _ParsedExpression:
"""One side of a pattern, as a list of axis groups."""
def __init__(self, expression: str, *, allow_underscore: bool = False) -> None:
self.has_ellipsis = False
self.has_ellipsis_parenthesized = False
self.identifiers: set[Any] = set()
self.composition: list[Union[list[Any], str]] = []
if "." in expression:
if "..." not in expression:
raise ValueError(
"Expression may contain dots only inside ellipsis (...)"
)
if str.count(expression, "...") != 1 or str.count(expression, ".") != 3:
raise ValueError(
"Expression may contain dots only inside ellipsis (...); only "
"one ellipsis for tensor is allowed"
)
expression = expression.replace("...", _ELLIPSIS)
self.has_ellipsis = True
bracket_group: list[Any] | None = None
def add_axis_name(name: str) -> None:
if name in self.identifiers:
if not (allow_underscore and name == "_"):
raise ValueError(f"Indexing expression contains duplicate axis {name}")
if name == _ELLIPSIS:
self.identifiers.add(_ELLIPSIS)
if bracket_group is None:
self.composition.append(_ELLIPSIS)
else:
bracket_group.append(_ELLIPSIS)
self.has_ellipsis_parenthesized = True
return
is_number = str.isdecimal(name)
if is_number and int(name) == 1:
# A literal 1 adds nothing to a group and is a bare axis alone.
if bracket_group is None:
self.composition.append([])
return
axis_name: Any = _AnonymousAxis(name) if is_number else name
if not (is_number or _is_valid_identifier(name)):
raise ValueError(f"Invalid axis identifier: {name}")
self.identifiers.add(axis_name)
if bracket_group is None:
self.composition.append([axis_name])
else:
bracket_group.append(axis_name)
current_identifier = None
for char in expression:
if char in "() ":
if current_identifier is not None:
add_axis_name(current_identifier)
current_identifier = None
if char == "(":
if bracket_group is not None:
raise ValueError("Axis composition is one-level (brackets are not allowed inside brackets)")
bracket_group = []
elif char == ")":
if bracket_group is None:
raise ValueError("Brackets are not balanced")
self.composition.append(bracket_group)
bracket_group = None
elif str.isalnum(char) or char in ["_", _ELLIPSIS]:
if current_identifier is None:
current_identifier = char
else:
current_identifier += char
else:
raise ValueError(f"Unknown character '{char}'")
if bracket_group is not None:
raise ValueError(f"Imbalanced parentheses in expression: {expression}")
if current_identifier is not None:
add_axis_name(current_identifier)
def _report(pattern: str, message: str) -> ValueError:
return ValueError(f"{message}\n Expression: '{pattern}'")
@functools.lru_cache(256)
def _parse_pattern(pattern: str) -> tuple[_ParsedExpression, _ParsedExpression]:
if "->" not in pattern:
raise ValueError(f"Pattern must contain '->'\n Expression: '{pattern}'")
left_str, right_str = pattern.split("->")
left = _ParsedExpression(left_str)
right = _ParsedExpression(right_str)
if not left.has_ellipsis and right.has_ellipsis:
raise _report(pattern, f"Ellipsis found in right side, but not left side of a pattern")
if left.has_ellipsis and left.has_ellipsis_parenthesized:
raise _report(pattern, f"Ellipsis is parenthesis in the left side is not allowed")
return left, right
[docs]
def rearrange(
tensor: Any,
pattern: str,
**axes_lengths: int,
) -> Any:
"""Reshapes and permutes ``tensor`` according to ``pattern``.
Args:
tensor: the tensor to rearrange, or a sequence of tensors, which is
stacked along a new leading axis first.
pattern (str): ``" ->