diff --git a/doc/changelog.rst b/doc/changelog.rst index 7a12dce8..d8e1d91a 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -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 `__ + - 640,000 + - 1,020,000 + - **1.6 x** + * - `Barkour v0 `__ + - 1,290,000 + - 1,750,000 + - **1.35 x** + * - `Shadow Hand `__ + - 215,000 + - 270,000 + - **1.25 x** + + Humanoid is the standard MuJoCo humanoid, + `Google Barkour `__ and the Shadow Hand + are both available in the :ref:`MuJoCo 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 ^^^^^^^^^^^^^^^ diff --git a/doc/mjx.rst b/doc/mjx.rst index 7d1b937a..481da5c8 100644 --- a/doc/mjx.rst +++ b/doc/mjx.rst @@ -198,7 +198,7 @@ The following features are **fully supported** in MJX: * - :ref:`Condim ` - 3 * - :ref:`Solver ` - - ``CG`` + - ``CG``, ``NEWTON`` * - Fluid Model - :ref:`flInertia` @@ -226,8 +226,6 @@ The following features are **in development** and coming soon: - ``ELLIPTIC`` * - :ref:`Condim ` - 1, 4, 6 - * - :ref:`Solver ` - - ``NEWTON`` * - Fluid Model - :ref:`flEllipsoid` * - :ref:`Tendons ` @@ -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 ` 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 diff --git a/mjx/mujoco/mjx/_src/device.py b/mjx/mujoco/mjx/_src/device.py index 072755ea..c39860ef 100644 --- a/mjx/mujoco/mjx/_src/device.py +++ b/mjx/mujoco/mjx/_src/device.py @@ -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.') diff --git a/mjx/mujoco/mjx/_src/device_test.py b/mjx/mujoco/mjx/_src/device_test.py index e8aa09b9..e6eb7562 100644 --- a/mjx/mujoco/mjx/_src/device_test.py +++ b/mjx/mujoco/mjx/_src/device_test.py @@ -110,11 +110,10 @@ class ValidateInputTest(absltest.TestCase): def test_solver(self): m = mujoco.MjModel.from_xml_string( - '' + '' ) - 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( diff --git a/mjx/mujoco/mjx/_src/forward.py b/mjx/mujoco/mjx/_src/forward.py index 42faeca2..06bec894 100644 --- a/mjx/mujoco/mjx/_src/forward.py +++ b/mjx/mujoco/mjx/_src/forward.py @@ -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 diff --git a/mjx/mujoco/mjx/_src/solver.py b/mjx/mujoco/mjx/_src/solver.py index 998df547..567cb709 100644 --- a/mjx/mujoco/mjx/_src/solver.py +++ b/mjx/mujoco/mjx/_src/solver.py @@ -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, diff --git a/mjx/mujoco/mjx/_src/types.py b/mjx/mujoco/mjx/_src/types.py index d43e36b1..5875ad7b 100644 --- a/mjx/mujoco/mjx/_src/types.py +++ b/mjx/mujoco/mjx/_src/types.py @@ -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): diff --git a/mjx/mujoco/mjx/benchmark/benchmark.py b/mjx/mujoco/mjx/benchmark/benchmark.py index 61112953..c3145f57 100644 --- a/mjx/mujoco/mjx/benchmark/benchmark.py +++ b/mjx/mujoco/mjx/benchmark/benchmark.py @@ -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) diff --git a/mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/barkour_v0_mjx.xml b/mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/barkour_v0_mjx.xml index 692780dc..a98db27b 100644 --- a/mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/barkour_v0_mjx.xml +++ b/mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/barkour_v0_mjx.xml @@ -2,7 +2,7 @@ - diff --git a/mjx/mujoco/mjx/benchmark/model/humanoid/humanoid.xml b/mjx/mujoco/mjx/benchmark/model/humanoid/humanoid.xml index 0de6f774..0851b8ec 100644 --- a/mjx/mujoco/mjx/benchmark/model/humanoid/humanoid.xml +++ b/mjx/mujoco/mjx/benchmark/model/humanoid/humanoid.xml @@ -14,7 +14,7 @@ --> - @@ -195,7 +195,7 @@ - + diff --git a/mjx/mujoco/mjx/benchmark/model/shadow_hand/right_hand.xml b/mjx/mujoco/mjx/benchmark/model/shadow_hand/right_hand.xml index 60139cbf..2ec862e5 100644 --- a/mjx/mujoco/mjx/benchmark/model/shadow_hand/right_hand.xml +++ b/mjx/mujoco/mjx/benchmark/model/shadow_hand/right_hand.xml @@ -1,7 +1,7 @@ - diff --git a/mjx/mujoco/mjx/test_data/humanoid.xml b/mjx/mujoco/mjx/test_data/humanoid.xml index 545204a1..2d7158ee 100644 --- a/mjx/mujoco/mjx/test_data/humanoid.xml +++ b/mjx/mujoco/mjx/test_data/humanoid.xml @@ -5,7 +5,7 @@ -