Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 17 additions & 0 deletions grassmann_tensor/tensor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Comment thread
hzhangxyz marked this conversation as resolved.
50 changes: 50 additions & 0 deletions tests/clone_test.py
Original file line number Diff line number Diff line change
@@ -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)