From a3ad4bf180cb76d1a8242ae2b8f0348e0a72d460 Mon Sep 17 00:00:00 2001 From: N-0-MAD Date: Tue, 15 Sep 2026 03:25:48 +0530 Subject: [PATCH] Fix PSD solve rewrite for vector RHS --- pytensor/tensor/rewriting/linalg/solvers.py | 9 ++++++--- tests/tensor/rewriting/linalg/test_solvers.py | 9 +++++---- 2 files changed, 11 insertions(+), 7 deletions(-) diff --git a/pytensor/tensor/rewriting/linalg/solvers.py b/pytensor/tensor/rewriting/linalg/solvers.py index f3439bd3c6..0db3a0627f 100644 --- a/pytensor/tensor/rewriting/linalg/solvers.py +++ b/pytensor/tensor/rewriting/linalg/solvers.py @@ -111,7 +111,9 @@ def batched_vector_b_solve_to_matrix_b_solve(fgraph, node): @register_stabilize -@node_rewriter([blockwise_of(OpPattern(Solve, b_ndim=2))]) +@node_rewriter( + [blockwise_of(OpPattern(Solve, b_ndim=1)), blockwise_of(OpPattern(Solve, b_ndim=2))] +) def psd_solve_to_chol_solve(fgraph, node): """Rewrite solve(A, b) → triangular solves via Cholesky when A is positive-definite.""" assume_a = node.op.core_op.assume_a @@ -121,9 +123,10 @@ def psd_solve_to_chol_solve(fgraph, node): or getattr(A.tag, "psd", None) is True or check_assumption(fgraph, A, POSITIVE_DEFINITE) ): + b_ndim = node.op.core_op.b_ndim L = cholesky(A) - Li_b = solve_triangular(L, b, lower=True, b_ndim=2) - x = solve_triangular((L.mT), Li_b, lower=False, b_ndim=2) + Li_b = solve_triangular(L, b, lower=True, b_ndim=b_ndim) + x = solve_triangular((L.mT), Li_b, lower=False, b_ndim=b_ndim) return [x] diff --git a/tests/tensor/rewriting/linalg/test_solvers.py b/tests/tensor/rewriting/linalg/test_solvers.py index 3dd70ae394..0ffc33354d 100644 --- a/tests/tensor/rewriting/linalg/test_solvers.py +++ b/tests/tensor/rewriting/linalg/test_solvers.py @@ -77,17 +77,18 @@ def test_generic_solve_to_solve_triangular(): ) -def test_psd_solve_with_chol(): +@pytest.mark.parametrize("b_ndim", [1, 2], ids=lambda x: f"b_ndim={x}") +def test_psd_solve_with_chol(b_ndim): """Test that solve(A, b) with PSD A gets rewritten to cholesky + cho_solve.""" A = matrix("A") - b = matrix("b") + b = pt.vector("b") if b_ndim == 1 else matrix("b") A_psd = assume(A, positive_definite=True) - out = pt.linalg.solve(A_psd, b) + out = pt.linalg.solve(A_psd, b, b_ndim=b_ndim) rewritten = rewrite_graph(out, include=("canonicalize", "stabilize", "specialize")) L = cholesky(A_psd) - expected = cho_solve((L, True), b, b_ndim=2) + expected = cho_solve((L, True), b, b_ndim=b_ndim) assert_equal_computations([rewritten], [expected])