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:
Baruch Tabanpour
2025-06-25 10:14:03 -07:00
committed by Copybara-Service
parent 18d5d5d03f
commit 48205fb570
+3 -1
View File
@@ -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}')