From b56d7d6e308d128c687e2108f28ad42c4b32b668 Mon Sep 17 00:00:00 2001 From: Baruch Tabanpour Date: Tue, 23 Apr 2024 15:46:37 -0700 Subject: [PATCH] Add more info to doc about static fields. Fixes #1615. PiperOrigin-RevId: 627528282 Change-Id: I88aa1ece0e530a53a3ae18c8dc935d3ea9c56c09 --- doc/mjx.rst | 21 ++++++++++++++------- 1 file changed, 14 insertions(+), 7 deletions(-) diff --git a/doc/mjx.rst b/doc/mjx.rst index 1628c8ea..a2e79d8e 100644 --- a/doc/mjx.rst +++ b/doc/mjx.rst @@ -66,25 +66,32 @@ directly from the top-level ``mjx`` module. Structs ------- -Before running MJX functions on an accelerator device, structs must be copied onto the device via the ``mjx.device_put`` -function. Placing an :ref:`mjModel` on device yields an ``mjx.Model``. Placing an :ref:`mjData` on device yields +Before running MJX functions on an accelerator device, structs must be copied onto the device via the ``mjx.put_model`` and ``mjx.put_data`` +functions. Placing an :ref:`mjModel` on device yields an ``mjx.Model``. Placing an :ref:`mjData` on device yields an ``mjx.Data``: .. code-block:: python model = mujoco.MjModel.from_xml_string("...") data = mujoco.MjData(model) - mjx_model = mjx.device_put(model) - mjx_data = mjx.device_put(data) + mjx_model = mjx.put_model(model) + mjx_data = mjx.put_data(data) -These MJX variants mirror their MuJoCo counterparts but have three key differences: +These MJX variants mirror their MuJoCo counterparts but have a few key differences: -#. Fields in ``mjx.Model`` and ``mjx.Data`` are JAX arrays copied onto device, instead of numpy arrays. +#. ``mjx.Model`` and ``mjx.Data`` contain JAX arrays that are copied onto device. #. Some fields are missing from ``mjx.Model`` and ``mjx.Data`` for features that are :ref:`unsupported ` in MJX. -#. Arrays in ``mjx.Model`` and ``mjx.Data`` support adding batch dimensions. Batch dimensions are a natural way to +#. JAX arrays in ``mjx.Model`` and ``mjx.Data`` support adding batch dimensions. Batch dimensions are a natural way to express domain randomization (in the case of ``mjx.Model``) or high-throughput simulation for reinforcement learning (in the case of ``mjx.Data``). +#. Numpy arrays in ``mjx.Model`` and ``mjx.Data`` are structural fields that control the output of JIT compilation. + Modifying these arrays will force JAX to recompile MJX functions. As an example, + ``jnt_limited`` is a numpy array passed by reference from :ref:`mjModel`, which determines if joint limit + constraints should be applied. If ``jnt_limited`` is modified, JAX will + re-compile MJX functions. + On the other hand, ``jnt_range`` is a JAX array that can be modified at runtime, and will only apply to joints with limits + as specified by the ``jnt_limited`` field. Neither ``mjx.Model`` nor ``mjx.Data`` are meant to be constructed manually. An ``mjx.Data`` may be created by calling