From dd63bfd2bf87b6d97e57620d89a3ec8727a48c00 Mon Sep 17 00:00:00 2001 From: Hao Zhang Date: Fri, 8 Aug 2025 16:40:38 +0800 Subject: [PATCH 1/2] Add clone function. --- grassmann_tensor/tensor.py | 17 +++++++++++++++++ 1 file changed, 17 insertions(+) diff --git a/grassmann_tensor/tensor.py b/grassmann_tensor/tensor.py index 30ec73e..9c5d109 100644 --- a/grassmann_tensor/tensor.py +++ b/grassmann_tensor/tensor.py @@ -358,3 +358,20 @@ def __itruediv__(self, other: typing.Any) -> GrassmannTensor: else: self._tensor /= other return self + + def clone(self) -> GrassmannTensor: + """ + Create a deep copy of the Grassmann tensor. + """ + return dataclasses.replace( + self, + _tensor=self._tensor.clone(), + _parity=tuple(parity.clone() for parity in self._parity) if self._parity is not None else None, + _mask=self._mask.clone() if self._mask is not None else None, + ) + + def __copy__(self) -> GrassmannTensor: + return self.clone() + + def __deepcopy__(self, memo: dict) -> GrassmannTensor: + return self.clone() From 40862f068477cb116971bc5f1c73569deb5dda4d Mon Sep 17 00:00:00 2001 From: Hao Zhang Date: Fri, 8 Aug 2025 16:56:02 +0800 Subject: [PATCH 2/2] Add tests for clone. --- tests/clone_test.py | 50 +++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 50 insertions(+) create mode 100644 tests/clone_test.py diff --git a/tests/clone_test.py b/tests/clone_test.py new file mode 100644 index 0000000..6a67bcc --- /dev/null +++ b/tests/clone_test.py @@ -0,0 +1,50 @@ +import typing +import copy +import pytest +import torch +from grassmann_tensor import GrassmannTensor + + +@pytest.mark.parametrize("parity,mask", [(True, True), (True, False), (False, False)]) +@pytest.mark.parametrize("which", ["clone", "copy", "deepcopy"]) +def test_clone( + parity: bool, + mask: bool, + which: typing.Literal["clone", "copy", "deepcopy"], +) -> None: + original_tensor = GrassmannTensor( + _arrow=(False, True), + _edges=((2, 2), (1, 3)), + _tensor=torch.randn([4, 4]), + ) + + if parity: + _ = original_tensor.parity + if mask: + _ = original_tensor.mask + + match which: + case "clone": + cloned_tensor = original_tensor.clone() + case "copy": + cloned_tensor = copy.copy(original_tensor) + case "deepcopy": + cloned_tensor = copy.deepcopy(original_tensor) + + assert cloned_tensor._arrow == original_tensor._arrow + assert cloned_tensor._edges == original_tensor._edges + assert torch.equal(cloned_tensor._tensor, original_tensor._tensor) + if parity: + assert cloned_tensor._parity is not None + assert original_tensor._parity is not None + assert all(torch.equal(c, o) for c, o in zip(cloned_tensor._parity, original_tensor._parity)) + else: + assert cloned_tensor._parity is original_tensor._parity + if mask: + assert cloned_tensor._mask is not None + assert original_tensor._mask is not None + assert torch.equal(cloned_tensor._mask, original_tensor._mask) + else: + assert cloned_tensor._mask is original_tensor._mask + + assert id(original_tensor.tensor) != id(cloned_tensor.tensor)