From 06b14b070b136d8ae36335aedd57b42c8cfa426e Mon Sep 17 00:00:00 2001 From: Seunghoon Lee Date: Sun, 28 Sep 2025 18:43:37 +0900 Subject: [PATCH] hijack torch.linalg.cholesky() --- modules/rocm.py | 11 ++++++++++- 1 file changed, 10 insertions(+), 1 deletion(-) diff --git a/modules/rocm.py b/modules/rocm.py index 9728313e0..0d489b72a 100644 --- a/modules/rocm.py +++ b/modules/rocm.py @@ -245,7 +245,7 @@ if sys.platform == "win32": cholesky_ex_gpu = torch.linalg.cholesky_ex @wraps(cholesky_ex_gpu) - def cholesky_ex(A: torch.Tensor, upper=False, check_errors=False, out=None) -> torch.Tensor: + def cholesky_ex(A: torch.Tensor, upper=False, check_errors=False, out=None) -> torch.return_types.linalg_cholesky_ex: if A.device.type != 'cpu': return cholesky_ex_gpu(A, upper=upper, check_errors=check_errors, out=out) @@ -269,6 +269,15 @@ if sys.platform == "win32": return torch.return_types.linalg_cholesky_ex((L, torch.tensor(0, dtype=torch.int32, device='cpu')), {}) torch.linalg.cholesky_ex = cholesky_ex + cholesky_gpu = torch.linalg.cholesky + @wraps(cholesky_gpu) + def cholesky(A: torch.Tensor, upper=False, out=None) -> torch.Tensor: + if A.device.type != 'cpu': + return cholesky_gpu(A, upper=upper, out=out) + L, _ = torch.linalg.cholesky_ex(A, upper=upper, out=out) + return L + torch.linalg.cholesky = cholesky + is_wsl: bool = False else: def get_agents() -> List[Agent]: