Adds support for Newton solver to MJX.
PiperOrigin-RevId: 578738033 Change-Id: Id0e68c7b961380cb1ed88b40d99f202c3fa8e228
This commit is contained in:
committed by
Copybara-Service
parent
e82cb9f420
commit
3c0a56c1e5
+28
-3
@@ -28,12 +28,37 @@ General
|
||||
|
||||
MJX
|
||||
^^^
|
||||
- Added support for Newton solver (``mjSOL_NEWTON`` in :ref:`mjtSolver`). The Newton solver significantly speeds up
|
||||
simulation on GPU. See updated benchmarks on an Nvidia A100 (reported numbers are steps-per-second):
|
||||
|
||||
.. list-table:: Steps-per-second, Conjugate Gradient vs. Newton on A100
|
||||
:header-rows: 1
|
||||
:align: left
|
||||
|
||||
* - Model
|
||||
- CG
|
||||
- Newton
|
||||
- Speedup
|
||||
* - `Humanoid <https://github.com/google-deepmind/mujoco/tree/main/mjx/mujoco/mjx/benchmark/model/humanoid>`__
|
||||
- 640,000
|
||||
- 1,020,000
|
||||
- **1.6 x**
|
||||
* - `Barkour v0 <https://github.com/google-deepmind/mujoco/tree/main/mjx/mujoco/mjx/benchmark/model/barkour_v0>`__
|
||||
- 1,290,000
|
||||
- 1,750,000
|
||||
- **1.35 x**
|
||||
* - `Shadow Hand <https://github.com/google-deepmind/mujoco/tree/main/mjx/mujoco/mjx/benchmark/model/shadow_hand>`__
|
||||
- 215,000
|
||||
- 270,000
|
||||
- **1.25 x**
|
||||
|
||||
Humanoid is the standard MuJoCo humanoid,
|
||||
`Google Barkour <https://blog.research.google/2023/05/barkour-benchmarking-animal-level.html>`__ and the Shadow Hand
|
||||
are both available in the :ref:`MuJoCo Menagerie<Menagerie>`.
|
||||
- Added support for joint equality constraints (``mjEQ_JOINT`` in :ref:`mjtEq`).
|
||||
- Fixed bug where mixed ``jnt_limited`` joints were not being constrained correctly.
|
||||
- Made ``device_put`` type validation more verbose (fixes :github:issue:`1113`).
|
||||
- Removed empty EFC rows from `MJX`, for joints with no limits (fixes :github:issue:`1117`).
|
||||
- Fixed bug where equality constraints became inactive (fixes :github:issue:`1129`).
|
||||
- Added an error when loading a model with tendons (fixes :github:issue:`1149`).
|
||||
- Removed empty EFC rows from ``MJX``, for joints with no limits (fixes :github:issue:`1117`).
|
||||
|
||||
Python bindings
|
||||
^^^^^^^^^^^^^^^
|
||||
|
||||
+6
-7
@@ -198,7 +198,7 @@ The following features are **fully supported** in MJX:
|
||||
* - :ref:`Condim <coContact>`
|
||||
- 3
|
||||
* - :ref:`Solver <mjtSolver>`
|
||||
- ``CG``
|
||||
- ``CG``, ``NEWTON``
|
||||
* - Fluid Model
|
||||
- :ref:`flInertia`
|
||||
|
||||
@@ -226,8 +226,6 @@ The following features are **in development** and coming soon:
|
||||
- ``ELLIPTIC``
|
||||
* - :ref:`Condim <coContact>`
|
||||
- 1, 4, 6
|
||||
* - :ref:`Solver <mjtSolver>`
|
||||
- ``NEWTON``
|
||||
* - Fluid Model
|
||||
- :ref:`flEllipsoid`
|
||||
* - :ref:`Tendons <tendon>`
|
||||
@@ -306,10 +304,11 @@ Performance tuning
|
||||
For MJX to perform well, some configuration parameters should be adjusted from their default MuJoCo values:
|
||||
|
||||
:ref:`option` element
|
||||
For now, solver must be set to ``CG`` (but Newton is on its way!). The ``iterations`` and ``ls_iterations``
|
||||
attributes---which control solver and linesearch iterations, respectively---should be brought down to just low enough
|
||||
that the simulation remains stable. Accurate solver forces are not so important in reinforcement learning in which
|
||||
domain randomization is often used to add noise to physics for sim2real.
|
||||
The ``iterations`` and ``ls_iterations`` attributes---which control solver and linesearch iterations, respectively---
|
||||
should be brought down to just low enough that the simulation remains stable. Accurate solver forces are not so
|
||||
important in reinforcement learning in which domain randomization is often used to add noise to physics for sim-to-real.
|
||||
The ``NEWTON`` :ref:`Solver <mjtSolver>` often delivers reasonable convergence with one solver iteration, and performs
|
||||
well on GPU. ``CG`` is currently a better choice for TPU.
|
||||
|
||||
:ref:`contact-pair` element
|
||||
Consider explicitly marking geoms for collision detection to reduce the number of contacts that MJX must consider
|
||||
|
||||
@@ -40,6 +40,10 @@ _MJ_TYPE_ATTR = {
|
||||
mujoco.MjModel.opt,
|
||||
mujoco.MjOption.integrator,
|
||||
),
|
||||
mujoco.mjtSolver: (
|
||||
mujoco.MjModel.opt,
|
||||
mujoco.MjOption.solver,
|
||||
),
|
||||
}
|
||||
|
||||
_TYPE_MAP = {
|
||||
@@ -103,10 +107,6 @@ def _option_derived(value: types.Option) -> Dict[str, Any]:
|
||||
|
||||
def _validate(m: mujoco.MjModel):
|
||||
"""Validates that an mjModel is compatible with MJX."""
|
||||
if m.opt.solver not in set(types.SolverType):
|
||||
name = mujoco.mjtSolver(m.opt.solver).name
|
||||
warnings.warn(f'Solver {name} is not supported, reverting to CG.')
|
||||
m.opt.solver = mujoco.mjtSolver.mjSOL_CG.value
|
||||
|
||||
# check enum types
|
||||
for mj_type, attrs in _MJ_TYPE_ATTR.items():
|
||||
@@ -116,7 +116,7 @@ def _validate(m: mujoco.MjModel):
|
||||
|
||||
typs = set(val) if isinstance(val, Iterable) else {val}
|
||||
unsupported_typs = typs - set(_TYPE_MAP[mj_type])
|
||||
unsupported = [mj_type(t) for t in unsupported_typs]
|
||||
unsupported = [mj_type(t) for t in unsupported_typs] # pylint: disable=too-many-function-args
|
||||
if unsupported:
|
||||
raise NotImplementedError(f'{unsupported} not implemented.')
|
||||
|
||||
|
||||
@@ -110,11 +110,10 @@ class ValidateInputTest(absltest.TestCase):
|
||||
|
||||
def test_solver(self):
|
||||
m = mujoco.MjModel.from_xml_string(
|
||||
'<mujoco><option solver="Newton"/><worldbody/></mujoco>'
|
||||
'<mujoco><option solver="PGS"/><worldbody/></mujoco>'
|
||||
)
|
||||
with self.assertWarns(UserWarning):
|
||||
mx = mjx.device_put(m)
|
||||
self.assertEqual(mx.opt.solver, mujoco.mjtSolver.mjSOL_CG)
|
||||
with self.assertRaises(NotImplementedError):
|
||||
mjx.device_put(m)
|
||||
|
||||
def test_integrator(self):
|
||||
m = mujoco.MjModel.from_xml_string(
|
||||
|
||||
@@ -37,7 +37,6 @@ from mujoco.mjx._src.types import GainType
|
||||
from mujoco.mjx._src.types import IntegratorType
|
||||
from mujoco.mjx._src.types import JointType
|
||||
from mujoco.mjx._src.types import Model
|
||||
from mujoco.mjx._src.types import SolverType
|
||||
# pylint: enable=g-importing-member
|
||||
import numpy as np
|
||||
|
||||
@@ -333,10 +332,7 @@ def forward(m: Model, d: Data) -> Data:
|
||||
d = d.replace(qacc=d.qacc_smooth)
|
||||
return d
|
||||
|
||||
if m.opt.solver == SolverType.CG:
|
||||
d = named_scope(solver.cg_solve)(m, d)
|
||||
else:
|
||||
raise NotImplementedError(f'solver {m.opt.solver} not implemented.')
|
||||
d = named_scope(solver.solve)(m, d)
|
||||
|
||||
return d
|
||||
|
||||
|
||||
@@ -12,7 +12,7 @@
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# ==============================================================================
|
||||
"""CG and Newton solvers."""
|
||||
"""Constraint solvers."""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
@@ -26,11 +26,12 @@ from mujoco.mjx._src.dataclasses import PyTreeNode
|
||||
from mujoco.mjx._src.types import Data
|
||||
from mujoco.mjx._src.types import DisableBit
|
||||
from mujoco.mjx._src.types import Model
|
||||
from mujoco.mjx._src.types import SolverType
|
||||
# pylint: enable=g-importing-member
|
||||
|
||||
|
||||
class _CGContext(PyTreeNode):
|
||||
"""Data updated during each cg solver iteration.
|
||||
class _Context(PyTreeNode):
|
||||
"""Data updated during each solver iteration.
|
||||
|
||||
Attributes:
|
||||
qacc: acceleration (from Data) (nv,)
|
||||
@@ -44,7 +45,7 @@ class _CGContext(PyTreeNode):
|
||||
search: linesearch vector (nv,)
|
||||
gauss: gauss Cost
|
||||
cost: constraint + Gauss cost
|
||||
prev_cost: cost from previous cg iter
|
||||
prev_cost: cost from previous iter
|
||||
solver_niter: number of solver iterations
|
||||
"""
|
||||
|
||||
@@ -63,13 +64,13 @@ class _CGContext(PyTreeNode):
|
||||
solver_niter: jax.Array
|
||||
|
||||
@classmethod
|
||||
def create(cls, m: Model, d: Data, grad: bool = True) -> '_CGContext':
|
||||
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
|
||||
M = smooth.dense_m(m, d) if m.nv < 100 else None # pylint: disable=invalid-name
|
||||
ma = smooth.mul_m(m, d, d.qacc) if M is None else M @ d.qacc
|
||||
nv_0 = jp.zeros((m.nv,))
|
||||
ctx = _CGContext(
|
||||
ctx = _Context(
|
||||
qacc=d.qacc,
|
||||
qfrc_constraint=d.qfrc_constraint,
|
||||
Jaref=jaref,
|
||||
@@ -84,9 +85,9 @@ class _CGContext(PyTreeNode):
|
||||
prev_cost=0.0,
|
||||
solver_niter=0,
|
||||
)
|
||||
ctx = _cg_update_constraint(m, d, ctx)
|
||||
ctx = _update_constraint(m, d, ctx)
|
||||
if grad:
|
||||
ctx = _cg_update_gradient(m, d, ctx)
|
||||
ctx = _update_gradient(m, d, ctx)
|
||||
ctx = ctx.replace(search=-ctx.Mgrad) # start with preconditioned gradient
|
||||
|
||||
return ctx
|
||||
@@ -111,7 +112,7 @@ class _LSPoint(PyTreeNode):
|
||||
def create(
|
||||
cls,
|
||||
d: Data,
|
||||
ctx: _CGContext,
|
||||
ctx: _Context,
|
||||
alpha: jax.Array,
|
||||
jv: jax.Array,
|
||||
quad: jax.Array,
|
||||
@@ -132,7 +133,7 @@ class _LSPoint(PyTreeNode):
|
||||
|
||||
|
||||
class _LSContext(PyTreeNode):
|
||||
"""Data updated during each cg line search iteration.
|
||||
"""Data updated during each line search iteration.
|
||||
|
||||
Attributes:
|
||||
lo: low point bounding the line search interval
|
||||
@@ -163,15 +164,15 @@ 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 _cg_update_constraint(m: Model, d: Data, ctx: _CGContext) -> _CGContext:
|
||||
"""Updates constraint force and resulting cost given latst CG iteration.
|
||||
def _update_constraint(m: Model, d: Data, ctx: _Context) -> _Context:
|
||||
"""Updates constraint force and resulting cost given latst solver iteration.
|
||||
|
||||
Corresponds to CGupdateConstraint in mujoco/src/engine/engine_solver.c
|
||||
|
||||
Args:
|
||||
m: model defining constraints
|
||||
d: data which contains latest qacc and smooth terms
|
||||
ctx: current CG context
|
||||
ctx: current solver context
|
||||
|
||||
Returns:
|
||||
context with new constraint force and costs
|
||||
@@ -199,22 +200,34 @@ def _cg_update_constraint(m: Model, d: Data, ctx: _CGContext) -> _CGContext:
|
||||
return ctx
|
||||
|
||||
|
||||
def _cg_update_gradient(m: Model, d: Data, ctx: _CGContext) -> _CGContext:
|
||||
"""Updates grad and M / grad given latest CG iteration.
|
||||
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
|
||||
|
||||
Args:
|
||||
m: model defining constraints
|
||||
d: data which contains latest smooth terms
|
||||
ctx: current CG contet
|
||||
ctx: current solver context
|
||||
|
||||
Returns:
|
||||
context with new grad and M / grad
|
||||
Raises:
|
||||
NotImplementedError: for unsupported solver type
|
||||
"""
|
||||
|
||||
grad = ctx.Ma - d.qfrc_smooth - ctx.qfrc_constraint
|
||||
mgrad = smooth.solve_m(m, d, grad)
|
||||
|
||||
if m.opt.solver == SolverType.CG:
|
||||
mgrad = smooth.solve_m(m, d, grad)
|
||||
elif m.opt.solver == SolverType.NEWTON:
|
||||
active = (ctx.Jaref < 0).at[:d.ne + d.nf].set(True)
|
||||
h = (d.efc_J.T * d.efc_D * active) @ d.efc_J
|
||||
h = smooth.dense_m(m, d) + h
|
||||
h_ = jax.scipy.linalg.cho_factor(h)
|
||||
mgrad = jax.scipy.linalg.cho_solve(h_, grad)
|
||||
else:
|
||||
raise NotImplementedError(f"unsupported solver type: {m.opt.solver}")
|
||||
|
||||
ctx = ctx.replace(grad=grad, Mgrad=mgrad)
|
||||
|
||||
@@ -225,13 +238,13 @@ def _rescale(m: Model, value: jax.Array) -> jax.Array:
|
||||
return value / (m.stat.meaninertia * max(1, m.nv))
|
||||
|
||||
|
||||
def _cg_search(m: Model, d: Data, ctx: _CGContext) -> _CGContext:
|
||||
def _linesearch(m: Model, d: Data, ctx: _Context) -> _Context:
|
||||
"""Performs a zoom linesearch to find optimal search step size.
|
||||
|
||||
Args:
|
||||
m: model defining search options and other needed terms
|
||||
d: data with inertia matrix and other needed terms
|
||||
ctx: current CG context
|
||||
ctx: current solver context
|
||||
|
||||
Returns:
|
||||
updated context with new qacc, Ma, Jaref
|
||||
@@ -310,10 +323,10 @@ def _cg_search(m: Model, d: Data, ctx: _CGContext) -> _CGContext:
|
||||
return ctx
|
||||
|
||||
|
||||
def cg_solve(m: Model, d: Data) -> Data:
|
||||
def solve(m: Model, d: Data) -> Data:
|
||||
"""Finds forces that satisfy constraints using conjugate gradient descent."""
|
||||
|
||||
def cond(ctx: _CGContext) -> jax.Array:
|
||||
def cond(ctx: _Context) -> jax.Array:
|
||||
improvement = _rescale(m, ctx.prev_cost - ctx.cost)
|
||||
gradient = _rescale(m, math.norm(ctx.grad))
|
||||
|
||||
@@ -323,11 +336,11 @@ def cg_solve(m: Model, d: Data) -> Data:
|
||||
|
||||
return ~done
|
||||
|
||||
def body(ctx: _CGContext) -> _CGContext:
|
||||
ctx = _cg_search(m, d, ctx)
|
||||
def body(ctx: _Context) -> _Context:
|
||||
ctx = _linesearch(m, d, ctx)
|
||||
prev_grad, prev_Mgrad = ctx.grad, ctx.Mgrad # pylint: disable=invalid-name
|
||||
ctx = _cg_update_constraint(m, d, ctx)
|
||||
ctx = _cg_update_gradient(m, d, ctx)
|
||||
ctx = _update_constraint(m, d, ctx)
|
||||
ctx = _update_gradient(m, d, ctx)
|
||||
|
||||
# polak-ribiere:
|
||||
beta = jp.dot(ctx.grad, ctx.Mgrad - prev_Mgrad)
|
||||
@@ -341,12 +354,16 @@ def cg_solve(m: Model, d: Data) -> Data:
|
||||
# warmstart:
|
||||
qacc = d.qacc_smooth
|
||||
if not m.opt.disableflags & DisableBit.WARMSTART:
|
||||
warm = _CGContext.create(m, d.replace(qacc=d.qacc_warmstart), grad=False)
|
||||
smth = _CGContext.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 = jax.lax.while_loop(cond, body, _CGContext.create(m, d))
|
||||
ctx = _Context.create(m, d)
|
||||
if m.opt.iterations == 1:
|
||||
ctx = body(ctx)
|
||||
else:
|
||||
ctx = jax.lax.while_loop(cond, body, ctx)
|
||||
|
||||
d = d.replace(
|
||||
qacc_warmstart=ctx.qacc,
|
||||
|
||||
@@ -131,8 +131,9 @@ class SolverType(enum.IntEnum):
|
||||
Attributes:
|
||||
CG: Conjugate gradient (primal)
|
||||
"""
|
||||
# unsupported: PGS, NEWTON
|
||||
# unsupported: PGS
|
||||
CG = mujoco.mjtSolver.mjSOL_CG
|
||||
NEWTON = mujoco.mjtSolver.mjSOL_NEWTON
|
||||
|
||||
|
||||
class EqType(enum.IntEnum):
|
||||
|
||||
@@ -33,21 +33,34 @@ _PATHS = {
|
||||
'shadow_hand': 'benchmark/model/shadow_hand/scene_right.xml',
|
||||
}
|
||||
|
||||
|
||||
_BATCH_SIZE = {
|
||||
('humanoid', 'TPU v5 lite'): 1024,
|
||||
('barkour', 'TPU v5 lite'): 1024,
|
||||
('shadow_hand', 'TPU v5 lite'): 1024,
|
||||
('humanoid', 'NVIDIA A100-SXM4-40GB'): 8192,
|
||||
('barkour', 'NVIDIA A100-SXM4-40GB'): 8192,
|
||||
('shadow_hand', 'NVIDIA A100-SXM4-40GB'): 4096,
|
||||
('humanoid', 'cpu'): 64,
|
||||
('barkour', 'tpu_v5e'): 1024,
|
||||
('humanoid', 'tpu_v5e'): 1024,
|
||||
('shadow_hand', 'tpu_v5e'): 1024,
|
||||
('barkour', 'gpu_a100'): 8192,
|
||||
('humanoid', 'gpu_a100'): 8192,
|
||||
('shadow_hand', 'gpu_a100'): 4096,
|
||||
('barkour', 'cpu'): 64,
|
||||
('humanoid', 'cpu'): 64,
|
||||
('shadow_hand', 'cpu'): 64,
|
||||
}
|
||||
|
||||
_SOLVER_CONFIG = {
|
||||
('barkour', 'tpu_v5e'): (mujoco.mjtSolver.mjSOL_CG, 4, 6),
|
||||
('humanoid', 'tpu_v5e'): (mujoco.mjtSolver.mjSOL_CG, 6, 6),
|
||||
('shadow_hand', 'tpu_v5e'): (mujoco.mjtSolver.mjSOL_CG, 8, 6),
|
||||
('humanoid', 'gpu_a100'): (mujoco.mjtSolver.mjSOL_NEWTON, 1, 4),
|
||||
('barkour', 'gpu_a100'): (mujoco.mjtSolver.mjSOL_NEWTON, 1, 4),
|
||||
('shadow_hand', 'gpu_a100'): (mujoco.mjtSolver.mjSOL_NEWTON, 1, 4),
|
||||
('barkour', 'cpu'): (mujoco.mjtSolver.mjSOL_NEWTON, 1, 4),
|
||||
('humanoid', 'cpu'): (mujoco.mjtSolver.mjSOL_NEWTON, 1, 4),
|
||||
('shadow_hand', 'cpu'): (mujoco.mjtSolver.mjSOL_NEWTON, 1, 4),
|
||||
}
|
||||
|
||||
|
||||
flags.DEFINE_string('model', 'humanoid', 'Model to benchmark')
|
||||
flags.DEFINE_string('device', 'cpu', 'Device benchmark is running on')
|
||||
flags.DEFINE_enum('device', 'cpu', ('cpu', 'tpu_v5e', 'gpu_a100'),
|
||||
'Device benchmark is running on')
|
||||
|
||||
|
||||
def _measure_fn(state, init_fn, step_fn, batch_size: int = 1024) -> float:
|
||||
@@ -95,6 +108,9 @@ def _run(state: benchmark.State):
|
||||
|
||||
f = epath.resource_path('mujoco.mjx') / _PATHS[FLAGS.model]
|
||||
m = mujoco.MjModel.from_xml_path(f.as_posix())
|
||||
m.opt.solver, m.opt.iterations, m.opt.ls_iterations = _SOLVER_CONFIG[
|
||||
(FLAGS.model, FLAGS.device)
|
||||
]
|
||||
m = mjx.device_put(m)
|
||||
|
||||
def init(rng):
|
||||
@@ -106,7 +122,7 @@ def _run(state: benchmark.State):
|
||||
def step(d):
|
||||
return mjx.step(m, d)
|
||||
|
||||
batch_size = _BATCH_SIZE[(FLAGS.model, jax.devices()[0].device_kind)]
|
||||
batch_size = _BATCH_SIZE[(FLAGS.model, FLAGS.device)]
|
||||
_measure_fn(state, init, step, batch_size=batch_size)
|
||||
|
||||
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
|
||||
<compiler angle="radian" meshdir="." texturedir="assets" autolimits="true"/>
|
||||
|
||||
<option timestep="0.002" iterations="4" ls_iterations="6" solver="CG">
|
||||
<option timestep="0.002" iterations="1" ls_iterations="4">
|
||||
<flag eulerdamp="disable"/>
|
||||
</option>
|
||||
|
||||
|
||||
@@ -14,7 +14,7 @@
|
||||
-->
|
||||
|
||||
<mujoco model="Humanoid">
|
||||
<option timestep="0.005" solver="CG" iterations="6" ls_iterations="6">
|
||||
<option timestep="0.005" iterations="1" ls_iterations="4">
|
||||
<flag eulerdamp="disable"/>
|
||||
</option>
|
||||
|
||||
@@ -195,7 +195,7 @@
|
||||
<pair geom1="foot2_right" geom2="floor"/>
|
||||
</contact>
|
||||
|
||||
<tendon>
|
||||
<!-- <tendon>
|
||||
<fixed name="hamstring_right" limited="true" range="-0.3 2">
|
||||
<joint joint="hip_y_right" coef=".5"/>
|
||||
<joint joint="knee_right" coef="-.5"/>
|
||||
@@ -204,7 +204,7 @@
|
||||
<joint joint="hip_y_left" coef=".5"/>
|
||||
<joint joint="knee_left" coef="-.5"/>
|
||||
</fixed>
|
||||
</tendon>
|
||||
</tendon> -->
|
||||
|
||||
<actuator>
|
||||
<motor name="abdomen_y" gear="40" joint="abdomen_y"/>
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
<mujoco model="right_shadow_hand">
|
||||
<compiler angle="radian" meshdir="assets" autolimits="true"/>
|
||||
|
||||
<option impratio="10" solver="CG" iterations="8" ls_iterations="6">
|
||||
<option impratio="10" iterations="1" ls_iterations="4">
|
||||
<flag eulerdamp="disable"/>
|
||||
</option>
|
||||
|
||||
|
||||
@@ -5,7 +5,7 @@
|
||||
<geom conaffinity="0" condim="3" contype="0" material="geom"/>
|
||||
<motor ctrllimited="true" ctrlrange="-.4 .4"/>
|
||||
</default>
|
||||
<option iterations="8" timestep="0.003" solver="CG"/>
|
||||
<option iterations="8" timestep="0.003"/>
|
||||
<size nkey="5" nuser_geom="1"/>
|
||||
<visual>
|
||||
<map fogend="5" fogstart="3"/>
|
||||
|
||||
@@ -38,10 +38,6 @@ def main(argv: Sequence[str]) -> None:
|
||||
m = mujoco.MjModel.from_xml_path(_MODEL_PATH.value)
|
||||
d = mujoco.MjData(m)
|
||||
|
||||
# Override the solver option to CG since that is currently the only one
|
||||
# supported by MJX.
|
||||
m.opt.solver = mujoco.mjtSolver.mjSOL_CG
|
||||
|
||||
mx = mjx.device_put(m)
|
||||
dx = mjx.make_data(mx)
|
||||
|
||||
|
||||
+6
-1
@@ -356,6 +356,9 @@
|
||||
" )\n",
|
||||
" mj_model = mujoco.MjModel.from_xml_path(\n",
|
||||
" (path / 'humanoid.xml').as_posix())\n",
|
||||
" mj_model.opt.solver = mujoco.mjtSolver.mjSOL_CG\n",
|
||||
" mj_model.opt.iterations = 6\n",
|
||||
" mj_model.opt.ls_iterations = 6\n",
|
||||
"\n",
|
||||
" physics_steps_per_control_step = 5\n",
|
||||
" kwargs['physics_steps_per_control_step'] = kwargs.get(\n",
|
||||
@@ -834,7 +837,6 @@
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"cellView": "form",
|
||||
"id": "y79PoJOCIl-O"
|
||||
},
|
||||
"outputs": [],
|
||||
@@ -912,6 +914,9 @@
|
||||
" )\n",
|
||||
" mj_model = mujoco.MjModel.from_xml_path(\n",
|
||||
" (path / 'barkour_v0_mjx.xml').as_posix())\n",
|
||||
" mj_model.opt.solver = mujoco.mjtSolver.mjSOL_CG\n",
|
||||
" mj_model.opt.iterations = 4\n",
|
||||
" mj_model.opt.ls_iterations = 6\n",
|
||||
"\n",
|
||||
" physics_steps_per_control_step = 10\n",
|
||||
" kwargs['physics_steps_per_control_step'] = kwargs.get(\n",
|
||||
|
||||
Reference in New Issue
Block a user