diff --git a/doc/changelog.rst b/doc/changelog.rst index fa30391f..d094094e 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -1,6 +1,14 @@ ========= Changelog ========= + +Upcoming version (not yet release) +---------------------------------- + +MJX +^^^ +- Added inverse dynamics. + Version 3.3.1 (Apr 9, 2025) ---------------------------- diff --git a/doc/mjx.rst b/doc/mjx.rst index 98fb36b3..0c00525f 100644 --- a/doc/mjx.rst +++ b/doc/mjx.rst @@ -235,6 +235,8 @@ The following features are **fully supported** in MJX: - 1, 3, 4, 6 (1 is not supported with ``ELLIPTIC``) * - :ref:`Solver ` - ``CG``, ``NEWTON`` + * - Dynamics + - :ref:`Inverse ` * - Fluid Model - :ref:`flInertia` * - :ref:`Tendons ` @@ -262,8 +264,6 @@ The following features are **in development** and coming soon: (``BOX``, ``MESH``, ``HFIELD``) and ``ELLIPSOID``. * - :ref:`Integrator ` - ``IMPLICIT`` - * - Dynamics - - :ref:`Inverse ` * - Fluid Model - :ref:`flEllipsoid` * - :ref:`Sensors ` diff --git a/mjx/mujoco/mjx/__init__.py b/mjx/mujoco/mjx/__init__.py index 6a031b57..c60c785c 100644 --- a/mjx/mujoco/mjx/__init__.py +++ b/mjx/mujoco/mjx/__init__.py @@ -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 diff --git a/mjx/mujoco/mjx/_src/derivative.py b/mjx/mujoco/mjx/_src/derivative.py new file mode 100644 index 00000000..b0fc65f7 --- /dev/null +++ b/mjx/mujoco/mjx/_src/derivative.py @@ -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 diff --git a/mjx/mujoco/mjx/_src/forward.py b/mjx/mujoco/mjx/_src/forward.py index 6ad1ae71..967aa639 100644 --- a/mjx/mujoco/mjx/_src/forward.py +++ b/mjx/mujoco/mjx/_src/forward.py @@ -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: diff --git a/mjx/mujoco/mjx/_src/inverse.py b/mjx/mujoco/mjx/_src/inverse.py new file mode 100644 index 00000000..5ad3e1ad --- /dev/null +++ b/mjx/mujoco/mjx/_src/inverse.py @@ -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) diff --git a/mjx/mujoco/mjx/_src/inverse_test.py b/mjx/mujoco/mjx/_src/inverse_test.py new file mode 100644 index 00000000..19a5deef --- /dev/null +++ b/mjx/mujoco/mjx/_src/inverse_test.py @@ -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(""" + + + """) + 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() diff --git a/mjx/mujoco/mjx/_src/io.py b/mjx/mujoco/mjx/_src/io.py index aac18731..54ac3dac 100644 --- a/mjx/mujoco/mjx/_src/io.py +++ b/mjx/mujoco/mjx/_src/io.py @@ -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) diff --git a/mjx/mujoco/mjx/_src/smooth.py b/mjx/mujoco/mjx/_src/smooth.py index 87750c04..92e46702 100644 --- a/mjx/mujoco/mjx/_src/smooth.py +++ b/mjx/mujoco/mjx/_src/smooth.py @@ -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) diff --git a/mjx/mujoco/mjx/_src/smooth_test.py b/mjx/mujoco/mjx/_src/smooth_test.py index 01ce8abb..eca05068 100644 --- a/mjx/mujoco/mjx/_src/smooth_test.py +++ b/mjx/mujoco/mjx/_src/smooth_test.py @@ -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 diff --git a/mjx/mujoco/mjx/_src/solver.py b/mjx/mujoco/mjx/_src/solver.py index 6d3d4f21..efb58c41 100644 --- a/mjx/mujoco/mjx/_src/solver.py +++ b/mjx/mujoco/mjx/_src/solver.py @@ -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: diff --git a/mjx/mujoco/mjx/_src/solver_test.py b/mjx/mujoco/mjx/_src/solver_test.py index 57d4a116..6da06526 100644 --- a/mjx/mujoco/mjx/_src/solver_test.py +++ b/mjx/mujoco/mjx/_src/solver_test.py @@ -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') diff --git a/mjx/mujoco/mjx/_src/types.py b/mjx/mujoco/mjx/_src/types.py index 3ee10499..94587478 100644 --- a/mjx/mujoco/mjx/_src/types.py +++ b/mjx/mujoco/mjx/_src/types.py @@ -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.