3636)
3737from pytensor .tensor .elemwise import DimShuffle
3838from pytensor .tensor .exceptions import NotScalarConstantError
39- from pytensor .tensor .linalg import cholesky , det , eigh , solve_triangular , trace
39+ from pytensor .tensor .linalg import det , eigh , solve_triangular , trace
4040from pytensor .tensor .linalg import inv as matrix_inverse
4141from pytensor .tensor .random import chisquare
4242from pytensor .tensor .random .basic import MvNormalRV , dirichlet , multinomial , multivariate_normal
@@ -124,13 +124,6 @@ def simplex_cont_transform(op, rv):
124124 return transforms .simplex
125125
126126
127- # Step methods and advi do not catch LinAlgErrors at the
128- # moment. We work around that by using a cholesky op
129- # that returns a nan as first entry instead of raising
130- # an error.
131- nan_lower_cholesky = partial (cholesky , lower = True , on_error = "nan" )
132-
133-
134127def quaddist_matrix (cov = None , chol = None , tau = None , lower = True , * args , ** kwargs ):
135128 if len ([i for i in [tau , cov , chol ] if i is not None ]) != 1 :
136129 raise ValueError ("Incompatible parameterization. Specify exactly one of tau, cov, or chol." )
@@ -176,13 +169,9 @@ def quaddist_chol(value, mu, cov):
176169 else :
177170 onedim = False
178171
179- chol_cov = nan_lower_cholesky (cov )
172+ chol_cov = pt . linalg . cholesky (cov , lower = True )
180173 logdet , posdef = _logdet_from_cholesky (chol_cov )
181174
182- # solve_triangular will raise if there are nans
183- # (which happens if the cholesky fails)
184- chol_cov = pt .switch (posdef [..., None , None ], chol_cov , 1 )
185-
186175 delta = value - mu
187176 delta_trans = solve_lower (chol_cov , delta , b_ndim = 1 )
188177 quaddist = (delta_trans ** 2 ).sum (axis = - 1 )
@@ -347,7 +336,7 @@ def precision_mv_normal_logp(op: PrecisionMvNormalRV, value, rng, size, mean, ta
347336
348337 delta = value - mean
349338 quadratic_form = delta .T @ tau @ delta
350- logdet , posdef = _logdet_from_cholesky (nan_lower_cholesky (tau ))
339+ logdet , posdef = _logdet_from_cholesky (pt . linalg . cholesky (tau , lower = True ))
351340 logp = - 0.5 * (k * pt .log (2 * np .pi ) + quadratic_form ) + logdet
352341
353342 return check_parameters (
@@ -1861,8 +1850,6 @@ def dist(
18611850 * args ,
18621851 ** kwargs ,
18631852 ):
1864- lower_cholesky = partial (cholesky , lower = True , on_error = "raise" )
1865-
18661853 # Among-row matrices
18671854 if len ([i for i in [rowcov , rowchol ] if i is not None ]) != 1 :
18681855 raise ValueError (
@@ -1871,7 +1858,7 @@ def dist(
18711858 if rowcov is not None :
18721859 if rowcov .ndim != 2 :
18731860 raise ValueError ("rowcov must be two dimensional." )
1874- rowchol_cov = lower_cholesky (rowcov )
1861+ rowchol_cov = pt . linalg . cholesky (rowcov , lower = True )
18751862 else :
18761863 if rowchol .ndim != 2 :
18771864 raise ValueError ("rowchol must be two dimensional." )
@@ -1886,7 +1873,7 @@ def dist(
18861873 colcov = pt .as_tensor_variable (colcov )
18871874 if colcov .ndim != 2 :
18881875 raise ValueError ("colcov must be two dimensional." )
1889- colchol_cov = lower_cholesky (colcov )
1876+ colchol_cov = pt . linalg . cholesky (colcov , lower = True )
18901877 else :
18911878 if colchol .ndim != 2 :
18921879 raise ValueError ("colchol must be two dimensional." )
0 commit comments