From 4fa4e4ea568a2921e8ec0e2c7a4d2ea1eb9c9631 Mon Sep 17 00:00:00 2001 From: Gausshj Date: Mon, 3 Nov 2025 23:33:47 +0800 Subject: [PATCH] - Fix svd arrow handling logic - Add support for contract - Add basic test cases for contract --- grassmann_tensor/tensor.py | 257 +++++++++++++++++++++++++++++++------ tests/contract_test.py | 34 +++++ tests/exponential_test.py | 93 ++++++++++++-- tests/identity_test.py | 83 ++++++++++++ tests/reshape_test.py | 2 + tests/svd_test.py | 28 +--- 6 files changed, 426 insertions(+), 71 deletions(-) create mode 100644 tests/contract_test.py create mode 100644 tests/identity_test.py diff --git a/grassmann_tensor/tensor.py b/grassmann_tensor/tensor.py index 21d6e5b..ee4cef3 100644 --- a/grassmann_tensor/tensor.py +++ b/grassmann_tensor/tensor.py @@ -308,7 +308,10 @@ def reshape(self, new_shape: tuple[int | tuple[int, int], ...]) -> GrassmannTens if (isinstance(new_shape_check, int) and new_shape_check == 1) or ( new_shape_check == (1, 0) ): - arrow.append(False) + if cursor_plan < len(self.arrow): + arrow.append(self.arrow[cursor_plan]) + else: + arrow.append(False) edges.append((1, 0)) shape.append(1) cursor_plan += 1 @@ -573,6 +576,42 @@ def _group_edges( self, pairs: tuple[int, ...] | tuple[tuple[int, ...], tuple[int, ...]], ) -> tuple[GrassmannTensor, tuple[int, ...], tuple[int, ...]]: + return self.group_edges(self, pairs) + + @staticmethod + def group_edges( + tensor: GrassmannTensor, + pairs: tuple[int, ...] | tuple[tuple[int, ...], tuple[int, ...]], + ) -> tuple[GrassmannTensor, tuple[int, ...], tuple[int, ...]]: + left_legs, right_legs = GrassmannTensor.get_legs_pair(tensor.tensor.dim(), pairs) + + order = left_legs + right_legs + + tensor = tensor.permute(order) + + left_dim = math.prod(tensor.tensor.shape[: len(left_legs)]) + right_dim = math.prod(tensor.tensor.shape[len(left_legs) :]) + + tensor = tensor.reshape((left_dim, right_dim)) + + return tensor, left_legs, right_legs + + @staticmethod + def get_legs_pair( + dim: int, pairs: tuple[int, ...] | tuple[tuple[int, ...], tuple[int, ...]] + ) -> tuple[tuple[int, ...], tuple[int, ...]]: + def check_pairs_coverage(dim: int, pairs: tuple[tuple[int, ...], tuple[int, ...]]) -> bool: + set0 = set(pairs[0]) + set1 = set(pairs[1]) + + are_disjoint = set0.isdisjoint(set1) + + is_complete_union = (set0 | set1) == set(range(dim)) + + no_duplicates = len(pairs[0]) + len(pairs[1]) == dim + + return are_disjoint and is_complete_union and no_duplicates + if (isinstance(pairs, tuple) and len(pairs)) and all( isinstance(x, tuple) and all(isinstance(i, int) for i in x) for x in pairs ): @@ -580,22 +619,18 @@ def _group_edges( right_legs = typing.cast(tuple[int, ...], pairs[1]) else: left_legs = typing.cast(tuple[int, ...], pairs) - right_legs = tuple(i for i in range(self.tensor.dim()) if i not in left_legs) + right_legs = tuple(i for i in range(dim) if i not in left_legs) - assert self._check_pairs_coverage((left_legs, right_legs)), ( + assert check_pairs_coverage(dim, (left_legs, right_legs)), ( f"Input pairs must cover all dimension and disjoint, but got {(left_legs, right_legs)}" ) - order = left_legs + right_legs - - tensor = self.permute(order) + return left_legs, right_legs - left_dim = math.prod(tensor.tensor.shape[: len(left_legs)]) - right_dim = math.prod(tensor.tensor.shape[len(left_legs) :]) - - tensor = tensor.reshape((left_dim, right_dim)) - - return tensor, left_legs, right_legs + def _get_legs_pair( + self, pairs: tuple[int, ...] | tuple[tuple[int, ...], tuple[int, ...]] + ) -> tuple[tuple[int, ...], tuple[int, ...]]: + return self.get_legs_pair(self.tensor.dim(), pairs) def svd( self, @@ -622,7 +657,25 @@ def svd( if isinstance(cutoff, tuple): assert len(cutoff) == 2, "The length of cutoff must be 2 if cutoff is a tuple." - tensor, left_legs, right_legs = self._group_edges(free_names_u) + left_legs, right_legs = self._get_legs_pair(free_names_u) + order = left_legs + right_legs + tensor = self.permute(order) + + arrow_reverse = tuple(i for i, current in enumerate(tensor.arrow) if current) + if arrow_reverse: + tensor = tensor.reverse(arrow_reverse).reverse(arrow_reverse).reverse(arrow_reverse) + + left_dim = math.prod(tensor.tensor.shape[: len(left_legs)]) + right_dim = math.prod(tensor.tensor.shape[len(left_legs) :]) + tensor = tensor.reshape((left_dim, right_dim)) + + origin_arrow_left = tuple(self.arrow[i] for i in left_legs) + origin_arrow_right = tuple(self.arrow[i] for i in right_legs) + + arrow_reverse_left = tuple(i for i, current in enumerate(origin_arrow_left) if current) + arrow_reverse_right = tuple( + i + 1 for i, current in enumerate(origin_arrow_right) if current + ) (even_left, odd_left) = tensor.edges[0] (even_right, odd_right) = tensor.edges[1] @@ -700,7 +753,7 @@ def svd( (Vh_even_trunc.shape[1], Vh_odd_trunc.shape[1]), ) - U = GrassmannTensor(_arrow=(True, True), _edges=U_edges, _tensor=U_tensor) + U = GrassmannTensor(_arrow=(False, True), _edges=U_edges, _tensor=U_tensor) S = GrassmannTensor( _arrow=( False, @@ -709,52 +762,141 @@ def svd( _edges=S_edges, _tensor=torch.diag(S_tensor), ) - Vh = GrassmannTensor(_arrow=(False, True), _edges=Vh_edges, _tensor=Vh_tensor) - # Split - left_arrow = [self.arrow[i] for i in left_legs] - left_edges = [self.edges[i] for i in left_legs] + Vh = GrassmannTensor(_arrow=(False, False), _edges=Vh_edges, _tensor=Vh_tensor) - right_arrow = [self.arrow[i] for i in right_legs] + left_edges = [self.edges[i] for i in left_legs] right_edges = [self.edges[i] for i in right_legs] U = U.reshape((*left_edges, U_edges[1])) - U._arrow = tuple(left_arrow + [True]) + U = U.reverse(arrow_reverse_left) Vh = Vh.reshape((Vh_edges[0], *right_edges)) - Vh._arrow = tuple([False] + right_arrow) + Vh = Vh.reverse(arrow_reverse_right) return U, S, Vh - def _get_inv_order(self, order: tuple[int, ...]) -> tuple[int, ...]: - inv = [0] * self.tensor.dim() + @staticmethod + def get_inv_order(dim: int, order: tuple[int, ...]) -> tuple[int, ...]: + inv = [0] * dim for new_position, origin_idx in enumerate(order): inv[origin_idx] = new_position return tuple(inv) - def _check_pairs_coverage(self, pairs: tuple[tuple[int, ...], tuple[int, ...]]) -> bool: - set0 = set(pairs[0]) - set1 = set(pairs[1]) + def _get_inv_order(self, order: tuple[int, ...]) -> tuple[int, ...]: + return self.get_inv_order(self.tensor.dim(), order) + + @staticmethod + def contract( + a: GrassmannTensor, + b: GrassmannTensor, + a_leg: int | tuple[int, ...], + b_leg: int | tuple[int, ...], + ) -> GrassmannTensor: + contract_lengths = [] + for leg in (a_leg, b_leg): + if isinstance(leg, int): + contract_lengths.append(1) + else: + contract_lengths.append(len(leg)) + assert all(a.arrow[i] == a.arrow[leg[0]] for i in leg), ( + "All the legs that need to be contracted must have the same arrow" + ) + + contract_length_a, contract_length_b = contract_lengths + + a_leg_tuple = (a_leg,) if isinstance(a_leg, int) else a_leg + b_leg_tuple = (b_leg,) if isinstance(b_leg, int) else b_leg - are_disjoint = set0.isdisjoint(set1) + a_range_list = tuple(range(a.tensor.dim())) + b_range_list = tuple(range(b.tensor.dim())) - is_complete_union = (set0 | set1) == set(range(self.tensor.dim())) + a_contract_set = set(a_leg_tuple) + b_contract_set = set(b_leg_tuple) - return are_disjoint and is_complete_union + order_a = tuple(i for i in a_range_list if i not in a_contract_set) + a_leg_tuple + order_b = b_leg_tuple + tuple(i for i in b_range_list if i not in b_contract_set) + + tensor_a = a.permute(order_a) + tensor_b = b.permute(order_b) + + assert (tensor_a.arrow[-1], tensor_b.arrow[0]) in ((False, True), (True, False)), ( + f"Contract requires arrow (False, True) or (True, False), but got {tensor_a.arrow[-1], tensor_b.arrow[0]}" + ) + + arrow_after_permute_a = tensor_a.arrow + arrow_after_permute_b = tensor_b.arrow + + edge_after_permute_a = tensor_a.edges + edge_after_permute_b = tensor_b.edges + + arrow_expected_a = [i >= a.tensor.dim() - contract_length_a for i in range(a.tensor.dim())] + arrow_expected_b = [i >= contract_length_b for i in range(b.tensor.dim())] + + arrow_reverse_a = tuple( + i + for i, (cur, exp) in enumerate(zip(arrow_after_permute_a, arrow_expected_a)) + if cur != exp + ) + arrow_reverse_b = tuple( + i + for i, (cur, exp) in enumerate(zip(arrow_after_permute_b, arrow_expected_b)) + if cur != exp + ) + + if arrow_reverse_a: + tensor_a = ( + tensor_a.reverse(arrow_reverse_a).reverse(arrow_reverse_a).reverse(arrow_reverse_a) + ) + if arrow_reverse_b: + tensor_b = ( + tensor_b.reverse(arrow_reverse_b).reverse(arrow_reverse_b).reverse(arrow_reverse_b) + ) + + tensor_a = tensor_a.reshape( + ( + math.prod(tensor_a.tensor.shape[:-contract_length_a]), + math.prod(tensor_a.tensor.shape[-contract_length_a:]), + ) + ) + tensor_b = tensor_b.reshape( + ( + math.prod(tensor_b.tensor.shape[:contract_length_b]), + math.prod(tensor_b.tensor.shape[contract_length_b:]), + ) + ) + + c = tensor_a @ tensor_b + + c = c.reshape( + (edge_after_permute_a[:-contract_length_a] + edge_after_permute_b[contract_length_b:]) + ) + + arrow_reverse_c = tuple( + [i for i in arrow_reverse_a if i < a.tensor.dim() - contract_length_a] + + [ + (a.tensor.dim() - contract_length_a) + (i - contract_length_b) + for i in arrow_reverse_b + if i >= contract_length_b + ] + ) + c = c.reverse(arrow_reverse_c) + return c def exponential(self, pairs: tuple[tuple[int, ...], tuple[int, ...]]) -> GrassmannTensor: tensor, left_legs, right_legs = self._group_edges(pairs) - arrow_order = (False, True) - edges_to_reverse = tuple( - i for i, arrow in enumerate(arrow_order) if tensor.arrow[i] != arrow + assert tensor.arrow in ((False, True), (True, False)), ( + f"Exponentiation requires arrow (False, True) or (True, False), but got {tensor.arrow}" ) - if edges_to_reverse: - tensor = tensor.reverse(edges_to_reverse) + + tensor_reverse_flag = tensor.arrow != (False, True) + if tensor_reverse_flag: + tensor = tensor.reverse((0, 1)) left_dim, right_dim = tensor.tensor.shape assert left_dim == right_dim, ( - f"Exponential requires a square operator, but got {left_dim} x {right_dim}." + f"Exponentiation requires a square operator, but got {left_dim} x {right_dim}." ) (even_left, odd_left) = tensor.edges[0] @@ -774,8 +916,8 @@ def exponential(self, pairs: tuple[tuple[int, ...], tuple[int, ...]]) -> Grassma tensor_exp = dataclasses.replace(tensor, _tensor=tensor_exp) - if edges_to_reverse: - tensor_exp = tensor_exp.reverse(tuple(edges_to_reverse)) + if tensor_reverse_flag: + tensor_exp = tensor_exp.reverse((0, 1)) order = left_legs + right_legs edges_after_permute = tuple(self.edges[i] for i in order) @@ -787,6 +929,47 @@ def exponential(self, pairs: tuple[tuple[int, ...], tuple[int, ...]]) -> Grassma return tensor_exp + def identity(self, pairs: tuple[tuple[int, ...], tuple[int, ...]]) -> GrassmannTensor: + tensor, left_legs, right_legs = self._group_edges(pairs) + + assert tensor.arrow in ((False, True), (True, False)), ( + f"Identity requires arrow (False, True) or (True, False), but got {tensor.arrow}" + ) + + tensor_reverse_flag = tensor.arrow != (False, True) + if tensor_reverse_flag: + tensor = tensor.reverse((0, 1)) + + left_dim, right_dim = tensor.tensor.shape + + assert left_dim == right_dim, ( + f"Identity requires a square operator, but got {left_dim} x {right_dim}." + ) + + (even_left, odd_left) = tensor.edges[0] + (even_right, odd_right) = tensor.edges[1] + + assert even_left == even_right and odd_left == odd_right, ( + f"Parity blocks must be square, but got L=({even_left},{odd_left}), R=({even_right},{odd_right})" + ) + + I = torch.eye(left_dim, dtype=tensor.tensor.dtype, device=tensor.tensor.device) # noqa: E741 + + tensor_identity = dataclasses.replace(tensor, _tensor=I) + + if tensor_reverse_flag: + tensor_identity = tensor_identity.reverse((0, 1)) + + order = left_legs + right_legs + edges_after_permute = tuple(self.edges[i] for i in order) + tensor_identity = tensor_identity.reshape(edges_after_permute) + + inv_order = self._get_inv_order(order) + + tensor_identity = tensor_identity.permute(inv_order) + + return tensor_identity + def __post_init__(self) -> None: assert len(self._arrow) == self._tensor.dim(), ( f"Arrow length ({len(self._arrow)}) must match tensor dimensions ({self._tensor.dim()})." diff --git a/tests/contract_test.py b/tests/contract_test.py new file mode 100644 index 0000000..501b3de --- /dev/null +++ b/tests/contract_test.py @@ -0,0 +1,34 @@ +import torch +import pytest + +from grassmann_tensor import GrassmannTensor + + +def test_contract() -> None: + a = GrassmannTensor( + (False, False, False, True), + ((2, 2), (4, 4), (8, 8), (32, 32)), + torch.randn(4, 8, 16, 64, dtype=torch.float64), + ) + b = GrassmannTensor( + (False, True, True, True), + ((2, 2), (4, 4), (4, 4), (32, 32)), + torch.randn(4, 8, 8, 64, dtype=torch.float64), + ) + _ = GrassmannTensor.contract(a, b, (0, 2), 3) + _ = GrassmannTensor.contract(a, b, (0, 2), (1, 2)) + + +def test_contract_assertion() -> None: + a = GrassmannTensor((False, True), ((1, 0), (1, 0)), torch.randn(1, 1, dtype=torch.float64)) + b = GrassmannTensor( + (False, True, False, True), + ((2, 2), (4, 4), (8, 8), (16, 16)), + torch.randn(4, 8, 16, 32, dtype=torch.float64), + ) + with pytest.raises(AssertionError, match="Contract requires arrow"): + _ = a.contract(a, b, 0, 0) + with pytest.raises( + AssertionError, match="All the legs that need to be contracted must have the same arrow" + ): + _ = a.contract(a, b, 0, (0, 1)) diff --git a/tests/exponential_test.py b/tests/exponential_test.py index 438cba8..9735115 100644 --- a/tests/exponential_test.py +++ b/tests/exponential_test.py @@ -1,24 +1,17 @@ import torch import pytest +from typing import TypeAlias from grassmann_tensor import GrassmannTensor - -def test_exponential() -> None: - a = GrassmannTensor( - (True, True, True, True), - ((4, 4), (8, 8), (4, 4), (8, 8)), - torch.randn(8, 16, 8, 16, dtype=torch.float64), - ) - b = a.exponential(((0, 3), (1, 2))) - c = a.exponential(((0, 3), (2, 1))) - assert not torch.allclose(b.tensor, c.tensor) +Tensor: TypeAlias = GrassmannTensor +Pairs: TypeAlias = tuple[tuple[int, ...], tuple[int, ...]] def test_exponential_with_empty_parity_block() -> None: - a = GrassmannTensor((False, True), ((1, 0), (1, 0)), torch.randn(1, 1)) + a = GrassmannTensor((False, True), ((1, 0), (1, 0)), torch.randn(1, 1, dtype=torch.float64)) a.exponential(((0,), (1,))) - b = GrassmannTensor((False, True), ((0, 1), (0, 1)), torch.randn(1, 1)) + b = GrassmannTensor((False, True), ((0, 1), (0, 1)), torch.randn(1, 1, dtype=torch.float64)) b.exponential(((0,), (1,))) @@ -28,5 +21,79 @@ def test_exponential_assertation() -> None: ((2, 2), (4, 4), (8, 8), (16, 16)), torch.randn(4, 8, 16, 32, dtype=torch.float64), ) - with pytest.raises(AssertionError, match="Exponential requires a square operator"): + with pytest.raises(AssertionError, match="Exponentiation requires arrow"): a.exponential(((0, 2), (1, 3))) + + b = GrassmannTensor( + (False, True, False, True), + ((2, 2), (4, 4), (8, 8), (16, 16)), + torch.randn(4, 8, 16, 32, dtype=torch.float64), + ) + with pytest.raises(AssertionError, match="Exponentiation requires a square operator"): + b.exponential(((0, 2), (1, 3))) + + c = GrassmannTensor( + (False, True, False, True), + ((1, 3), (3, 1), (3, 1), (3, 1)), + torch.randn(4, 4, 4, 4, dtype=torch.float64), + ) + with pytest.raises(AssertionError, match="Parity blocks must be square"): + c.exponential(((0, 2), (1, 3))) + + +@pytest.mark.parametrize( + "tensor, pairs", + [ + ( + GrassmannTensor( + (False, True), ((4, 4), (4, 4)), torch.randn(8, 8, dtype=torch.float64) + ), + ((0,), (1,)), + ), + ( + GrassmannTensor( + (True, False), ((4, 4), (4, 4)), torch.randn(8, 8, dtype=torch.float64) + ), + ((0,), (1,)), + ), + ( + GrassmannTensor( + (False, False, True), + ((4, 4), (4, 4), (32, 32)), + torch.randn(8, 8, 64, dtype=torch.float64), + ), + ((0, 1), (2,)), + ), + ( + GrassmannTensor( + (False, False, True, True), + ((4, 4), (8, 8), (4, 4), (8, 8)), + torch.randn(8, 16, 8, 16, dtype=torch.float64), + ), + ((0, 1), (2, 3)), + ), + ], +) +def test_exponential_via_taylor_expansion( + tensor: Tensor, + pairs: Pairs, +) -> None: + tensor_exp = tensor.exponential(pairs) + iter_tensor = tensor.identity(pairs) + iter_tensor, _, _ = iter_tensor._group_edges(pairs) + iter_tensor = iter_tensor.update_mask() + tensor_group_edges, left_legs, right_legs = tensor._group_edges(pairs) + tensor_group_edges = tensor_group_edges.update_mask() + + tensor_taylor_expansion = iter_tensor + for i in range(1, 50): + iter_tensor @= tensor_group_edges / i + tensor_taylor_expansion += iter_tensor + + order = left_legs + right_legs + edges_after_permute = tuple(tensor.edges[i] for i in order) + tensor_taylor_expansion = tensor_taylor_expansion.reshape(edges_after_permute) + inv_order = tensor._get_inv_order(order) + tensor_taylor_expansion = tensor_taylor_expansion.permute(inv_order) + + assert torch.allclose(tensor_taylor_expansion.tensor, tensor_exp.tensor) diff --git a/tests/identity_test.py b/tests/identity_test.py new file mode 100644 index 0000000..b1699d6 --- /dev/null +++ b/tests/identity_test.py @@ -0,0 +1,83 @@ +import pytest +import torch +from typing import TypeAlias + +from grassmann_tensor import GrassmannTensor + +Tensor: TypeAlias = GrassmannTensor +Pairs: TypeAlias = tuple[tuple[int, ...], tuple[int, ...]] + + +def test_identity_assertation() -> None: + a = GrassmannTensor( + (True, True, True, True), + ((2, 2), (4, 4), (8, 8), (16, 16)), + torch.randn(4, 8, 16, 32, dtype=torch.float64), + ) + with pytest.raises(AssertionError, match="Identity requires arrow"): + a.identity(((0, 2), (1, 3))) + + b = GrassmannTensor( + (False, True, False, True), + ((2, 2), (4, 4), (8, 8), (16, 16)), + torch.randn(4, 8, 16, 32, dtype=torch.float64), + ) + with pytest.raises(AssertionError, match="Identity requires a square operator"): + b.identity(((0, 2), (1, 3))) + + c = GrassmannTensor( + (False, True, False, True), + ((1, 3), (3, 1), (3, 1), (3, 1)), + torch.randn(4, 4, 4, 4, dtype=torch.float64), + ) + with pytest.raises(AssertionError, match="Parity blocks must be square"): + c.identity(((0, 2), (1, 3))) + + +@pytest.mark.parametrize( + "tensor, pairs", + [ + ( + GrassmannTensor( + (False, True), ((4, 4), (4, 4)), torch.randn(8, 8, dtype=torch.float64) + ), + ((0,), (1,)), + ), + ( + GrassmannTensor( + (True, False), ((4, 4), (4, 4)), torch.randn(8, 8, dtype=torch.float64) + ), + ((0,), (1,)), + ), + ( + GrassmannTensor( + (False, False, True), + ((4, 4), (4, 4), (32, 32)), + torch.randn(8, 8, 64, dtype=torch.float64), + ), + ((0, 1), (2,)), + ), + ( + GrassmannTensor( + (False, False, True, True), + ((4, 4), (8, 8), (4, 4), (8, 8)), + torch.randn(8, 16, 8, 16, dtype=torch.float64), + ), + ((0, 1), (2, 3)), + ), + ], +) +def test_identity_via_self_multiplication( + tensor: Tensor, + pairs: Pairs, +) -> None: + identity = tensor.identity(pairs) + identity, _, _ = identity._group_edges(pairs) + tensor, _, _ = tensor._group_edges(pairs) + tensor_reverse_flag = tensor.arrow != (False, True) + if tensor_reverse_flag: + identity = identity.reverse((0, 1)) + tensor = tensor.reverse((0, 1)) + assert torch.allclose((identity @ identity).tensor, identity.tensor) + assert torch.allclose((identity @ tensor).tensor, tensor.tensor) + assert torch.allclose((tensor @ identity).tensor, tensor.tensor) diff --git a/tests/reshape_test.py b/tests/reshape_test.py index 2b3b06c..77f1823 100644 --- a/tests/reshape_test.py +++ b/tests/reshape_test.py @@ -225,6 +225,8 @@ def test_reshape_with_one_dimension( assert ( len(a.arrow) == len(shape) and len(a.edges) == len(shape) and a.tensor.dim() == len(shape) ) + if len(shape) > len(arrow): + assert all(not a.arrow[i] for i in range(len(arrow), len(shape))) def test_reshape_trailing_nontrivial_dim_raises() -> None: diff --git a/tests/svd_test.py b/tests/svd_test.py index 4660385..dcff3b9 100644 --- a/tests/svd_test.py +++ b/tests/svd_test.py @@ -1,7 +1,6 @@ import torch import pytest from _pytest.mark.structures import ParameterSet -import math import itertools from typing import TypeAlias, Iterable, Any @@ -43,6 +42,7 @@ def choose_free_names(n_edges: int, limit: int = 8) -> list[FreeNamesU]: BASE_GT_CASES: list[tuple[Arrow, Edges, Tensor]] = [ ((True, True), ((2, 2), (4, 4)), torch.randn(4, 8, dtype=torch.float64)), + ((False, False), ((2, 2), (4, 4)), torch.randn(4, 8, dtype=torch.float64)), ((True, True, True), ((2, 2), (4, 4), (8, 8)), torch.randn(4, 8, 16, dtype=torch.float64)), ( (True, True, True, True), @@ -106,27 +106,13 @@ def test_svd( gt = GrassmannTensor(arrow, edges, tensor) U, S, Vh = gt.svd(free_names_u, cutoff=cutoff) - # reshape U - left_dim = math.prod(U.tensor.shape[:-1]) - left_edge = list(U.edges[:-1]) - U = U.reshape((left_dim, -1)) + US = GrassmannTensor.contract(U, S, U.tensor.dim() - 1, 0) + USV = GrassmannTensor.contract(US, Vh, US.tensor.dim() - 1, 0) - # reshape Vh - right_dim = math.prod(Vh.tensor.shape[1:]) - right_edge = list(Vh.edges[1:]) - Vh = Vh.reshape((-1, right_dim)) - - US = GrassmannTensor.matmul(U, S) - USV = GrassmannTensor.matmul(US, Vh) - - set_all = set(range(len(edges))) - set_u = set(free_names_u) - set_v = sorted(set_all - set_u) - perm_order = list(free_names_u) + list(set_v) - inv_perm = [perm_order.index(i) for i in range(len(edges))] - - USV = USV.reshape(tuple(left_edge + right_edge)) - USV = USV.permute(tuple(inv_perm)) + left_legs, right_legs = gt.get_legs_pair(len(edges), free_names_u) + order = left_legs + right_legs + inv_order = gt.get_inv_order(USV.tensor.dim(), order) + USV = USV.permute(inv_order) masked = gt.update_mask().tensor den = masked.norm()