Adds support for Newton solver to MJX.

PiperOrigin-RevId: 578738033
Change-Id: Id0e68c7b961380cb1ed88b40d99f202c3fa8e228
This commit is contained in:
Erik Frey
2023-11-01 22:04:02 -07:00
committed by Copybara-Service
parent e82cb9f420
commit 3c0a56c1e5
14 changed files with 128 additions and 74 deletions
+28 -3
View File
@@ -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
View File
@@ -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
+5 -5
View File
@@ -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.')
+3 -4
View File
@@ -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(
+1 -5
View File
@@ -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
+45 -28
View File
@@ -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,
+2 -1
View File
@@ -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):
+26 -10
View File
@@ -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>
+1 -1
View File
@@ -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"/>
-4
View File
@@ -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
View File
@@ -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",