Add more info to doc about static fields. Fixes #1615.

PiperOrigin-RevId: 627528282
Change-Id: I88aa1ece0e530a53a3ae18c8dc935d3ea9c56c09
This commit is contained in:
Baruch Tabanpour
2024-04-23 15:46:37 -07:00
committed by Copybara-Service
parent 80810391ab
commit b56d7d6e30
+14 -7
View File
@@ -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 <mjxFeatureParity>` 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