From 24f8303055c933bf6bfca3e0e6c6c176ce9d08e4 Mon Sep 17 00:00:00 2001 From: Gausshj Date: Fri, 21 Nov 2025 17:35:33 +0800 Subject: [PATCH] refactor(get-inv-order): Refactor get inverse order method - Refactor get inverse order method, remove redundant parameters - Modify the affected parts --- grassmann_tensor/tensor.py | 11 ++++------- tests/exponential_test.py | 2 +- tests/svd_test.py | 2 +- 3 files changed, 6 insertions(+), 9 deletions(-) diff --git a/grassmann_tensor/tensor.py b/grassmann_tensor/tensor.py index 47a75f9..19da2c7 100644 --- a/grassmann_tensor/tensor.py +++ b/grassmann_tensor/tensor.py @@ -777,15 +777,12 @@ def svd( return U, S, Vh @staticmethod - def get_inv_order(dim: int, order: tuple[int, ...]) -> tuple[int, ...]: - inv = [0] * dim + def get_inv_order(order: tuple[int, ...]) -> tuple[int, ...]: + inv = [0] * len(order) for new_position, origin_idx in enumerate(order): inv[origin_idx] = new_position return tuple(inv) - def _get_inv_order(self, order: tuple[int, ...]) -> tuple[int, ...]: - return self.get_inv_order(self.tensor.dim(), order) - def contract( self, b: GrassmannTensor, @@ -925,7 +922,7 @@ def exponential(self, pairs: tuple[tuple[int, ...], tuple[int, ...]]) -> Grassma edges_after_permute = tuple(self.edges[i] for i in order) tensor_exp = tensor_exp.reshape(edges_after_permute) - inv_order = self._get_inv_order(order) + inv_order = self.get_inv_order(order) tensor_exp = tensor_exp.permute(inv_order) @@ -966,7 +963,7 @@ def identity(self, pairs: tuple[tuple[int, ...], tuple[int, ...]]) -> GrassmannT 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) + inv_order = self.get_inv_order(order) tensor_identity = tensor_identity.permute(inv_order) diff --git a/tests/exponential_test.py b/tests/exponential_test.py index 9735115..1f24a46 100644 --- a/tests/exponential_test.py +++ b/tests/exponential_test.py @@ -93,7 +93,7 @@ def test_exponential_via_taylor_expansion( 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) + 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/svd_test.py b/tests/svd_test.py index 174d8ef..0e3da3b 100644 --- a/tests/svd_test.py +++ b/tests/svd_test.py @@ -111,7 +111,7 @@ def test_svd( 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) + inv_order = gt.get_inv_order(order) USV = USV.permute(inv_order) masked = gt.update_mask().tensor