From e59605135a7aabed9c8a277bdf8396364015bcce Mon Sep 17 00:00:00 2001 From: Erik Frey Date: Wed, 24 Jan 2024 17:03:13 -0800 Subject: [PATCH] Make benchmarks and viewer more accessible. - Drop google_benchmark module dependency, wasn't really using it. - Move benchmark models into test_data - Rename benchmark => testspeed, make default output more like MJ testspeed. - Update viewer to use new io functions and add get_data_into. - Add "mjx-testspeed" and "mjx-viewer" bin scripts PiperOrigin-RevId: 601277135 Change-Id: I782f666c58a96ac1a0a1c5ffc8ecf29d9fcd810f --- mjx/MANIFEST.in | 1 - mjx/mujoco/mjx/__init__.py | 1 + mjx/mujoco/mjx/_src/io.py | 59 ++++++--- mjx/mujoco/mjx/_src/io_test.py | 21 +++ mjx/mujoco/mjx/benchmark/__init__.py | 14 -- mjx/mujoco/mjx/benchmark/benchmark.py | 109 ---------------- .../model => test_data}/barkour_v0/README.md | 0 .../barkour_v0/assets/abduction.stl | Bin .../barkour_v0/assets/barkour_v0_mjx.xml | 0 .../barkour_v0/assets/body.stl | Bin .../barkour_v0/assets/foot.stl | Bin .../barkour_v0/assets/handle.stl | Bin .../barkour_v0/assets/head.stl | Bin .../barkour_v0/assets/head_mount.stl | Bin .../barkour_v0/assets/joystick.gif | Bin .../barkour_v0/assets/lower_leg_1to1.stl | Bin .../barkour_v0/assets/powercable.stl | Bin .../barkour_v0/assets/upper_left_1.stl | Bin .../barkour_v0/assets/upper_left_2.stl | Bin .../barkour_v0/assets/upper_left_3.stl | Bin .../barkour_v0/assets/upper_right_1.stl | Bin .../barkour_v0/assets/upper_right_2.stl | Bin .../barkour_v0/assets/upper_right_3.stl | Bin .../model => test_data}/humanoid/README.md | 0 .../model => test_data}/humanoid/humanoid.png | Bin .../model => test_data}/humanoid/humanoid.xml | 0 .../model => test_data}/shadow_hand/README.md | 0 .../shadow_hand/assets/f_distal_pst.obj | 0 .../shadow_hand/assets/f_knuckle.obj | 0 .../shadow_hand/assets/f_middle.obj | 0 .../shadow_hand/assets/f_proximal.obj | 0 .../shadow_hand/assets/forearm_0.obj | 0 .../shadow_hand/assets/forearm_1.obj | 0 .../shadow_hand/assets/forearm_collision.obj | 0 .../shadow_hand/assets/lf_metacarpal.obj | 0 .../shadow_hand/assets/mounting_plate.obj | 0 .../shadow_hand/assets/palm.obj | 0 .../shadow_hand/assets/th_distal_pst.obj | 0 .../shadow_hand/assets/th_middle.obj | 0 .../shadow_hand/assets/th_proximal.obj | 0 .../shadow_hand/assets/wrist.obj | 0 .../shadow_hand/right_hand.xml | 0 .../shadow_hand/scene_right.xml | 0 .../shadow_hand/shadow_hand.png | Bin mjx/mujoco/mjx/testspeed.py | 123 ++++++++++++++++++ mjx/mujoco/mjx/viewer.py | 30 ++--- mjx/pyproject.toml | 4 + 47 files changed, 204 insertions(+), 158 deletions(-) delete mode 100644 mjx/mujoco/mjx/benchmark/__init__.py delete mode 100644 mjx/mujoco/mjx/benchmark/benchmark.py rename mjx/mujoco/mjx/{benchmark/model => test_data}/barkour_v0/README.md (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/barkour_v0/assets/abduction.stl (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/barkour_v0/assets/barkour_v0_mjx.xml (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/barkour_v0/assets/body.stl (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/barkour_v0/assets/foot.stl (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/barkour_v0/assets/handle.stl (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/barkour_v0/assets/head.stl (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/barkour_v0/assets/head_mount.stl (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/barkour_v0/assets/joystick.gif (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/barkour_v0/assets/lower_leg_1to1.stl (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/barkour_v0/assets/powercable.stl (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/barkour_v0/assets/upper_left_1.stl (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/barkour_v0/assets/upper_left_2.stl (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/barkour_v0/assets/upper_left_3.stl (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/barkour_v0/assets/upper_right_1.stl (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/barkour_v0/assets/upper_right_2.stl (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/barkour_v0/assets/upper_right_3.stl (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/humanoid/README.md (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/humanoid/humanoid.png (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/humanoid/humanoid.xml (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/shadow_hand/README.md (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/shadow_hand/assets/f_distal_pst.obj (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/shadow_hand/assets/f_knuckle.obj (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/shadow_hand/assets/f_middle.obj (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/shadow_hand/assets/f_proximal.obj (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/shadow_hand/assets/forearm_0.obj (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/shadow_hand/assets/forearm_1.obj (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/shadow_hand/assets/forearm_collision.obj (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/shadow_hand/assets/lf_metacarpal.obj (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/shadow_hand/assets/mounting_plate.obj (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/shadow_hand/assets/palm.obj (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/shadow_hand/assets/th_distal_pst.obj (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/shadow_hand/assets/th_middle.obj (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/shadow_hand/assets/th_proximal.obj (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/shadow_hand/assets/wrist.obj (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/shadow_hand/right_hand.xml (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/shadow_hand/scene_right.xml (100%) rename mjx/mujoco/mjx/{benchmark/model => test_data}/shadow_hand/shadow_hand.png (100%) create mode 100644 mjx/mujoco/mjx/testspeed.py diff --git a/mjx/MANIFEST.in b/mjx/MANIFEST.in index b963ce6e..4bc2067d 100644 --- a/mjx/MANIFEST.in +++ b/mjx/MANIFEST.in @@ -1,2 +1 @@ recursive-include mujoco/mjx/test_data * -recursive-include mujoco/mjx/benchmark *.obj *.stl *.xml diff --git a/mjx/mujoco/mjx/__init__.py b/mjx/mujoco/mjx/__init__.py index 737fe2c5..8656e618 100644 --- a/mjx/mujoco/mjx/__init__.py +++ b/mjx/mujoco/mjx/__init__.py @@ -29,6 +29,7 @@ from mujoco.mjx._src.forward import fwd_velocity from mujoco.mjx._src.forward import rungekutta4 from mujoco.mjx._src.forward import step 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 from mujoco.mjx._src.io import put_data from mujoco.mjx._src.io import put_model diff --git a/mjx/mujoco/mjx/_src/io.py b/mjx/mujoco/mjx/_src/io.py index dd91de94..095c4eae 100644 --- a/mjx/mujoco/mjx/_src/io.py +++ b/mjx/mujoco/mjx/_src/io.py @@ -233,8 +233,33 @@ def get_data( m: mujoco.MjModel, d: types.Data ) -> Union[mujoco.MjData, List[mujoco.MjData]]: """Gets mjx.Data from a device, resulting in mujoco.MjData or List[MjData].""" - dx = jax.device_get(d) batched = len(d.qpos.shape) > 1 + batch_size = d.qpos.shape[0] if batched else 1 + + if batched: + result = [mujoco.MjData(m) for _ in range(batch_size)] + else: + result = mujoco.MjData(m) + + get_data_into(result, m, d) + + return result + + +def get_data_into( + result: Union[mujoco.MjData, List[mujoco.MjData]], + m: mujoco.MjModel, + d: types.Data, +): + """Gets mjx.Data from a device into an existing mujoco.MjData or list.""" + batched = isinstance(result, list) + if batched and len(d.qpos.shape) < 2: + raise ValueError('dst is a list, but d is not batched.') + if not batched and len(d.qpos.shape) >= 2: + raise ValueError('dst is a an MjData, but d is batched.') + + d = jax.device_get(d) + batch_size = d.qpos.shape[0] if batched else 1 ne, nf, nl, nc = constraint.count_constraints(m) efc_type = np.array([ @@ -252,26 +277,25 @@ def get_data( dof_j.append(j) j = m.dof_parentid[j] - ds = [] for i in range(batch_size): - dx_i = jax.tree_map(lambda x, i=i: x[i], dx) if batched else d - ncon = (dx_i.contact.dist <= 0).sum() - efc_active = (dx_i.efc_J != 0).any(axis=1) + d_i = jax.tree_map(lambda x, i=i: x[i], d) if batched else d + result_i = result[i] if batched else result + ncon = (d_i.contact.dist <= 0).sum() + efc_active = (d_i.efc_J != 0).any(axis=1) efc_con = efc_type == mujoco.mjtConstraint.mjCNSTR_CONTACT_PYRAMIDAL nefc, nc = efc_active.sum(), (efc_active & efc_con).sum() - d_i = mujoco.MjData(m) - d_i.nnzJ = nefc * m.nv - mujoco._functions._realloc_con_efc(d_i, ncon=ncon, nefc=nefc) # pylint: disable=protected-access - d_i.efc_J_rownnz[:] = np.repeat(m.nv, nefc) - d_i.efc_J_rowadr[:] = np.arange(0, nefc * m.nv, m.nv) - d_i.efc_J_colind[:] = np.tile(np.arange(m.nv), nefc) + result_i.nnzJ = nefc * m.nv + mujoco._functions._realloc_con_efc(result_i, ncon=ncon, nefc=nefc) # pylint: disable=protected-access + result_i.efc_J_rownnz[:] = np.repeat(m.nv, nefc) + result_i.efc_J_rowadr[:] = np.arange(0, nefc * m.nv, m.nv) + result_i.efc_J_colind[:] = np.tile(np.arange(m.nv), nefc) for field in types.Data.fields(): if field.name == 'contact': - _get_contact(d_i.contact, dx_i.contact, nefc - nc) + _get_contact(result_i.contact, d_i.contact, nefc - nc) continue - value = getattr(dx_i, field.name) + value = getattr(d_i, field.name) if field.name in ('xmat', 'ximat', 'geom_xmat', 'site_xmat'): value = value.reshape((-1, 9)) @@ -292,14 +316,11 @@ def get_data( value = np.ones(m.nv) if value.shape: - getattr(d_i, field.name)[:] = value + getattr(result_i, field.name)[:] = value else: - setattr(d_i, field.name, value) + setattr(result_i, field.name, value) - d_i.efc_type[:] = efc_type[efc_active] - ds.append(d_i) - - return ds if batched else ds[0] + result_i.efc_type[:] = efc_type[efc_active] def _put_contact( diff --git a/mjx/mujoco/mjx/_src/io_test.py b/mjx/mujoco/mjx/_src/io_test.py index b402c407..7e71db5d 100644 --- a/mjx/mujoco/mjx/_src/io_test.py +++ b/mjx/mujoco/mjx/_src/io_test.py @@ -412,6 +412,27 @@ class DataIOTest(parameterized.TestCase): self.assertEqual(ds[0].ncon, 1) self.assertEqual(ds[1].ncon, 0) + def test_get_data_into(self): + """Test that get_data_into correctly populates an MjData.""" + + m = mujoco.MjModel.from_xml_string(_MULTIPLE_CONSTRAINTS) + d = mujoco.MjData(m) + mujoco.mj_step(m, d, 2) + dx = mjx.put_data(m, d) + d_2 = mujoco.MjData(m) + mjx.get_data_into(d_2, m, dx) + + # check a few fields + np.testing.assert_allclose(d_2.qpos, d.qpos) + np.testing.assert_allclose(d_2.xpos, d.xpos) + np.testing.assert_allclose(d_2.qM, d.qM) + + # only 1 contact active + self.assertEqual(d_2.contact.dist.shape, (1,)) + self.assertEqual(d_2.ncon, 1) + np.testing.assert_allclose(d_2.contact.dist, d.contact.dist) + self.assertEqual(d_2.contact.frame.shape, (1, 9)) + np.testing.assert_allclose(d_2.contact.frame, d.contact.frame) if __name__ == '__main__': absltest.main() diff --git a/mjx/mujoco/mjx/benchmark/__init__.py b/mjx/mujoco/mjx/benchmark/__init__.py deleted file mode 100644 index 4ad51d69..00000000 --- a/mjx/mujoco/mjx/benchmark/__init__.py +++ /dev/null @@ -1,14 +0,0 @@ -# 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. -# ============================================================================== diff --git a/mjx/mujoco/mjx/benchmark/benchmark.py b/mjx/mujoco/mjx/benchmark/benchmark.py deleted file mode 100644 index ca29a7e1..00000000 --- a/mjx/mujoco/mjx/benchmark/benchmark.py +++ /dev/null @@ -1,109 +0,0 @@ -# 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. -# ============================================================================== -"""Run benchmarks on various devices.""" - -import sys -import time - -from absl import flags -from etils import epath -import google_benchmark as benchmark -import jax -from jax import numpy as jp -import mujoco -from mujoco import mjx - -FLAGS = flags.FLAGS - -flags.DEFINE_string('mjcf', None, 'path to model', required=True) -flags.DEFINE_integer('step_count', 1000, 'number of steps per rollout') -flags.DEFINE_integer('batch_size', 1024, 'number of parallel rollouts') -flags.DEFINE_integer('unroll', 1, 'loop unroll length') -flags.DEFINE_enum('solver', 'cg', ['cg', 'newton'], 'constraint solver') -flags.DEFINE_integer('iterations', 1, 'number of solver iterations') -flags.DEFINE_integer('ls_iterations', 4, 'number of linesearch iterations') - - -def _measure(state, init_fn, step_fn) -> float: - """Reports jit time and op time for a function.""" - - @jax.pmap - def run_batch(seed: jp.ndarray): - batch_size = FLAGS.batch_size // jax.device_count() - rngs = jax.random.split(jax.random.PRNGKey(seed), batch_size) - state = jax.vmap(init_fn)(rngs) - - @jax.vmap - def step(state, _): - state = step_fn(state) - return state, None - - state, _ = jax.lax.scan( - step, state, None, length=FLAGS.step_count, unroll=FLAGS.unroll - ) - return state - - # run once to jit - beg = time.perf_counter() - seed = 0 - seeds = jp.arange(seed, seed + jax.device_count(), dtype=int) - jax.tree_util.tree_map(lambda x: x.block_until_ready(), run_batch(seeds)) - first_t = time.perf_counter() - beg - - times = [] - while state: - seed += jax.device_count() - seeds = jp.arange(seed, seed + jax.device_count(), dtype=int) - beg = time.perf_counter() - jax.tree_util.tree_map(lambda x: x.block_until_ready(), run_batch(seeds)) - times.append(time.perf_counter() - beg) - - op_time = jp.mean(jp.array(times)) - batch_sps = FLAGS.batch_size * FLAGS.step_count / op_time - - state.counters['jit_time'] = first_t - op_time - state.counters['batch_sps'] = batch_sps - - -@benchmark.option.unit(benchmark.kSecond) -def _run(state: benchmark.State): - """Benchmark a model.""" - - f = epath.resource_path('mujoco.mjx') / 'benchmark/model' / FLAGS.mjcf - m = mujoco.MjModel.from_xml_path(f.as_posix()) - m.opt.solver = { - 'cg': mujoco.mjtSolver.mjSOL_CG, - 'newton': mujoco.mjtSolver.mjSOL_NEWTON, - }[FLAGS.solver.lower()] - m.opt.iterations = FLAGS.iterations - m.opt.ls_iterations = FLAGS.ls_iterations - m = mjx.device_put(m) - - def init(rng): - d = mjx.make_data(m) - qvel = 0.01 * jax.random.normal(rng, shape=(m.nv,)) - d = d.replace(qvel=qvel) - return d - - def step(d): - return mjx.step(m, d) - - _measure(state, init, step) - - -if __name__ == '__main__': - FLAGS(sys.argv) - benchmark.register(_run, name=sys.argv[0].split('/')[-1]) - benchmark.main() diff --git a/mjx/mujoco/mjx/benchmark/model/barkour_v0/README.md b/mjx/mujoco/mjx/test_data/barkour_v0/README.md similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/barkour_v0/README.md rename to mjx/mujoco/mjx/test_data/barkour_v0/README.md diff --git a/mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/abduction.stl b/mjx/mujoco/mjx/test_data/barkour_v0/assets/abduction.stl similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/abduction.stl rename to mjx/mujoco/mjx/test_data/barkour_v0/assets/abduction.stl diff --git a/mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/barkour_v0_mjx.xml b/mjx/mujoco/mjx/test_data/barkour_v0/assets/barkour_v0_mjx.xml similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/barkour_v0_mjx.xml rename to mjx/mujoco/mjx/test_data/barkour_v0/assets/barkour_v0_mjx.xml diff --git a/mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/body.stl b/mjx/mujoco/mjx/test_data/barkour_v0/assets/body.stl similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/body.stl rename to mjx/mujoco/mjx/test_data/barkour_v0/assets/body.stl diff --git a/mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/foot.stl b/mjx/mujoco/mjx/test_data/barkour_v0/assets/foot.stl similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/foot.stl rename to mjx/mujoco/mjx/test_data/barkour_v0/assets/foot.stl diff --git a/mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/handle.stl b/mjx/mujoco/mjx/test_data/barkour_v0/assets/handle.stl similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/handle.stl rename to mjx/mujoco/mjx/test_data/barkour_v0/assets/handle.stl diff --git a/mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/head.stl b/mjx/mujoco/mjx/test_data/barkour_v0/assets/head.stl similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/head.stl rename to mjx/mujoco/mjx/test_data/barkour_v0/assets/head.stl diff --git a/mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/head_mount.stl b/mjx/mujoco/mjx/test_data/barkour_v0/assets/head_mount.stl similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/head_mount.stl rename to mjx/mujoco/mjx/test_data/barkour_v0/assets/head_mount.stl diff --git a/mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/joystick.gif b/mjx/mujoco/mjx/test_data/barkour_v0/assets/joystick.gif similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/joystick.gif rename to mjx/mujoco/mjx/test_data/barkour_v0/assets/joystick.gif diff --git a/mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/lower_leg_1to1.stl b/mjx/mujoco/mjx/test_data/barkour_v0/assets/lower_leg_1to1.stl similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/lower_leg_1to1.stl rename to mjx/mujoco/mjx/test_data/barkour_v0/assets/lower_leg_1to1.stl diff --git a/mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/powercable.stl b/mjx/mujoco/mjx/test_data/barkour_v0/assets/powercable.stl similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/powercable.stl rename to mjx/mujoco/mjx/test_data/barkour_v0/assets/powercable.stl diff --git a/mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/upper_left_1.stl b/mjx/mujoco/mjx/test_data/barkour_v0/assets/upper_left_1.stl similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/upper_left_1.stl rename to mjx/mujoco/mjx/test_data/barkour_v0/assets/upper_left_1.stl diff --git a/mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/upper_left_2.stl b/mjx/mujoco/mjx/test_data/barkour_v0/assets/upper_left_2.stl similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/upper_left_2.stl rename to mjx/mujoco/mjx/test_data/barkour_v0/assets/upper_left_2.stl diff --git a/mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/upper_left_3.stl b/mjx/mujoco/mjx/test_data/barkour_v0/assets/upper_left_3.stl similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/upper_left_3.stl rename to mjx/mujoco/mjx/test_data/barkour_v0/assets/upper_left_3.stl diff --git a/mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/upper_right_1.stl b/mjx/mujoco/mjx/test_data/barkour_v0/assets/upper_right_1.stl similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/upper_right_1.stl rename to mjx/mujoco/mjx/test_data/barkour_v0/assets/upper_right_1.stl diff --git a/mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/upper_right_2.stl b/mjx/mujoco/mjx/test_data/barkour_v0/assets/upper_right_2.stl similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/upper_right_2.stl rename to mjx/mujoco/mjx/test_data/barkour_v0/assets/upper_right_2.stl diff --git a/mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/upper_right_3.stl b/mjx/mujoco/mjx/test_data/barkour_v0/assets/upper_right_3.stl similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/barkour_v0/assets/upper_right_3.stl rename to mjx/mujoco/mjx/test_data/barkour_v0/assets/upper_right_3.stl diff --git a/mjx/mujoco/mjx/benchmark/model/humanoid/README.md b/mjx/mujoco/mjx/test_data/humanoid/README.md similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/humanoid/README.md rename to mjx/mujoco/mjx/test_data/humanoid/README.md diff --git a/mjx/mujoco/mjx/benchmark/model/humanoid/humanoid.png b/mjx/mujoco/mjx/test_data/humanoid/humanoid.png similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/humanoid/humanoid.png rename to mjx/mujoco/mjx/test_data/humanoid/humanoid.png diff --git a/mjx/mujoco/mjx/benchmark/model/humanoid/humanoid.xml b/mjx/mujoco/mjx/test_data/humanoid/humanoid.xml similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/humanoid/humanoid.xml rename to mjx/mujoco/mjx/test_data/humanoid/humanoid.xml diff --git a/mjx/mujoco/mjx/benchmark/model/shadow_hand/README.md b/mjx/mujoco/mjx/test_data/shadow_hand/README.md similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/shadow_hand/README.md rename to mjx/mujoco/mjx/test_data/shadow_hand/README.md diff --git a/mjx/mujoco/mjx/benchmark/model/shadow_hand/assets/f_distal_pst.obj b/mjx/mujoco/mjx/test_data/shadow_hand/assets/f_distal_pst.obj similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/shadow_hand/assets/f_distal_pst.obj rename to mjx/mujoco/mjx/test_data/shadow_hand/assets/f_distal_pst.obj diff --git a/mjx/mujoco/mjx/benchmark/model/shadow_hand/assets/f_knuckle.obj b/mjx/mujoco/mjx/test_data/shadow_hand/assets/f_knuckle.obj similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/shadow_hand/assets/f_knuckle.obj rename to mjx/mujoco/mjx/test_data/shadow_hand/assets/f_knuckle.obj diff --git a/mjx/mujoco/mjx/benchmark/model/shadow_hand/assets/f_middle.obj b/mjx/mujoco/mjx/test_data/shadow_hand/assets/f_middle.obj similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/shadow_hand/assets/f_middle.obj rename to mjx/mujoco/mjx/test_data/shadow_hand/assets/f_middle.obj diff --git a/mjx/mujoco/mjx/benchmark/model/shadow_hand/assets/f_proximal.obj b/mjx/mujoco/mjx/test_data/shadow_hand/assets/f_proximal.obj similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/shadow_hand/assets/f_proximal.obj rename to mjx/mujoco/mjx/test_data/shadow_hand/assets/f_proximal.obj diff --git a/mjx/mujoco/mjx/benchmark/model/shadow_hand/assets/forearm_0.obj b/mjx/mujoco/mjx/test_data/shadow_hand/assets/forearm_0.obj similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/shadow_hand/assets/forearm_0.obj rename to mjx/mujoco/mjx/test_data/shadow_hand/assets/forearm_0.obj diff --git a/mjx/mujoco/mjx/benchmark/model/shadow_hand/assets/forearm_1.obj b/mjx/mujoco/mjx/test_data/shadow_hand/assets/forearm_1.obj similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/shadow_hand/assets/forearm_1.obj rename to mjx/mujoco/mjx/test_data/shadow_hand/assets/forearm_1.obj diff --git a/mjx/mujoco/mjx/benchmark/model/shadow_hand/assets/forearm_collision.obj b/mjx/mujoco/mjx/test_data/shadow_hand/assets/forearm_collision.obj similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/shadow_hand/assets/forearm_collision.obj rename to mjx/mujoco/mjx/test_data/shadow_hand/assets/forearm_collision.obj diff --git a/mjx/mujoco/mjx/benchmark/model/shadow_hand/assets/lf_metacarpal.obj b/mjx/mujoco/mjx/test_data/shadow_hand/assets/lf_metacarpal.obj similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/shadow_hand/assets/lf_metacarpal.obj rename to mjx/mujoco/mjx/test_data/shadow_hand/assets/lf_metacarpal.obj diff --git a/mjx/mujoco/mjx/benchmark/model/shadow_hand/assets/mounting_plate.obj b/mjx/mujoco/mjx/test_data/shadow_hand/assets/mounting_plate.obj similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/shadow_hand/assets/mounting_plate.obj rename to mjx/mujoco/mjx/test_data/shadow_hand/assets/mounting_plate.obj diff --git a/mjx/mujoco/mjx/benchmark/model/shadow_hand/assets/palm.obj b/mjx/mujoco/mjx/test_data/shadow_hand/assets/palm.obj similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/shadow_hand/assets/palm.obj rename to mjx/mujoco/mjx/test_data/shadow_hand/assets/palm.obj diff --git a/mjx/mujoco/mjx/benchmark/model/shadow_hand/assets/th_distal_pst.obj b/mjx/mujoco/mjx/test_data/shadow_hand/assets/th_distal_pst.obj similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/shadow_hand/assets/th_distal_pst.obj rename to mjx/mujoco/mjx/test_data/shadow_hand/assets/th_distal_pst.obj diff --git a/mjx/mujoco/mjx/benchmark/model/shadow_hand/assets/th_middle.obj b/mjx/mujoco/mjx/test_data/shadow_hand/assets/th_middle.obj similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/shadow_hand/assets/th_middle.obj rename to mjx/mujoco/mjx/test_data/shadow_hand/assets/th_middle.obj diff --git a/mjx/mujoco/mjx/benchmark/model/shadow_hand/assets/th_proximal.obj b/mjx/mujoco/mjx/test_data/shadow_hand/assets/th_proximal.obj similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/shadow_hand/assets/th_proximal.obj rename to mjx/mujoco/mjx/test_data/shadow_hand/assets/th_proximal.obj diff --git a/mjx/mujoco/mjx/benchmark/model/shadow_hand/assets/wrist.obj b/mjx/mujoco/mjx/test_data/shadow_hand/assets/wrist.obj similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/shadow_hand/assets/wrist.obj rename to mjx/mujoco/mjx/test_data/shadow_hand/assets/wrist.obj diff --git a/mjx/mujoco/mjx/benchmark/model/shadow_hand/right_hand.xml b/mjx/mujoco/mjx/test_data/shadow_hand/right_hand.xml similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/shadow_hand/right_hand.xml rename to mjx/mujoco/mjx/test_data/shadow_hand/right_hand.xml diff --git a/mjx/mujoco/mjx/benchmark/model/shadow_hand/scene_right.xml b/mjx/mujoco/mjx/test_data/shadow_hand/scene_right.xml similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/shadow_hand/scene_right.xml rename to mjx/mujoco/mjx/test_data/shadow_hand/scene_right.xml diff --git a/mjx/mujoco/mjx/benchmark/model/shadow_hand/shadow_hand.png b/mjx/mujoco/mjx/test_data/shadow_hand/shadow_hand.png similarity index 100% rename from mjx/mujoco/mjx/benchmark/model/shadow_hand/shadow_hand.png rename to mjx/mujoco/mjx/test_data/shadow_hand/shadow_hand.png diff --git a/mjx/mujoco/mjx/testspeed.py b/mjx/mujoco/mjx/testspeed.py new file mode 100644 index 00000000..66937a8b --- /dev/null +++ b/mjx/mujoco/mjx/testspeed.py @@ -0,0 +1,123 @@ +# 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. +# ============================================================================== +"""Run benchmarks on various devices.""" + +import time +from typing import Sequence, Tuple + +from absl import app +from absl import flags +from etils import epath +import jax +import mujoco +from mujoco import mjx + +FLAGS = flags.FLAGS + +flags.DEFINE_string('mjcf', None, 'path to model', required=True) +flags.DEFINE_integer('nstep', 1000, 'number of steps per rollout') +flags.DEFINE_integer('batch_size', 1024, 'number of parallel rollouts') +flags.DEFINE_integer('unroll', 1, 'loop unroll length') +flags.DEFINE_enum('solver', 'cg', ['cg', 'newton'], 'constraint solver') +flags.DEFINE_integer('iterations', 1, 'number of solver iterations') +flags.DEFINE_integer('ls_iterations', 4, 'number of linesearch iterations') +flags.DEFINE_enum('output', 'text', ['text', 'tsv'], 'format to print results') + + +def _measure(fn, *args) -> Tuple[float, float]: + """Reports jit time and op time for a function.""" + + beg = time.perf_counter() + compiled_fn = fn.lower(*args).compile() + end = time.perf_counter() + jit_time = end - beg + + beg = time.perf_counter() + result = compiled_fn(*args) + jax.block_until_ready(result) + end = time.perf_counter() + run_time = end - beg + + return jit_time, run_time + + +def _main(argv: Sequence[str]): + """Benchmark a model.""" + + f = epath.resource_path('mujoco.mjx') / 'test_data' / FLAGS.mjcf + m = mujoco.MjModel.from_xml_path(f.as_posix()) + m.opt.solver = { + 'cg': mujoco.mjtSolver.mjSOL_CG, + 'newton': mujoco.mjtSolver.mjSOL_NEWTON, + }[FLAGS.solver.lower()] + m.opt.iterations = FLAGS.iterations + m.opt.ls_iterations = FLAGS.ls_iterations + m = mjx.put_model(m) + + if FLAGS.output == 'text': + print(f"Rolling out {FLAGS.nstep} steps at dt = {m.opt.timestep:.3f}...") + + @jax.pmap + def init(key): + key = jax.random.split(key, FLAGS.batch_size // jax.device_count()) + + @jax.vmap + def random_init(key): + d = mjx.make_data(m) + qvel = 0.01 * jax.random.normal(key, shape=(m.nv,)) + d = d.replace(qvel=qvel) + return d + + return random_init(key) + + key = jax.random.split(jax.random.key(0), jax.device_count()) + d = init(key) + jax.block_until_ready(d) + + @jax.pmap + def unroll(d): + + @jax.vmap + def step(d, _): + d = mjx.step(m, d) + return d, None + + d, _ = jax.lax.scan(step, d, None, length=FLAGS.nstep, unroll=FLAGS.unroll) + + return d + + jit_time, run_time = _measure(unroll, d) + steps = FLAGS.nstep * FLAGS.batch_size + + if FLAGS.output == 'text': + print(f""" +Summary for {FLAGS.batch_size} parallel rollouts + + Total JIT time: {jit_time:.2f} s + Total simulation time: {run_time:.2f} s + Total steps per second: { steps / run_time:.0f} + Total realtime factor: { steps * m.opt.timestep / run_time:.2f} x + Total time per step: { 1e6 * run_time / steps:.2f} µs""") + elif FLAGS.output == 'tsv': + name = argv[0].split('/')[-1].replace('testspeed_', '') + print(f"{name}\tjit: {jit_time:.2f}s\tsteps/second: {steps / run_time:.0f}") + + +def main(): + app.run(_main) + + +if __name__ == '__main__': + main() diff --git a/mjx/mujoco/mjx/viewer.py b/mjx/mujoco/mjx/viewer.py index 7ca85cf4..72af4649 100644 --- a/mjx/mujoco/mjx/viewer.py +++ b/mjx/mujoco/mjx/viewer.py @@ -28,7 +28,8 @@ _MODEL_PATH = flags.DEFINE_string('mjcf', None, 'Path to a MuJoCo MJCF file.', required=True) -def main(argv: Sequence[str]) -> None: +def _main(argv: Sequence[str]) -> None: + """Launches MuJoCo passive viewer fed by MJX.""" if len(argv) > 1: raise app.UsageError('Too many command-line arguments.') @@ -37,20 +38,15 @@ def main(argv: Sequence[str]) -> None: print(f'Loading model from: {_MODEL_PATH.value}.') m = mujoco.MjModel.from_xml_path(_MODEL_PATH.value) d = mujoco.MjData(m) + mx = mjx.put_model(m) + dx = mjx.put_data(m, d) - mx = mjx.device_put(m) - dx = mjx.make_data(mx) - - dt = jax.device_get(mx.opt.timestep) - step_fn = jax.jit(mjx.step) - - print(f'JAX default backend: {jax.default_backend()}') - print('JIT-compiling the MJX step (this may take a while)...') + print(f'Default backend: {jax.default_backend()}') + print('JIT-compiling the model physics step...') start = time.time() - dx = step_fn(mx, dx) - mjx.device_get_into(d, dx) + step_fn = jax.jit(mjx.step).lower(dx).compile() elapsed = time.time() - start - print(f'JIT compilation took {elapsed}s.') + print(f'Compilation took {elapsed}s.') with mujoco.viewer.launch_passive(m, d) as v: while True: @@ -72,9 +68,13 @@ def main(argv: Sequence[str]) -> None: v.sync() elapsed = time.time() - start - if elapsed < dt: - time.sleep(dt - elapsed) + if elapsed < m.opt.timestep: + time.sleep(m.opt.timestep - elapsed) + + +def main(): + app.run(_main) if __name__ == '__main__': - app.run(main) + main() diff --git a/mjx/pyproject.toml b/mjx/pyproject.toml index b0672ead..0c5c342b 100644 --- a/mjx/pyproject.toml +++ b/mjx/pyproject.toml @@ -36,6 +36,10 @@ dependencies = [ "trimesh", ] +[project.scripts] +mjx-testspeed = "mujoco.mjx.testspeed:main" +mjx-viewer = "mujoco.mjx.viewer:main" + [project.urls] Homepage = "https://github.com/google-deepmind/mujoco/tree/main/mjx" Documentation = "https://mujoco.readthedocs.io/en/3.1.2"