Skip to content

Commit 5ff12c8

Browse files
jessegrabowskiricardoV94
authored andcommitted
Remove nan guards from cholesky/solve_triangular
1 parent 5d26f8b commit 5ff12c8

1 file changed

Lines changed: 5 additions & 18 deletions

File tree

pymc/distributions/multivariate.py

Lines changed: 5 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -36,7 +36,7 @@
3636
)
3737
from pytensor.tensor.elemwise import DimShuffle
3838
from 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
4040
from pytensor.tensor.linalg import inv as matrix_inverse
4141
from pytensor.tensor.random import chisquare
4242
from 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-
134127
def 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

Comments
 (0)