From 48205fb5709349c727e1c5a1c9697aa0f5f44bef Mon Sep 17 00:00:00 2001 From: Baruch Tabanpour Date: Wed, 25 Jun 2025 10:14:03 -0700 Subject: [PATCH] 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 --- mjx/mujoco/mjx/_src/solver.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/mjx/mujoco/mjx/_src/solver.py b/mjx/mujoco/mjx/_src/solver.py index 98ec3d27..90c1add8 100644 --- a/mjx/mujoco/mjx/_src/solver.py +++ b/mjx/mujoco/mjx/_src/solver.py @@ -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}')