Add filterexact to MJX.
PiperOrigin-RevId: 592026712 Change-Id: Ia73d67f36db5a4a1786a2532b5daa9e9a36f714c
This commit is contained in:
committed by
Copybara-Service
parent
7fb448613f
commit
80f50c943c
@@ -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
@@ -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>`
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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}'
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user