From 3dec35a91e2e84f4a2399d08b4c99ff309566228 Mon Sep 17 00:00:00 2001 From: Baruch Tabanpour Date: Thu, 21 Aug 2025 10:43:09 -0700 Subject: [PATCH] Do not update `qacc_warmstart` at the end of the solver call; instead, update it at the same time as all other state variables. This change makes `mj_forward` idempotent. PiperOrigin-RevId: 797826112 Change-Id: Ibc51624adc3ec42c265958bf77696b33231431c1 --- mjx/mujoco/mjx/_src/forward.py | 3 +++ mjx/mujoco/mjx/_src/solver.py | 1 - 2 files changed, 3 insertions(+), 1 deletion(-) diff --git a/mjx/mujoco/mjx/_src/forward.py b/mjx/mujoco/mjx/_src/forward.py index bb01727d..2a0e946e 100644 --- a/mjx/mujoco/mjx/_src/forward.py +++ b/mjx/mujoco/mjx/_src/forward.py @@ -334,6 +334,9 @@ def _advance( # advance time time = d.time + m.opt.timestep + # save qacc for next step warmstart + d = d.replace(qacc_warmstart=d.qacc) + return d.replace(act=act, qpos=qpos, time=time) diff --git a/mjx/mujoco/mjx/_src/solver.py b/mjx/mujoco/mjx/_src/solver.py index f57dbd1c..6ac37f19 100644 --- a/mjx/mujoco/mjx/_src/solver.py +++ b/mjx/mujoco/mjx/_src/solver.py @@ -602,7 +602,6 @@ def solve(m: Model, d: Data) -> Data: ctx = jax.lax.while_loop(cond, body, ctx) d = d.tree_replace({ - 'qacc_warmstart': ctx.qacc, 'qfrc_constraint': ctx.qfrc_constraint, 'qacc': ctx.qacc, '_impl.efc_force': ctx.efc_force,