diff --git a/doc/changelog.rst b/doc/changelog.rst index b70c7e9e..ad621fb4 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -2,6 +2,14 @@ Changelog ========= +Upcoming version (not yet released) +----------------------------------- + +MJX +^^^ + +1. Add :ref:`dyntype` ``filterexact``. + Version 3.1.1 (December 18, 2023) ----------------------------------- diff --git a/doc/mjx.rst b/doc/mjx.rst index 92a709b6..cc052182 100644 --- a/doc/mjx.rst +++ b/doc/mjx.rst @@ -183,7 +183,7 @@ The following features are **fully supported** in MJX: * - :ref:`Transmission ` - ``TRN_JOINT`` * - :ref:`Actuator Dynamics ` - - ``NONE``, ``INTEGRATOR``, ``FILTER`` + - ``NONE``, ``INTEGRATOR``, ``FILTER``, ``FILTEREXACT`` * - :ref:`Actuator Gain ` - ``FIXED``, ``AFFINE`` * - :ref:`Actuator Bias ` @@ -218,7 +218,7 @@ The following features are **in development** and coming soon: * - Dynamics - :ref:`Inverse ` * - :ref:`Transmission ` - - ``TRN_TENDON`` + - ``TRN_SITE``, ``TRN_TENDON`` * - :ref:`Actuator Dynamics ` - ``MUSCLE`` * - :ref:`Actuator Gain ` @@ -257,9 +257,9 @@ The following features are **unsupported**: * - Category - Feature * - :ref:`Transmission ` - - ``TRN_JOINTINPARENT``, ``TRN_SLIDERCRANK``, ``TRN_SITE``, ``TRN_BODY`` + - ``TRN_JOINTINPARENT``, ``TRN_SLIDERCRANK``, ``TRN_BODY`` * - :ref:`Actuator Dynamics ` - - ``FILTEREXACT``, ``USER`` + - ``USER`` * - :ref:`Actuator Gain ` - ``USER`` * - :ref:`Actuator Bias ` diff --git a/mjx/mujoco/mjx/_src/forward.py b/mjx/mujoco/mjx/_src/forward.py index 827c8f31..1ff268e7 100644 --- a/mjx/mujoco/mjx/_src/forward.py +++ b/mjx/mujoco/mjx/_src/forward.py @@ -107,7 +107,7 @@ def fwd_actuation(m: Model, d: Data) -> Data: act_dot = jp.array(0.0) elif dyn_typ == DynType.INTEGRATOR: act_dot = ctrl - elif dyn_typ == DynType.FILTER: + elif dyn_typ in (DynType.FILTER, DynType.FILTEREXACT): act_dot = (ctrl - act) / jp.clip(dyn_prm[0], mujoco.mjMINVAL) else: raise NotImplementedError(f'dyntype {dyn_typ.name} not implemented.') @@ -228,6 +228,34 @@ def _integrate_pos( return jp.concatenate(qs) if qs else jp.empty((0,)) +def _next_activation(m: Model, d: Data, act_dot: jax.Array) -> jax.Array: + """Returns the next act given the current act_dot, after clamping.""" + act = d.act + + if not m.na: + return act + + actrange = jp.where( + m.actuator_actlimited[:, None], + m.actuator_actrange, + jp.array([-jp.inf, jp.inf]), + ) + + def fn(dyntype, dynprm, act, act_dot, actrange): + if dyntype == DynType.FILTEREXACT: + tau = jp.clip(dynprm[0], a_min=mujoco.mjMINVAL) + act = act + act_dot * tau * (1 - jp.exp(-m.opt.timestep / tau)) + else: + act = act + act_dot * m.opt.timestep + act = jp.clip(act, actrange[0], actrange[1]) + return act + + args = (m.actuator_dyntype, m.actuator_dynprm, act, act_dot, actrange) + act = scan.flat(m, fn, 'uuaau', 'a', *args, group_by='u') + + return act.reshape(m.na) + + @named_scope def _advance( m: Model, @@ -237,16 +265,7 @@ def _advance( qvel: Optional[jax.Array] = None, ) -> Data: """Advance state and time given activation derivatives and acceleration.""" - act = d.act - if m.na: - act = d.act + act_dot * m.opt.timestep - actrange = jp.where( - m.actuator_actlimited[:, None], - m.actuator_actrange, - jp.array([-jp.inf, jp.inf]), - ) - fn = lambda act, actrange: jp.clip(act, actrange[0], actrange[1]) - act = scan.flat(m, fn, 'au', 'a', act, actrange, group_by='u') + act = _next_activation(m, d, act_dot) # advance velocities d = d.replace(qvel=d.qvel + qacc * m.opt.timestep) diff --git a/mjx/mujoco/mjx/_src/forward_test.py b/mjx/mujoco/mjx/_src/forward_test.py index fbe28c9d..cc767ae8 100644 --- a/mjx/mujoco/mjx/_src/forward_test.py +++ b/mjx/mujoco/mjx/_src/forward_test.py @@ -134,5 +134,45 @@ class ForwardTest(absltest.TestCase): np.testing.assert_allclose(dx.qvel, 1 + m.opt.timestep) +class ActuatorTest(absltest.TestCase): + _DYN_XML = """ + + + + + + + + + + + + + + + + + + + """ + + def test_dyntype(self): + m = mujoco.MjModel.from_xml_string(self._DYN_XML) + d = mujoco.MjData(m) + d.ctrl = np.array([1.5, 1.5, 1.5, 1.5]) + d.act = np.array([0.5, 0.5, 0.5]) + + mx = mjx.put_model(m) + dx = mjx.put_data(m, d) + + mujoco.mj_fwdActuation(m, d) + dx = jax.jit(mjx.fwd_actuation)(mx, dx) + _assert_attr_eq(d, dx, 'act_dot') + + mujoco.mj_Euler(m, d) + dx = jax.jit(mjx.euler)(mx, dx) + _assert_attr_eq(d, dx, 'act') + + if __name__ == '__main__': absltest.main() diff --git a/mjx/mujoco/mjx/_src/test_util.py b/mjx/mujoco/mjx/_src/test_util.py index 870621f7..85cdb5dc 100644 --- a/mjx/mujoco/mjx/_src/test_util.py +++ b/mjx/mujoco/mjx/_src/test_util.py @@ -29,6 +29,8 @@ TEST_FILES: List[str] = [ ] _ACTUATOR_TYPES = ['motor', 'velocity', 'position', 'general', 'intvelocity'] +_DYN_TYPES = ['none', 'integrator', 'filter', 'filterexact'] +_DYN_PRMS = ['0.189', '2.1'] _JOINT_TYPES = ['free', 'hinge', 'slide', 'ball'] _JOINT_AXES = ['1 0 0', '0 1 0', '0 0 1'] _FRICTIONS = ['1.2 0.003 0.0002', '0.2 0.0001 0.0005'] @@ -124,6 +126,8 @@ def _make_geom( def _make_actuator(actuator_type: str, joint: str) -> Dict[str, str]: """Returns attributes for an actuator.""" attr = {'joint': joint} + + # set actuator type if actuator_type == 'motor': attr['gear'] = np.random.choice(_GEARS) elif actuator_type == 'position': @@ -139,10 +143,18 @@ def _make_actuator(actuator_type: str, joint: str) -> Dict[str, str]: elif actuator_type == 'velocity': attr['kv'] = np.random.choice(_KV_VEL) + # set dyntype + if actuator_type == 'general': + attr['dyntype'] = np.random.choice(_DYN_TYPES) + if attr['dyntype'] != 'none': + attr['dynprm'] = np.random.choice(_DYN_PRMS) + + # ctrlrange if p(50) and actuator_type != 'intvelocity': lb, ub = -np.random.uniform(), np.random.uniform() attr['ctrlrange'] = f'{lb:.2f} {ub:.2f}' + # forcerange if p(50): lb, ub = -np.random.uniform(), np.random.uniform() attr['forcerange'] = f'{lb*10:.2f} {ub*10:.2f}' diff --git a/mjx/mujoco/mjx/_src/types.py b/mjx/mujoco/mjx/_src/types.py index 19b37630..9ccbe29a 100644 --- a/mjx/mujoco/mjx/_src/types.py +++ b/mjx/mujoco/mjx/_src/types.py @@ -20,10 +20,7 @@ from typing import Sequence import jax import jax.numpy as jp import mujoco -# pylint: disable=g-importing-member -from mujoco.mjx._src import dataclasses -from mujoco.mjx._src.dataclasses import PyTreeNode -# pylint: enable=g-importing-member +from mujoco.mjx._src.dataclasses import PyTreeNode # pylint: disable=g-importing-member import numpy as np @@ -167,11 +164,14 @@ class DynType(enum.IntEnum): Attributes: NONE: no internal dynamics; ctrl specifies force INTEGRATOR: integrator: da/dt = u + FILTER: linear filter: da/dt = (u-a) / tau + FILTEREXACT: linear filter: da/dt = (u-a) / tau, with exact integration """ NONE = mujoco.mjtDyn.mjDYN_NONE INTEGRATOR = mujoco.mjtDyn.mjDYN_INTEGRATOR FILTER = mujoco.mjtDyn.mjDYN_FILTER - # unsupported: FILTEREXACT, MUSCLE, USER + FILTEREXACT = mujoco.mjtDyn.mjDYN_FILTEREXACT + # unsupported: MUSCLE, USER class GainType(enum.IntEnum):