diff --git a/grassmann_tensor/tensor.py b/grassmann_tensor/tensor.py index ee4cef3..7b9dbb7 100644 --- a/grassmann_tensor/tensor.py +++ b/grassmann_tensor/tensor.py @@ -793,12 +793,12 @@ def contract( b_leg: int | tuple[int, ...], ) -> GrassmannTensor: contract_lengths = [] - for leg in (a_leg, b_leg): + for leg, tensor in ((a_leg, a), (b_leg, b)): 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), ( + assert all(tensor.arrow[i] == tensor.arrow[leg[0]] for i in leg), ( "All the legs that need to be contracted must have the same arrow" ) diff --git a/tests/contract_test.py b/tests/contract_test.py index 501b3de..6c494d8 100644 --- a/tests/contract_test.py +++ b/tests/contract_test.py @@ -7,14 +7,15 @@ 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), + ((2, 2), (4, 4), (8, 8), (8, 8)), + torch.randn(4, 8, 16, 16, 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), + ((8, 8), (4, 4), (4, 4), (32, 32)), + torch.randn(16, 8, 8, 64, dtype=torch.float64), ) + _ = GrassmannTensor.contract(a, b, 3, 0) _ = GrassmannTensor.contract(a, b, (0, 2), 3) _ = GrassmannTensor.contract(a, b, (0, 2), (1, 2))