Symmetrize dense matrices before cho_factor, since the jax.scipy.linalg.cho_factor API had a breaking change in https://github.com/jax-ml/jax/commit/a981e1c4b99b7628efe2bf70a75c0422a70e9524.
PiperOrigin-RevId: 775740310 Change-Id: I1a172ec4851b5d6a7f3d56179c54fec5d4c9e0c5
This commit is contained in:
committed by
Copybara-Service
parent
18d5d5d03f
commit
48205fb570
@@ -403,7 +403,9 @@ def _update_gradient(m: Model, d: Data, ctx: Context) -> Context:
|
||||
else:
|
||||
h = (d._impl.efc_J.T * d._impl.efc_D * ctx.active) @ d._impl.efc_J
|
||||
h = support.full_m(m, d) + h
|
||||
h_ = jax.scipy.linalg.cho_factor(h)
|
||||
# Symmetrize to reduce the chance of numerical issues in cholesky.
|
||||
h_sym = (h + h.T) * 0.5
|
||||
h_ = jax.scipy.linalg.cho_factor(h_sym)
|
||||
mgrad = jax.scipy.linalg.cho_solve(h_, grad)
|
||||
else:
|
||||
raise NotImplementedError(f'unsupported solver type: {m.opt.solver}')
|
||||
|
||||
Reference in New Issue
Block a user