From b1d4f8e8497172b07178221ee91638bcb1e98c6d Mon Sep 17 00:00:00 2001 From: Gausshj Date: Mon, 24 Nov 2025 16:50:38 +0800 Subject: [PATCH] fix(svd): fix svd dtype and device continuity - Fix svd dtype and device continuity issues - Add corresponding test --- grassmann_tensor/tensor.py | 4 ++++ tests/svd_test.py | 28 ++++++++++++++++++++++++++++ 2 files changed, 32 insertions(+) diff --git a/grassmann_tensor/tensor.py b/grassmann_tensor/tensor.py index 19da2c7..18eb773 100644 --- a/grassmann_tensor/tensor.py +++ b/grassmann_tensor/tensor.py @@ -741,6 +741,10 @@ def svd( S_tensor = torch.cat([S_even_trunc, S_odd_trunc], dim=0) Vh_tensor = torch.block_diag(Vh_even_trunc, Vh_odd_trunc) # type: ignore[no-untyped-call] + U_tensor = U_tensor.to(dtype=self.tensor.dtype, device=self.tensor.device) + S_tensor = S_tensor.to(dtype=self.tensor.dtype, device=self.tensor.device) + Vh_tensor = Vh_tensor.to(dtype=self.tensor.dtype, device=self.tensor.device) + U_edges = ( (U_even_trunc.shape[0], U_odd_trunc.shape[0]), (U_even_trunc.shape[1], U_odd_trunc.shape[1]), diff --git a/tests/svd_test.py b/tests/svd_test.py index 0e3da3b..84a2b8e 100644 --- a/tests/svd_test.py +++ b/tests/svd_test.py @@ -237,3 +237,31 @@ def test_svd_int_cutoff_odd_block_empty_select_from_even_only(a: int, b: int, k: assert U.edges[-1] == (expected_k, 0) assert Vh.edges[0] == (expected_k, 0) assert S.edges == ((expected_k, 0), (expected_k, 0)) + + +devices = [torch.device("cpu")] +if torch.cuda.is_available(): + devices.append(torch.device("cuda")) + + +@pytest.mark.parametrize( + "dtype", + [ + torch.float64, + torch.complex128, + ], +) +@pytest.mark.parametrize("device", devices) +def test_svd_dtype_device_continuity(dtype: torch.dtype, device: torch.device) -> None: + a = GrassmannTensor( + (True, True, True, True), + ((2, 2), (4, 4), (8, 8), (16, 16)), + torch.randn(4, 8, 16, 32, dtype=dtype, device=device), + ) + u, s, vh = a.svd((0,), cutoff=1) + assert u.tensor.dtype == dtype + assert s.tensor.dtype == dtype + assert vh.tensor.dtype == dtype + assert u.tensor.device == device + assert s.tensor.device == device + assert vh.tensor.device == device