From dd6cbd8bd996e9ad583af3927b156ac6bef83d6d Mon Sep 17 00:00:00 2001 From: Gausshj Date: Fri, 31 Oct 2025 16:30:40 +0800 Subject: [PATCH 1/2] fix(reshape): fix reshape with head int1 as trivial - Fix the issue when reshape with head int 1 - Add test cases --- grassmann_tensor/tensor.py | 17 +++++++++++++++-- tests/reshape_test.py | 32 ++++++++++++++++++++++++++++++++ 2 files changed, 47 insertions(+), 2 deletions(-) diff --git a/grassmann_tensor/tensor.py b/grassmann_tensor/tensor.py index e4de632..ddc9525 100644 --- a/grassmann_tensor/tensor.py +++ b/grassmann_tensor/tensor.py @@ -297,8 +297,8 @@ def reshape(self, new_shape: tuple[int | tuple[int, int], ...]) -> GrassmannTens return GrassmannTensor(_arrow=(), _edges=(), _tensor=tensor) if new_shape == (1,) and int(self.tensor.numel()) == 1: - eo = self._calculate_even_odd() - new_shape = (eo,) + even_self, odd_self = self._calculate_even_odd() + new_shape = ((even_self, odd_self),) cursor_plan: int = 0 cursor_self: int = 0 @@ -318,6 +318,19 @@ def reshape(self, new_shape: tuple[int | tuple[int, int], ...]) -> GrassmannTens f"edges={self.edges}, new_shape={new_shape}" ) + if cursor_plan != len(new_shape): + new_shape_check = new_shape[cursor_plan] + if ( + isinstance(new_shape_check, int) + and new_shape_check == 1 + and self.tensor.shape[cursor_self] != 1 + ): + arrow.append(False) + edges.append((1, 0)) + shape.append(1) + cursor_plan += 1 + continue + if cursor_plan != len(new_shape) and new_shape[cursor_plan] == -1: # Does not change arrow.append(self.arrow[cursor_self]) diff --git a/tests/reshape_test.py b/tests/reshape_test.py index 2d2d452..e9bbe91 100644 --- a/tests/reshape_test.py +++ b/tests/reshape_test.py @@ -231,3 +231,35 @@ def test_reshape_trailing_nontrivial_dim_raises() -> None: a = GrassmannTensor((True,), ((2, 2),), torch.randn([4])) with pytest.raises(AssertionError, match="New shape exceeds after exhausting self dimensions"): _ = a.reshape((-1, (2, 2))) + + +@pytest.mark.parametrize( + "tensor", + [ + GrassmannTensor( + (True, True, True, True), + ((1, 0), (1, 0), (2, 2), (8, 8)), + torch.randn(1, 1, 4, 16), + ), + ], +) +@pytest.mark.parametrize( + "shape", + [ + (1, 64), + ((1, 0), 64), + (-1, 64), + ], +) +def test_reshape_trivial_head_equivalence( + tensor: GrassmannTensor, + shape: tuple[int, ...], +) -> None: + baseline_tensor = tensor.reshape((1, 64)) + actual_tensor = tensor.reshape(shape) + + assert actual_tensor.edges == ((1, 0), (32, 32)) + assert torch.allclose(actual_tensor.tensor, baseline_tensor.tensor) + + roundtrip_tensor = actual_tensor.reshape(tensor.edges) + assert torch.allclose(roundtrip_tensor.tensor, tensor.tensor) From 1a135dfb534d5103b7fd9c2eb0fb7c617b199a8a Mon Sep 17 00:00:00 2001 From: Gausshj Date: Fri, 31 Oct 2025 16:43:03 +0800 Subject: [PATCH 2/2] fix(reshape): add test cases to cover all codes - Add test cases to cover all codes --- tests/reshape_test.py | 24 ++++++++++++++++++++++++ 1 file changed, 24 insertions(+) diff --git a/tests/reshape_test.py b/tests/reshape_test.py index e9bbe91..2b3b06c 100644 --- a/tests/reshape_test.py +++ b/tests/reshape_test.py @@ -263,3 +263,27 @@ def test_reshape_trivial_head_equivalence( roundtrip_tensor = actual_tensor.reshape(tensor.edges) assert torch.allclose(roundtrip_tensor.tensor, tensor.tensor) + + +def test_reshape_head_1_inserts_trivial_when_self_dim_not_one() -> None: + a = GrassmannTensor( + (True, True), + ((2, 2), (8, 8)), + torch.randn(4, 16), + ) + out = a.reshape((1, 64)) + assert out.edges == ((1, 0), (32, 32)) + assert out.tensor.shape == (1, 64) + assert out.arrow[0] is False + + +def test_reshape_plan_exhausted_then_skip_trivial_self_edges() -> None: + a = GrassmannTensor( + (False, False, False), + ((2, 2), (1, 0), (1, 0)), + torch.randn(4, 1, 1), + ) + out = a.reshape((4,)) + assert out.edges == ((2, 2),) + assert out.tensor.shape == (4,) + assert out.arrow == (False,)