diff --git a/mjx/mujoco/mjx/_src/collision_driver.py b/mjx/mujoco/mjx/_src/collision_driver.py index 70d3a59b..26d601b3 100644 --- a/mjx/mujoco/mjx/_src/collision_driver.py +++ b/mjx/mujoco/mjx/_src/collision_driver.py @@ -276,7 +276,7 @@ def _contact_groups(m: Model, d: Data) -> Dict[FunctionKey, Contact]: # pair contacts get their params from m.pair_* fields params.append(( m.pair_margin[ip] - m.pair_gap[ip], - jp.clip(m.pair_friction[ip], a_min=eps), + jp.clip(m.pair_friction[ip], min=eps), m.pair_solref[ip], m.pair_solreffriction[ip], m.pair_solimp[ip], diff --git a/mjx/mujoco/mjx/_src/forward.py b/mjx/mujoco/mjx/_src/forward.py index 02afca7a..8e0de506 100644 --- a/mjx/mujoco/mjx/_src/forward.py +++ b/mjx/mujoco/mjx/_src/forward.py @@ -301,7 +301,7 @@ def _next_activation(m: Model, d: Data, act_dot: jax.Array) -> jax.Array: def fn(dyntype, dynprm, act, act_dot, actrange): if dyntype == DynType.FILTEREXACT: - tau = jp.clip(dynprm[0], a_min=mujoco.mjMINVAL) + tau = jp.clip(dynprm[0], min=mujoco.mjMINVAL) act = act + act_dot * tau * (1 - jp.exp(-m.opt.timestep / tau)) else: act = act + act_dot * m.opt.timestep diff --git a/mjx/mujoco/mjx/_src/passive.py b/mjx/mujoco/mjx/_src/passive.py index dc574fbf..ee0aae9a 100644 --- a/mjx/mujoco/mjx/_src/passive.py +++ b/mjx/mujoco/mjx/_src/passive.py @@ -174,7 +174,7 @@ def _inertia_box_fluid_model( box = jp.repeat(inertia[None, :], 3, axis=0) box *= jp.ones((3, 3)) - 2 * jp.eye(3) - box = 6.0 * jp.clip(jp.sum(box, axis=-1), a_min=1e-12) + box = 6.0 * jp.clip(jp.sum(box, axis=-1), min=1e-12) box = jp.sqrt(box / jp.maximum(mass, 1e-12)) * (mass > 0.0) # transform to local coordinate frame