Add filterexact to MJX.

PiperOrigin-RevId: 592026712
Change-Id: Ia73d67f36db5a4a1786a2532b5daa9e9a36f714c
This commit is contained in:
Baruch Tabanpour
2023-12-18 15:27:16 -08:00
committed by Copybara-Service
parent 7fb448613f
commit 80f50c943c
6 changed files with 99 additions and 20 deletions
+8
View File
@@ -2,6 +2,14 @@
Changelog
=========
Upcoming version (not yet released)
-----------------------------------
MJX
^^^
1. Add :ref:`dyntype<actuator-general-dyntype>` ``filterexact``.
Version 3.1.1 (December 18, 2023)
-----------------------------------
+4 -4
View File
@@ -183,7 +183,7 @@ The following features are **fully supported** in MJX:
* - :ref:`Transmission <mjtTrn>`
- ``TRN_JOINT``
* - :ref:`Actuator Dynamics <mjtDyn>`
- ``NONE``, ``INTEGRATOR``, ``FILTER``
- ``NONE``, ``INTEGRATOR``, ``FILTER``, ``FILTEREXACT``
* - :ref:`Actuator Gain <mjtGain>`
- ``FIXED``, ``AFFINE``
* - :ref:`Actuator Bias <mjtBias>`
@@ -218,7 +218,7 @@ The following features are **in development** and coming soon:
* - Dynamics
- :ref:`Inverse <mj_inverse>`
* - :ref:`Transmission <mjtTrn>`
- ``TRN_TENDON``
- ``TRN_SITE``, ``TRN_TENDON``
* - :ref:`Actuator Dynamics <mjtDyn>`
- ``MUSCLE``
* - :ref:`Actuator Gain <mjtGain>`
@@ -257,9 +257,9 @@ The following features are **unsupported**:
* - Category
- Feature
* - :ref:`Transmission <mjtTrn>`
- ``TRN_JOINTINPARENT``, ``TRN_SLIDERCRANK``, ``TRN_SITE``, ``TRN_BODY``
- ``TRN_JOINTINPARENT``, ``TRN_SLIDERCRANK``, ``TRN_BODY``
* - :ref:`Actuator Dynamics <mjtDyn>`
- ``FILTEREXACT``, ``USER``
- ``USER``
* - :ref:`Actuator Gain <mjtGain>`
- ``USER``
* - :ref:`Actuator Bias <mjtBias>`
+30 -11
View File
@@ -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)
+40
View File
@@ -134,5 +134,45 @@ class ForwardTest(absltest.TestCase):
np.testing.assert_allclose(dx.qvel, 1 + m.opt.timestep)
class ActuatorTest(absltest.TestCase):
_DYN_XML = """
<mujoco>
<compiler autolimits="true"/>
<worldbody>
<body name="box">
<joint name="slide1" type="slide" axis="1 0 0" />
<joint name="slide2" type="slide" axis="0 1 0" />
<joint name="slide3" type="slide" axis="0 0 1" />
<joint name="slide4" type="slide" axis="1 1 0" />
<geom type="box" size=".05 .05 .05" mass="1"/>
</body>
</worldbody>
<actuator>
<general joint="slide1" dynprm="0.1" gainprm="1.1" />
<general joint="slide2" dyntype="integrator" dynprm="0.1" gainprm="1.1" />
<general joint="slide3" dyntype="filter" dynprm="0.1" gainprm="1.1" />
<general joint="slide4" dyntype="filterexact" dynprm="0.1" gainprm="1.1" />
</actuator>
</mujoco>
"""
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()
+12
View File
@@ -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}'
+5 -5
View File
@@ -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):