Add inverse dynamics to MJX.
PiperOrigin-RevId: 745998273 Change-Id: I203af89332ace5d7a60fff3ec6ff4cc88c02340d
This commit is contained in:
committed by
Copybara-Service
parent
d3664b5d01
commit
51c489fc30
@@ -1,6 +1,14 @@
|
||||
=========
|
||||
Changelog
|
||||
=========
|
||||
|
||||
Upcoming version (not yet release)
|
||||
----------------------------------
|
||||
|
||||
MJX
|
||||
^^^
|
||||
- Added inverse dynamics.
|
||||
|
||||
Version 3.3.1 (Apr 9, 2025)
|
||||
----------------------------
|
||||
|
||||
|
||||
+2
-2
@@ -235,6 +235,8 @@ The following features are **fully supported** in MJX:
|
||||
- 1, 3, 4, 6 (1 is not supported with ``ELLIPTIC``)
|
||||
* - :ref:`Solver <mjtSolver>`
|
||||
- ``CG``, ``NEWTON``
|
||||
* - Dynamics
|
||||
- :ref:`Inverse <mj_inverse>`
|
||||
* - Fluid Model
|
||||
- :ref:`flInertia`
|
||||
* - :ref:`Tendons <tendon>`
|
||||
@@ -262,8 +264,6 @@ The following features are **in development** and coming soon:
|
||||
(``BOX``, ``MESH``, ``HFIELD``) and ``ELLIPSOID``.
|
||||
* - :ref:`Integrator <mjtIntegrator>`
|
||||
- ``IMPLICIT``
|
||||
* - Dynamics
|
||||
- :ref:`Inverse <mj_inverse>`
|
||||
* - Fluid Model
|
||||
- :ref:`flEllipsoid`
|
||||
* - :ref:`Sensors <mjtSensor>`
|
||||
|
||||
@@ -17,6 +17,7 @@
|
||||
# pylint:disable=g-importing-member
|
||||
from mujoco.mjx._src.collision_driver import collision
|
||||
from mujoco.mjx._src.constraint import make_constraint
|
||||
from mujoco.mjx._src.derivative import deriv_smooth_vel
|
||||
from mujoco.mjx._src.forward import euler
|
||||
from mujoco.mjx._src.forward import forward
|
||||
from mujoco.mjx._src.forward import fwd_acceleration
|
||||
@@ -26,6 +27,7 @@ from mujoco.mjx._src.forward import fwd_velocity
|
||||
from mujoco.mjx._src.forward import implicit
|
||||
from mujoco.mjx._src.forward import rungekutta4
|
||||
from mujoco.mjx._src.forward import step
|
||||
from mujoco.mjx._src.inverse import inverse
|
||||
from mujoco.mjx._src.io import get_data
|
||||
from mujoco.mjx._src.io import get_data_into
|
||||
from mujoco.mjx._src.io import make_data
|
||||
|
||||
@@ -0,0 +1,59 @@
|
||||
# Copyright 2025 DeepMind Technologies Limited
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# ==============================================================================
|
||||
"""Derivative functions."""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
import jax
|
||||
from jax import numpy as jp
|
||||
# pylint: disable=g-importing-member
|
||||
from mujoco.mjx._src.types import BiasType
|
||||
from mujoco.mjx._src.types import Data
|
||||
from mujoco.mjx._src.types import DisableBit
|
||||
from mujoco.mjx._src.types import DynType
|
||||
from mujoco.mjx._src.types import GainType
|
||||
from mujoco.mjx._src.types import Model
|
||||
|
||||
|
||||
def deriv_smooth_vel(m: Model, d: Data) -> Optional[jax.Array]:
|
||||
"""Analytical derivative of smooth forces w.r.t velocities."""
|
||||
|
||||
qderiv = None
|
||||
|
||||
# qDeriv += d qfrc_actuator / d qvel
|
||||
if not m.opt.disableflags & DisableBit.ACTUATION:
|
||||
affine_bias = m.actuator_biastype == BiasType.AFFINE
|
||||
bias_vel = m.actuator_biasprm[:, 2] * affine_bias
|
||||
affine_gain = m.actuator_gaintype == GainType.AFFINE
|
||||
gain_vel = m.actuator_gainprm[:, 2] * affine_gain
|
||||
ctrl = d.ctrl.at[m.actuator_dyntype != DynType.NONE].set(d.act)
|
||||
vel = bias_vel + gain_vel * ctrl
|
||||
qderiv = d.actuator_moment.T @ jax.vmap(jp.multiply)(d.actuator_moment, vel)
|
||||
|
||||
# qDeriv += d qfrc_passive / d qvel
|
||||
if not m.opt.disableflags & DisableBit.PASSIVE:
|
||||
if qderiv is None:
|
||||
qderiv = -jp.diag(m.dof_damping)
|
||||
else:
|
||||
qderiv -= jp.diag(m.dof_damping)
|
||||
if m.ntendon:
|
||||
qderiv -= d.ten_J.T @ jp.diag(m.tendon_damping) @ d.ten_J
|
||||
# TODO(robotics-simulation): fluid drag model
|
||||
if m.opt.has_fluid_params:
|
||||
raise NotImplementedError('fluid drag not supported for implicitfast')
|
||||
|
||||
# TODO(team): rne derivative
|
||||
|
||||
return qderiv
|
||||
@@ -22,6 +22,7 @@ from jax import numpy as jp
|
||||
import mujoco
|
||||
from mujoco.mjx._src import collision_driver
|
||||
from mujoco.mjx._src import constraint
|
||||
from mujoco.mjx._src import derivative
|
||||
from mujoco.mjx._src import math
|
||||
from mujoco.mjx._src import passive
|
||||
from mujoco.mjx._src import scan
|
||||
@@ -392,29 +393,7 @@ def rungekutta4(m: Model, d: Data) -> Data:
|
||||
def implicit(m: Model, d: Data) -> Data:
|
||||
"""Integrates fully implicit in velocity."""
|
||||
|
||||
qderiv = None
|
||||
|
||||
# qDeriv += d qfrc_actuator / d qvel
|
||||
if not m.opt.disableflags & DisableBit.ACTUATION:
|
||||
affine_bias = m.actuator_biastype == BiasType.AFFINE
|
||||
bias_vel = m.actuator_biasprm[:, 2] * affine_bias
|
||||
affine_gain = m.actuator_gaintype == GainType.AFFINE
|
||||
gain_vel = m.actuator_gainprm[:, 2] * affine_gain
|
||||
ctrl = d.ctrl.at[m.actuator_dyntype != DynType.NONE].set(d.act)
|
||||
vel = bias_vel + gain_vel * ctrl
|
||||
qderiv = d.actuator_moment.T @ jp.diag(vel) @ d.actuator_moment
|
||||
|
||||
# qDeriv += d qfrc_passive / d qvel
|
||||
if not m.opt.disableflags & DisableBit.PASSIVE:
|
||||
if qderiv is None:
|
||||
qderiv = -jp.diag(m.dof_damping)
|
||||
else:
|
||||
qderiv -= jp.diag(m.dof_damping)
|
||||
if m.ntendon:
|
||||
qderiv -= d.ten_J.T @ jp.diag(m.tendon_damping) @ d.ten_J
|
||||
# TODO(robotics-simulation): fluid drag model
|
||||
if m.opt.has_fluid_params:
|
||||
raise NotImplementedError('fluid drag not supported for implicitfast')
|
||||
qderiv = derivative.deriv_smooth_vel(m, d)
|
||||
|
||||
qacc = d.qacc
|
||||
if qderiv is not None:
|
||||
|
||||
@@ -0,0 +1,106 @@
|
||||
# Copyright 2025 DeepMind Technologies Limited
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# ==============================================================================
|
||||
"""Inverse dynamics functions."""
|
||||
|
||||
from jax import numpy as jp
|
||||
from mujoco.mjx._src import derivative
|
||||
from mujoco.mjx._src import forward
|
||||
from mujoco.mjx._src import sensor
|
||||
from mujoco.mjx._src import smooth
|
||||
from mujoco.mjx._src import solver
|
||||
from mujoco.mjx._src import support
|
||||
# pylint: disable=g-importing-member
|
||||
from mujoco.mjx._src.types import Data
|
||||
from mujoco.mjx._src.types import DisableBit
|
||||
from mujoco.mjx._src.types import EnableBit
|
||||
from mujoco.mjx._src.types import IntegratorType
|
||||
from mujoco.mjx._src.types import Model
|
||||
|
||||
|
||||
def discrete_acc(m: Model, d: Data) -> Data:
|
||||
"""Convert discrete-time qacc to continuous-time qacc."""
|
||||
|
||||
if m.opt.integrator == IntegratorType.RK4:
|
||||
raise RuntimeError(
|
||||
'discrete inverse dynamics is not supported by RK4 integrator'
|
||||
)
|
||||
elif m.opt.integrator == IntegratorType.EULER:
|
||||
dsbl_eulerdamp = m.opt.disableflags & DisableBit.EULERDAMP
|
||||
no_dof_damping = (m.dof_damping == 0).all()
|
||||
if dsbl_eulerdamp or no_dof_damping:
|
||||
return d
|
||||
|
||||
# set qfrc = (M + h*diag(B)) * qacc
|
||||
qfrc = support.mul_m(m, d, d.qacc)
|
||||
qfrc += m.opt.timestep * m.dof_damping * d.qacc
|
||||
elif m.opt.integrator == IntegratorType.IMPLICITFAST:
|
||||
qm = support.full_m(m, d)
|
||||
|
||||
# compute analytical derivative qDeriv; skip rne derivative
|
||||
qderiv = derivative.deriv_smooth_vel(m, d)
|
||||
if qderiv is not None:
|
||||
# M = M - dt*qDeriv
|
||||
qm -= m.opt.timestep * qderiv
|
||||
|
||||
# set qfrc = (M - dt*qDeriv) * qacc
|
||||
qfrc = qm @ d.qacc
|
||||
else:
|
||||
raise NotImplementedError(f'integrator {m.opt.integrator} not implemented.')
|
||||
|
||||
# solve for qacc: qfrc = M * qacc
|
||||
qacc = smooth.solve_m(m, d, qfrc)
|
||||
|
||||
return d.replace(qacc=qacc)
|
||||
|
||||
|
||||
def inv_constraint(m: Model, d: Data) -> Data:
|
||||
"""Inverse constraint solver."""
|
||||
|
||||
# no constraints
|
||||
if d.efc_J.size == 0:
|
||||
return d.replace(qfrc_constraint=jp.zeros(m.nv))
|
||||
|
||||
# update
|
||||
ctx = solver.Context.create(m, d, grad=False)
|
||||
|
||||
return d.replace(
|
||||
qfrc_constraint=ctx.qfrc_constraint,
|
||||
efc_force=ctx.efc_force,
|
||||
)
|
||||
|
||||
|
||||
def inverse(m: Model, d: Data) -> Data:
|
||||
"""Inverse dynamics."""
|
||||
d = forward.fwd_position(m, d)
|
||||
d = sensor.sensor_pos(m, d)
|
||||
d = forward.fwd_velocity(m, d)
|
||||
d = sensor.sensor_vel(m, d)
|
||||
|
||||
qacc = d.qacc
|
||||
if m.opt.enableflags & EnableBit.INVDISCRETE:
|
||||
d = discrete_acc(m, d)
|
||||
|
||||
d = inv_constraint(m, d)
|
||||
d = smooth.rne(m, d, flg_acc=True)
|
||||
d = sensor.sensor_acc(m, d)
|
||||
|
||||
qfrc_inverse = (
|
||||
d.qfrc_bias + m.dof_armature * d.qacc - d.qfrc_passive - d.qfrc_constraint
|
||||
)
|
||||
|
||||
if m.opt.enableflags & EnableBit.INVDISCRETE:
|
||||
return d.replace(qfrc_inverse=qfrc_inverse, qacc=qacc)
|
||||
else:
|
||||
return d.replace(qfrc_inverse=qfrc_inverse)
|
||||
@@ -0,0 +1,131 @@
|
||||
# Copyright 2023 DeepMind Technologies Limited
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# ==============================================================================
|
||||
"""Tests for inverse dynamics functions."""
|
||||
from absl.testing import absltest
|
||||
from absl.testing import parameterized
|
||||
import jax
|
||||
from jax import numpy as jp
|
||||
import mujoco
|
||||
from mujoco import mjx
|
||||
from mujoco.mjx._src import support
|
||||
from mujoco.mjx._src import test_util
|
||||
import numpy as np
|
||||
|
||||
# tolerance for difference between MuJoCo and MJX calculations - mostly
|
||||
# due to float precision
|
||||
_TOLERANCE = 1e-5
|
||||
|
||||
|
||||
def _assert_eq(a, b, name, tol=_TOLERANCE):
|
||||
tol = tol * 10 # avoid test noise
|
||||
err_msg = f'mismatch: {name}'
|
||||
np.testing.assert_allclose(a, b, err_msg=err_msg, atol=tol, rtol=tol)
|
||||
|
||||
|
||||
class InverseTest(parameterized.TestCase):
|
||||
|
||||
@parameterized.parameters(
|
||||
(mujoco.mjtIntegrator.mjINT_EULER, False, False),
|
||||
(mujoco.mjtIntegrator.mjINT_EULER, False, True),
|
||||
(mujoco.mjtIntegrator.mjINT_EULER, True, False),
|
||||
(mujoco.mjtIntegrator.mjINT_EULER, True, True),
|
||||
(mujoco.mjtIntegrator.mjINT_IMPLICITFAST, False, False),
|
||||
(mujoco.mjtIntegrator.mjINT_IMPLICITFAST, True, False),
|
||||
)
|
||||
def test_forward_inverse_match(self, integrator, invdiscrete, eulerdamp):
|
||||
m = mujoco.MjModel.from_xml_string("""
|
||||
<mujoco>
|
||||
<option timestep=".005" gravity="-1 -1 -10"/>
|
||||
<worldbody>
|
||||
<geom type="plane" size="10 10 .001"/>
|
||||
<body pos="0 0 1">
|
||||
<geom type="sphere" size=".1" pos=".1 .2 .3"/>
|
||||
<joint name="jnt1" type="hinge" axis="0 1 0" stiffness=".25" damping=".125"/>
|
||||
<body pos="0 0 1">
|
||||
<geom type="sphere" size=".1" pos=".1 .2 .3"/>
|
||||
<joint name="jnt2" type="hinge" axis="0 1 0" stiffness=".6" damping=".3"/>
|
||||
</body>
|
||||
</body>
|
||||
</worldbody>
|
||||
<actuator>
|
||||
<motor joint="jnt1"/>
|
||||
</actuator>
|
||||
<equality>
|
||||
<joint joint1="jnt1" joint2="jnt2"/>
|
||||
</equality>
|
||||
</mujoco>
|
||||
""")
|
||||
m.opt.integrator = integrator
|
||||
if invdiscrete:
|
||||
m.opt.enableflags |= mujoco.mjtEnableBit.mjENBL_INVDISCRETE
|
||||
if not eulerdamp:
|
||||
m.opt.disableflags |= mujoco.mjtDisableBit.mjDSBL_EULERDAMP
|
||||
|
||||
d = mujoco.MjData(m)
|
||||
d.qvel = np.random.uniform(low=-0.01, high=0.01, size=d.qvel.shape)
|
||||
d.ctrl = np.random.uniform(low=-0.01, high=0.01, size=d.ctrl.shape)
|
||||
d.qfrc_applied = np.random.uniform(
|
||||
low=-0.01, high=0.01, size=d.qfrc_applied.shape
|
||||
)
|
||||
d.xfrc_applied = np.random.uniform(
|
||||
low=-0.01, high=0.01, size=d.xfrc_applied.shape
|
||||
)
|
||||
mujoco.mj_step(m, d, 100)
|
||||
|
||||
mx = mjx.put_model(m)
|
||||
dx = mjx.put_data(m, d)
|
||||
dx_next = mjx.step(mx, dx)
|
||||
qacc_fd = (dx_next.qvel - dx.qvel) / mx.opt.timestep
|
||||
|
||||
dx = mjx.forward(mx, dx)
|
||||
|
||||
if invdiscrete:
|
||||
dx = dx.replace(qacc=qacc_fd)
|
||||
|
||||
dxinv = mjx.inverse(mx, dx)
|
||||
|
||||
fwdinv0 = jp.linalg.norm(
|
||||
dxinv.qfrc_constraint - dx.qfrc_constraint, ord=np.inf
|
||||
)
|
||||
fwdinv1 = jp.linalg.norm(
|
||||
dxinv.qfrc_inverse
|
||||
- (
|
||||
dx.qfrc_applied + dx.qfrc_actuator + support.xfrc_accumulate(mx, dx)
|
||||
),
|
||||
ord=np.inf,
|
||||
)
|
||||
|
||||
self.assertLess(fwdinv0, 1.0e-3)
|
||||
self.assertLess(fwdinv1, 1.0e-3)
|
||||
_assert_eq(dxinv.qacc, dx.qacc, 'qacc')
|
||||
|
||||
def test_tendon_force_clamp(self):
|
||||
m = test_util.load_test_file('actuator/tendon_force_clamp.xml')
|
||||
d = mujoco.MjData(m)
|
||||
mx = mjx.put_model(m)
|
||||
dx = mjx.put_data(m, d)
|
||||
|
||||
dx = dx.replace(ctrl=jp.array([1.0, 1.0, 1.0, -1.0, 1.0, -20.0, 5.0, -5.0]))
|
||||
dx = mjx.forward(mx, dx)
|
||||
|
||||
_assert_eq(
|
||||
dx.actuator_force,
|
||||
jp.array([1.0, 1.0, 1.0, -1.0, 1.0, -10.0, 5.0, -5.0]),
|
||||
'actuator_force',
|
||||
)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
absltest.main()
|
||||
@@ -56,8 +56,8 @@ def _make_option(
|
||||
raise NotImplementedError(f'{mujoco.mjtSolver(o.solver)}')
|
||||
|
||||
for i in range(mujoco.mjtEnableBit.mjNENABLE):
|
||||
if o.enableflags & 2**i:
|
||||
raise NotImplementedError(f'{mujoco.mjtEnableBit(2 ** i)}')
|
||||
if o.enableflags & 2**i and 2**i not in set(types.EnableBit):
|
||||
raise NotImplementedError(f'{mujoco.mjtEnableBit(2**i)}')
|
||||
|
||||
has_fluid_params = o.density > 0 or o.viscosity > 0 or o.wind.any()
|
||||
implicitfast = o.integrator == mujoco.mjtIntegrator.mjINT_IMPLICITFAST
|
||||
@@ -72,6 +72,7 @@ def _make_option(
|
||||
fields['solver'] = types.SolverType(o.solver)
|
||||
fields['disableflags'] = types.DisableBit(o.disableflags)
|
||||
fields['has_fluid_params'] = has_fluid_params
|
||||
fields['enableflags'] = types.EnableBit(o.enableflags)
|
||||
|
||||
return types.Option(**fields)
|
||||
|
||||
|
||||
@@ -532,11 +532,14 @@ def subtree_vel(m: Model, d: Data) -> Data:
|
||||
return d.replace(subtree_linvel=subtree_linvel, subtree_angmom=subtree_angmom)
|
||||
|
||||
|
||||
def rne(m: Model, d: Data) -> Data:
|
||||
"""Computes inverse dynamics using the recursive Newton-Euler algorithm."""
|
||||
def rne(m: Model, d: Data, flg_acc: bool = False) -> Data:
|
||||
"""Computes inverse dynamics using the recursive Newton-Euler algorithm.
|
||||
|
||||
flg_acc=False removes inertial term.
|
||||
"""
|
||||
|
||||
# forward scan over tree: accumulate link center of mass acceleration
|
||||
def cacc_fn(cacc, cdof_dot, qvel):
|
||||
def cacc_fn(cacc, cdof_dot, qvel, cdof, qacc):
|
||||
if cacc is None:
|
||||
if m.opt.disableflags & DisableBit.GRAVITY:
|
||||
cacc = jp.zeros((6,))
|
||||
@@ -545,9 +548,15 @@ def rne(m: Model, d: Data) -> Data:
|
||||
|
||||
cacc += jp.sum(jax.vmap(jp.multiply)(cdof_dot, qvel), axis=0)
|
||||
|
||||
# cacc += cdof * qacc
|
||||
if flg_acc:
|
||||
cacc += jp.sum(jax.vmap(jp.multiply)(cdof, qacc), axis=0)
|
||||
|
||||
return cacc
|
||||
|
||||
cacc = scan.body_tree(m, cacc_fn, 'vv', 'b', d.cdof_dot, d.qvel)
|
||||
cacc = scan.body_tree(
|
||||
m, cacc_fn, 'vvvv', 'b', d.cdof_dot, d.qvel, d.cdof, d.qacc
|
||||
)
|
||||
|
||||
def frc(cinert, cacc, cvel):
|
||||
frc = math.inert_mul(cinert, cacc)
|
||||
|
||||
@@ -105,6 +105,13 @@ class SmoothTest(absltest.TestCase):
|
||||
# rne
|
||||
dx = jax.jit(mjx.rne)(mx, mjx.put_data(m, d))
|
||||
_assert_attr_eq(d, dx, 'qfrc_bias')
|
||||
# rne (flg_acc=True)
|
||||
qfrc_bias = np.zeros(m.nv)
|
||||
mujoco.mj_rne(m, d, 1, qfrc_bias)
|
||||
dx = jax.jit(mjx.rne, static_argnums=(2,))(
|
||||
mx, mjx.put_data(m, d), flg_acc=True
|
||||
)
|
||||
_assert_eq(dx.qfrc_bias, qfrc_bias, 'qfrc_bias')
|
||||
|
||||
# set dense jacobian for tendon:
|
||||
m.opt.jacobian = mujoco.mjtJacobian.mjJAC_DENSE
|
||||
|
||||
@@ -30,7 +30,7 @@ from mujoco.mjx._src.types import SolverType
|
||||
# pylint: enable=g-importing-member
|
||||
|
||||
|
||||
class _Context(PyTreeNode):
|
||||
class Context(PyTreeNode):
|
||||
"""Data updated during each solver iteration.
|
||||
|
||||
Attributes:
|
||||
@@ -72,7 +72,7 @@ class _Context(PyTreeNode):
|
||||
h: jax.Array
|
||||
|
||||
@classmethod
|
||||
def create(cls, m: Model, d: Data, grad: bool = True) -> '_Context':
|
||||
def create(cls, m: Model, d: Data, grad: bool = True) -> 'Context':
|
||||
jaref = d.efc_J @ d.qacc - d.efc_aref
|
||||
# TODO(robotics-team): determine nv at which sparse mul is faster
|
||||
ma = support.mul_m(m, d, d.qacc)
|
||||
@@ -86,7 +86,7 @@ class _Context(PyTreeNode):
|
||||
for condim in (3, 4, 6):
|
||||
fri = fri.at[dim == condim, condim:].set(0)
|
||||
|
||||
ctx = _Context(
|
||||
ctx = Context(
|
||||
qacc=d.qacc,
|
||||
qfrc_constraint=d.qfrc_constraint,
|
||||
Jaref=jaref,
|
||||
@@ -133,7 +133,7 @@ class _LSPoint(PyTreeNode):
|
||||
cls,
|
||||
m: Model,
|
||||
d: Data,
|
||||
ctx: _Context,
|
||||
ctx: Context,
|
||||
alpha: jax.Array,
|
||||
jv: jax.Array,
|
||||
quad: jax.Array,
|
||||
@@ -241,7 +241,7 @@ def _while_loop_scan(cond_fun, body_fun, init_val, max_iter):
|
||||
return jax.lax.scan(_fun, init, None, length=max_iter)[0][0]
|
||||
|
||||
|
||||
def _update_constraint(m: Model, d: Data, ctx: _Context) -> _Context:
|
||||
def _update_constraint(m: Model, d: Data, ctx: Context) -> Context:
|
||||
"""Updates constraint force and resulting cost given last solver iteration.
|
||||
|
||||
Corresponds to CGupdateConstraint in mujoco/src/engine/engine_solver.c
|
||||
@@ -356,7 +356,7 @@ def _update_constraint(m: Model, d: Data, ctx: _Context) -> _Context:
|
||||
return ctx
|
||||
|
||||
|
||||
def _update_gradient(m: Model, d: Data, ctx: _Context) -> _Context:
|
||||
def _update_gradient(m: Model, d: Data, ctx: Context) -> Context:
|
||||
"""Updates grad and M / grad given latest solver iteration.
|
||||
|
||||
Corresponds to CGupdateGradient in mujoco/src/engine/engine_solver.c
|
||||
@@ -403,7 +403,7 @@ def _rescale(m: Model, value: jax.Array) -> jax.Array:
|
||||
return value / (m.stat.meaninertia * max(1, m.nv))
|
||||
|
||||
|
||||
def _linesearch(m: Model, d: Data, ctx: _Context) -> _Context:
|
||||
def _linesearch(m: Model, d: Data, ctx: Context) -> Context:
|
||||
"""Performs a zoom linesearch to find optimal search step size.
|
||||
|
||||
Args:
|
||||
@@ -529,7 +529,7 @@ def _linesearch(m: Model, d: Data, ctx: _Context) -> _Context:
|
||||
def solve(m: Model, d: Data) -> Data:
|
||||
"""Finds forces that satisfy constraints using conjugate gradient descent."""
|
||||
|
||||
def cond(ctx: _Context) -> jax.Array:
|
||||
def cond(ctx: Context) -> jax.Array:
|
||||
improvement = _rescale(m, ctx.prev_cost - ctx.cost)
|
||||
gradient = _rescale(m, math.norm(ctx.grad))
|
||||
|
||||
@@ -539,7 +539,7 @@ def solve(m: Model, d: Data) -> Data:
|
||||
|
||||
return ~done
|
||||
|
||||
def body(ctx: _Context) -> _Context:
|
||||
def body(ctx: Context) -> Context:
|
||||
ctx = _linesearch(m, d, ctx)
|
||||
prev_grad, prev_Mgrad = ctx.grad, ctx.Mgrad # pylint: disable=invalid-name
|
||||
ctx = _update_constraint(m, d, ctx)
|
||||
@@ -560,12 +560,12 @@ def solve(m: Model, d: Data) -> Data:
|
||||
# warmstart:
|
||||
qacc = d.qacc_smooth
|
||||
if not m.opt.disableflags & DisableBit.WARMSTART:
|
||||
warm = _Context.create(m, d.replace(qacc=d.qacc_warmstart), grad=False)
|
||||
smth = _Context.create(m, d.replace(qacc=d.qacc_smooth), grad=False)
|
||||
warm = Context.create(m, d.replace(qacc=d.qacc_warmstart), grad=False)
|
||||
smth = Context.create(m, d.replace(qacc=d.qacc_smooth), grad=False)
|
||||
qacc = jp.where(warm.cost < smth.cost, d.qacc_warmstart, d.qacc_smooth)
|
||||
d = d.replace(qacc=qacc)
|
||||
|
||||
ctx = _Context.create(m, d)
|
||||
ctx = Context.create(m, d)
|
||||
if m.opt.iterations == 1:
|
||||
ctx = body(ctx)
|
||||
else:
|
||||
|
||||
@@ -74,7 +74,7 @@ class SolverTest(parameterized.TestCase):
|
||||
|
||||
# compare costs
|
||||
mj_cost = cost(d.qacc)
|
||||
ctx = solver._Context.create(mjx.put_model(m), mjx.put_data(m, d))
|
||||
ctx = solver.Context.create(mjx.put_model(m), mjx.put_data(m, d))
|
||||
mjx_cost = ctx.cost - ctx.gauss
|
||||
_assert_eq(mj_cost, mjx_cost, 'cost')
|
||||
|
||||
|
||||
@@ -65,6 +65,17 @@ class DisableBit(enum.IntFlag):
|
||||
# unsupported: MIDPHASE
|
||||
|
||||
|
||||
class EnableBit(enum.IntFlag):
|
||||
"""Enable optional feature bitflags.
|
||||
|
||||
Members:
|
||||
INVDISCRETE: discrete-time inverse dynamics
|
||||
"""
|
||||
|
||||
INVDISCRETE = mujoco.mjtEnableBit.mjENBL_INVDISCRETE
|
||||
# unsupported: OVERRIDE, ENERGY, FWDINV, MULTICCD, ISLAND
|
||||
|
||||
|
||||
class JointType(enum.IntEnum):
|
||||
"""Type of degree of freedom.
|
||||
|
||||
|
||||
Reference in New Issue
Block a user