From a0f014593b1ebaacc4db0194168966fa348fbca2 Mon Sep 17 00:00:00 2001 From: Yuval Tassa Date: Mon, 8 Jun 2026 06:58:54 -0700 Subject: [PATCH] Support float32 in mujoco.rollout Python wrapper and tests. - Update rollout.py to use mujoco.MJTNUM_DTYPE instead of np.float64. - Update rollout_test.py to use DTYPE for all state/control/sensor arrays to avoid copies and type mismatches. PiperOrigin-RevId: 928539946 Change-Id: Iee2250c049de799cb0e23ffe5558fe6eea8b3fd8 --- python/mujoco/rollout.cc | 16 +-- python/mujoco/rollout.py | 8 +- python/mujoco/rollout_test.py | 186 ++++++++++++++++++---------------- 3 files changed, 114 insertions(+), 96 deletions(-) diff --git a/python/mujoco/rollout.cc b/python/mujoco/rollout.cc index 2bb4dbcf..0928e09b 100644 --- a/python/mujoco/rollout.cc +++ b/python/mujoco/rollout.cc @@ -13,6 +13,7 @@ // limitations under the License. #include +#include #include #include #include @@ -177,10 +178,11 @@ void _unsafe_rollout(std::vector& m, mjData* d, int start_roll, } // C-style threaded version of _unsafe_rollout -void _unsafe_rollout_threaded(std::vector& m, std::vector& d, - int nbatch, int nstep, unsigned int control_spec, - const mjtNum* state0, const mjtNum* warmstart0, - const mjtNum* control, mjtNum* state, mjtNum* sensordata, +void _unsafe_rollout_threaded(std::vector& m, + std::vector& d, int nbatch, int nstep, + unsigned int control_spec, const mjtNum* state0, + const mjtNum* warmstart0, const mjtNum* control, + mjtNum* state, mjtNum* sensordata, ThreadPool* pool, int chunk_size) { int nfulljobs = nbatch / chunk_size; int chunk_remainder = nbatch % chunk_size; @@ -209,7 +211,7 @@ void _unsafe_rollout_threaded(std::vector& m, std::vectorSchedule(task); } - // wait for job counter to incremented up to the number of jobs submitted by this thread + // wait for counter to increment up to number of jobs submitted by this thread pool->WaitCount(njobs); } @@ -227,8 +229,8 @@ mjtNum* get_array_ptr(std::optional> arg, py::buffer_info info = arg->request(); // check size - size_t expected_size = - static_cast(nbatch) * static_cast(nstep) * static_cast(dim); + size_t expected_size = static_cast(nbatch) * + static_cast(nstep) * static_cast(dim); if (info.size != expected_size) { std::ostringstream msg; msg << name << ".size should be " << expected_size << ", got " << info.size; diff --git a/python/mujoco/rollout.py b/python/mujoco/rollout.py index 24ed27e6..daf55504 100644 --- a/python/mujoco/rollout.py +++ b/python/mujoco/rollout.py @@ -219,9 +219,9 @@ class Rollout: # allocate output if not provided if state is None: - state = np.empty((nbatch, nstep, nstate)) + state = np.empty((nbatch, nstep, nstate), dtype=mujoco.MJTNUM_DTYPE) if sensordata is None: - sensordata = np.empty((nbatch, nstep, nsensordata)) + sensordata = np.empty((nbatch, nstep, nsensordata), dtype=mujoco.MJTNUM_DTYPE) # call rollout self.rollout_.rollout( @@ -376,7 +376,7 @@ def _ensure_2d(arg): if arg is None: return None else: - return np.ascontiguousarray(np.atleast_2d(arg), dtype=np.float64) + return np.ascontiguousarray(np.atleast_2d(arg), dtype=mujoco.MJTNUM_DTYPE) def _ensure_3d(arg): @@ -390,7 +390,7 @@ def _ensure_3d(arg): arg = arg[np.newaxis, np.newaxis, ...] elif arg.ndim == 2: arg = arg[np.newaxis, ...] - return np.ascontiguousarray(arg, dtype=np.float64) + return np.ascontiguousarray(arg, dtype=mujoco.MJTNUM_DTYPE) def _infer_dimension(dim, value, **kwargs): diff --git a/python/mujoco/rollout_test.py b/python/mujoco/rollout_test.py index f2b5358f..c069dfba 100644 --- a/python/mujoco/rollout_test.py +++ b/python/mujoco/rollout_test.py @@ -21,11 +21,27 @@ import threading from absl.testing import absltest from absl.testing import parameterized -import numpy as np import mujoco from mujoco import rollout +import numpy as np + + +DTYPE = mujoco.MJTNUM_DTYPE + + +def randn(*shape): + return np.random.standard_normal(shape).astype(DTYPE) + + +def empty(*shape): + return np.empty(shape, dtype=DTYPE) + + +def zeros(*shape): + return np.zeros(shape, dtype=DTYPE) + # -------------------------- models used for testing --------------------------- TEST_XML = r""" @@ -154,8 +170,8 @@ class MuJoCoRolloutTest(parameterized.TestCase): nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS) data = mujoco.MjData(model) - initial_state = np.random.randn(nstate) - control = np.random.randn(model.nu) + initial_state = randn(nstate) + control = randn(model.nu) state, sensordata = rollout.rollout(model, data, initial_state, control) mujoco.mj_resetData(model, data) @@ -171,8 +187,8 @@ class MuJoCoRolloutTest(parameterized.TestCase): nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS) data = mujoco.MjData(model) - initial_state = np.random.randn(nstate) - control = np.random.randn(nstep, model.nu) + initial_state = randn(nstate) + control = randn(nstep, model.nu) state, sensordata = rollout.rollout(model, data, initial_state, control) py_state, py_sensordata = py_rollout(model, data, initial_state, control) @@ -188,8 +204,8 @@ class MuJoCoRolloutTest(parameterized.TestCase): nbatch = 5 # number of rollouts nstep = 1 # number of steps - initial_state = np.random.randn(nbatch, nstate) - control = np.random.randn(nbatch, nstep, model.nu) + initial_state = randn(nbatch, nstate) + control = randn(nbatch, nstep, model.nu) state, sensordata = rollout.rollout(model, data, initial_state, control) mujoco.mj_resetData(model, data) @@ -206,8 +222,8 @@ class MuJoCoRolloutTest(parameterized.TestCase): nbatch = 5 # number of rollouts nstep = 1 # number of steps - initial_state = np.random.randn(nbatch, nstate) - control = np.random.randn(nstep, model.nu) + initial_state = randn(nbatch, nstate) + control = randn(nstep, model.nu) state, sensordata = rollout.rollout(model, data, initial_state, control) mujoco.mj_resetData(model, data) @@ -225,8 +241,8 @@ class MuJoCoRolloutTest(parameterized.TestCase): nbatch = 5 # number of rollouts nstep = 1 # number of steps - initial_state = np.random.randn(nstate) - control = np.random.randn(nbatch, nstep, model.nu) + initial_state = randn(nstate) + control = randn(nbatch, nstep, model.nu) state, sensordata = rollout.rollout(model, data, initial_state, control) mujoco.mj_resetData(model, data) @@ -244,8 +260,8 @@ class MuJoCoRolloutTest(parameterized.TestCase): nbatch = 5 # number of rollouts nstep = 1 # number of steps - initial_state = np.random.randn(nstate) - control = np.random.randn(nstep, model.nu) + initial_state = randn(nstate) + control = randn(nstep, model.nu) initial_warmstart = np.tile(data.qacc_warmstart.copy(), (nbatch, 1)) state, sensordata = rollout.rollout( model, data, initial_state, control, initial_warmstart=initial_warmstart @@ -267,9 +283,9 @@ class MuJoCoRolloutTest(parameterized.TestCase): nbatch = 5 # number of rollouts nstep = 1 # number of steps - initial_state = np.random.randn(nstate) - control = np.random.randn(nstep, model.nu) - state = np.empty((nbatch, nstep, nstate)) + initial_state = randn(nstate) + control = randn(nstep, model.nu) + state = empty(nbatch, nstep, nstate) state, sensordata = rollout.rollout( model, data, initial_state, control, state=state ) @@ -290,9 +306,9 @@ class MuJoCoRolloutTest(parameterized.TestCase): nbatch = 5 # number of rollouts nstep = 1 # number of steps - initial_state = np.random.randn(nstate) - control = np.random.randn(nstep, model.nu) - sensordata = np.empty((nbatch, nstep, model.nsensordata)) + initial_state = randn(nstate) + control = randn(nstep, model.nu) + sensordata = empty(nbatch, nstep, model.nsensordata) state, sensordata = rollout.rollout( model, data, initial_state, control, sensordata=sensordata ) @@ -313,10 +329,10 @@ class MuJoCoRolloutTest(parameterized.TestCase): nbatch = 1 # number of rollouts nstep = 3 # number of steps - initial_state = np.random.randn(nstate) - control = np.random.randn(model.nu) - state = np.empty((nbatch, nstep, nstate)) - sensordata = np.empty((nbatch, nstep, model.nsensordata)) + initial_state = randn(nstate) + control = randn(model.nu) + state = empty(nbatch, nstep, nstate) + sensordata = empty(nbatch, nstep, model.nsensordata) rollout.rollout( model, data, initial_state, control, state=state, sensordata=sensordata ) @@ -335,8 +351,8 @@ class MuJoCoRolloutTest(parameterized.TestCase): nbatch = 2 # number of initial states nstep = 3 # number of timesteps - initial_state = np.random.randn(nbatch, nstate) - control = np.random.randn(nbatch, nstep, model.nu) + initial_state = randn(nbatch, nstate) + control = randn(nbatch, nstep, model.nu) state, sensordata = rollout.rollout(model, data, initial_state, control) py_state, py_sensordata = py_rollout(model, data, initial_state, control) @@ -363,8 +379,8 @@ class MuJoCoRolloutTest(parameterized.TestCase): nstate = mujoco.mj_stateSize(model[0], mujoco.mjtState.mjSTATE_FULLPHYSICS) data = mujoco.MjData(model[0]) - initial_state = np.random.randn(nbatch, nstate) - control = np.random.randn(nbatch, nstep, model[0].nu) + initial_state = randn(nbatch, nstate) + control = randn(nbatch, nstep, model[0].nu) state, sensordata = rollout.rollout(model, data, initial_state, control) py_state, py_sensordata = py_rollout(model, data, initial_state, control) @@ -380,9 +396,9 @@ class MuJoCoRolloutTest(parameterized.TestCase): nbatch = 2 # number of rollouts nstep = 3 # number of timesteps - initial_state = np.random.randn(nbatch, nstate) - control = np.random.randn(nbatch, 1, model.nu) - state = np.empty((nbatch, nstep, nstate)) + initial_state = randn(nbatch, nstate) + control = randn(nbatch, 1, model.nu) + state = empty(nbatch, nstep, nstate) state, sensordata = rollout.rollout( model, data, initial_state, control, state=state ) @@ -401,7 +417,7 @@ class MuJoCoRolloutTest(parameterized.TestCase): nbatch = 4 # number of rollouts nstep = 3 # number of timesteps - initial_state = np.random.randn(nbatch, nstate) + initial_state = randn(nbatch, nstate) control_spec = ( mujoco.mjtState.mjSTATE_CTRL @@ -409,7 +425,7 @@ class MuJoCoRolloutTest(parameterized.TestCase): | mujoco.mjtState.mjSTATE_XFRC_APPLIED ) ncontrol = mujoco.mj_stateSize(model, control_spec) - control = np.random.randn(nbatch, nstep, ncontrol) + control = randn(nbatch, nstep, ncontrol) state, sensordata = rollout.rollout( model, data, initial_state, control, control_spec=control_spec @@ -427,7 +443,7 @@ class MuJoCoRolloutTest(parameterized.TestCase): data = mujoco.MjData(model) nbatch = 4 # number of rollouts - initial_state = np.empty((nbatch, nstate)) + initial_state = empty(nbatch, nstate) # get diverging (0, 2) and non-diverging (1, 3) states mujoco.mj_getState( @@ -446,7 +462,7 @@ class MuJoCoRolloutTest(parameterized.TestCase): nstep = 10000 # divergence after ~15s, timestep = 2e-3 - state = np.random.randn(nbatch, nstep, nstate) + state = randn(nbatch, nstep, nstate) rollout.rollout(model, data, initial_state, state=state) @@ -466,10 +482,10 @@ class MuJoCoRolloutTest(parameterized.TestCase): num_workers = 32 nbatch = 100 nstep = 5 - initial_state = np.random.randn(nbatch, nstate) - state = np.empty((nbatch, nstep, nstate)) - sensordata = np.empty((nbatch, nstep, model.nsensordata)) - control = np.random.randn(nbatch, nstep, model.nu) + initial_state = randn(nbatch, nstate) + state = empty(nbatch, nstep, nstate) + sensordata = empty(nbatch, nstep, model.nsensordata) + control = randn(nbatch, nstep, model.nu) thread_local = threading.local() @@ -528,10 +544,10 @@ class MuJoCoRolloutTest(parameterized.TestCase): num_workers = 32 nbatch = 100 nstep = 5 - initial_state = np.random.randn(nbatch, nstate) - state = np.empty((nbatch, nstep, nstate)) - sensordata = np.empty((nbatch, nstep, model.nsensordata)) - control = np.random.randn(nbatch, nstep, model.nu) + initial_state = randn(nbatch, nstate) + state = empty(nbatch, nstep, nstate) + sensordata = empty(nbatch, nstep, model.nsensordata) + control = randn(nbatch, nstep, model.nu) model_list = [copy.copy(model) for _ in range(nbatch)] data_list = [mujoco.MjData(model) for _ in range(num_workers)] @@ -557,10 +573,10 @@ class MuJoCoRolloutTest(parameterized.TestCase): num_workers = 32 nbatch = 100 nstep = 5 - initial_state = np.random.randn(nbatch, nstate) - state = np.empty((nbatch, nstep, nstate)) - sensordata = np.empty((nbatch, nstep, model.nsensordata)) - control = np.random.randn(nbatch, nstep, model.nu) + initial_state = randn(nbatch, nstate) + state = empty(nbatch, nstep, nstate) + sensordata = empty(nbatch, nstep, model.nsensordata) + control = randn(nbatch, nstep, model.nu) model_list = [copy.copy(model) for _ in range(nbatch)] data_list = [mujoco.MjData(model) for _ in range(num_workers)] @@ -606,10 +622,10 @@ class MuJoCoRolloutTest(parameterized.TestCase): num_workers = 32 nbatch = 100 nstep = 5 - initial_state = np.random.randn(nbatch, nstate) - state = np.empty((nbatch, nstep, nstate)) - sensordata = np.empty((nbatch, nstep, model.nsensordata)) - control = np.random.randn(nbatch, nstep, model.nu) + initial_state = randn(nbatch, nstate) + state = empty(nbatch, nstep, nstate) + sensordata = empty(nbatch, nstep, model.nsensordata) + control = randn(nbatch, nstep, model.nu) model_list = [copy.copy(model) for _ in range(nbatch)] data_list = [mujoco.MjData(model) for _ in range(num_workers)] @@ -640,8 +656,8 @@ class MuJoCoRolloutTest(parameterized.TestCase): data = mujoco.MjData(model) # take one step, save the state - state0 = np.zeros(nstate) - control = np.zeros(model.nu) + state0 = zeros(nstate) + control = zeros(model.nu) state1, _ = step(model, data, state0, control) # save qacc_warmstart @@ -671,17 +687,17 @@ class MuJoCoRolloutTest(parameterized.TestCase): nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS) data = mujoco.MjData(model) - initial_state = np.zeros(nstate) + initial_state = zeros(nstate) control_spec = ( mujoco.mjtState.mjSTATE_MOCAP_POS | mujoco.mjtState.mjSTATE_MOCAP_QUAT ) - pos1 = np.array((1.0, 2.0, 3.0)) - quat1 = np.array((1.0, 2.0, 3.0, 4.0)) + pos1 = np.array((1.0, 2.0, 3.0), dtype=DTYPE) + quat1 = np.array((1.0, 2.0, 3.0, 4.0), dtype=DTYPE) quat1 /= np.linalg.norm(quat1) - pos2 = np.array((2.0, 3.0, 4.0)) - quat2 = np.array((2.0, 3.0, 4.0, 5.0)) + pos2 = np.array((2.0, 3.0, 4.0), dtype=DTYPE) + quat2 = np.array((2.0, 3.0, 4.0, 5.0), dtype=DTYPE) quat2 /= np.linalg.norm(quat2) control = np.hstack((pos1, pos2, quat1, quat2)) @@ -702,8 +718,8 @@ class MuJoCoRolloutTest(parameterized.TestCase): nbatch = 1 nstep = 3 - initial_state = np.zeros((nbatch, nstate)) - ctrl = np.zeros((nbatch, nstep, model.nu)) + initial_state = zeros(nbatch, nstate) + ctrl = zeros(nbatch, nstep, model.nu) model.opt.solver = 10 # invalid solver type with self.assertRaisesWithLiteralMatch( @@ -718,7 +734,7 @@ class MuJoCoRolloutTest(parameterized.TestCase): nbatch = 1 - initial_state = np.zeros((nbatch, nstate)) + initial_state = zeros(nbatch, nstate) control = 'string' with self.assertRaisesWithLiteralMatch( @@ -726,7 +742,7 @@ class MuJoCoRolloutTest(parameterized.TestCase): ): rollout.rollout(model, data, initial_state, control) - control = np.zeros((2, 3, 4, 5)) + control = zeros(2, 3, 4, 5) with self.assertRaisesWithLiteralMatch( ValueError, 'control can have at most 3 dimensions' ): @@ -740,28 +756,28 @@ class MuJoCoRolloutTest(parameterized.TestCase): nbatch = 1 nstep = 3 - initial_state = np.random.randn(nbatch, nstate + 1) + initial_state = randn(nbatch, nstate + 1) with self.assertRaisesWithLiteralMatch( ValueError, 'trailing dimension of initial_state must be 6, got 7' ): rollout.rollout(model, data, initial_state) - initial_state = np.random.randn(nbatch, nstate) - control = np.random.randn(1, nstep, model.nu + 1) + initial_state = randn(nbatch, nstate) + control = randn(1, nstep, model.nu + 1) with self.assertRaisesWithLiteralMatch( ValueError, 'trailing dimension of control must be 2, got 3' ): rollout.rollout(model, data, initial_state, control) - control = np.random.randn(nbatch, nstep, model.nu) - state = np.random.randn(nbatch, nstep + 1, nstate) # incompatible nstep + control = randn(nbatch, nstep, model.nu) + state = randn(nbatch, nstep + 1, nstate) # incompatible nstep with self.assertRaisesWithLiteralMatch( ValueError, 'dimension 1 inferred as 3 but state has 4' ): rollout.rollout(model, data, initial_state, control, state=state) - initial_state = np.random.randn(nbatch, nstate) - control = np.random.randn(nbatch, nstep, model.nu) + initial_state = randn(nbatch, nstate) + control = randn(nbatch, nstep, model.nu) bad_spec = mujoco.mjtState.mjSTATE_ACT with self.assertRaisesWithLiteralMatch( ValueError, 'control_spec can only contain bits in mjSTATE_USER' @@ -776,8 +792,8 @@ class MuJoCoRolloutTest(parameterized.TestCase): data = mujoco.MjData(model) # step with a clean mjData - initial_state = np.random.randn(nstate) - control = np.random.randn(3, 3, model.nu) + initial_state = randn(nstate) + control = randn(3, 3, model.nu) state, sensordata = rollout.rollout(model, data, initial_state, control) # fill user fields with random values @@ -788,7 +804,7 @@ class MuJoCoRolloutTest(parameterized.TestCase): 'mocap_pos', 'mocap_quat', ]: - setattr(data, attr, np.random.randn(*getattr(data, attr).shape)) + setattr(data, attr, randn(*getattr(data, attr).shape)) # roll out again state2, sensordata2 = rollout.rollout(model, data, initial_state, control) @@ -802,8 +818,8 @@ class MuJoCoRolloutTest(parameterized.TestCase): nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS) data = mujoco.MjData(model) - initial_state = np.random.randn(nstate) - control = np.random.randn(3, 3, model.nu) + initial_state = randn(nstate) + control = randn(3, 3, model.nu) state, sensordata = rollout.rollout(model, data, initial_state, control) state2, sensordata2 = rollout.rollout([model], data, initial_state, control) @@ -817,8 +833,8 @@ class MuJoCoRolloutTest(parameterized.TestCase): nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS) data = mujoco.MjData(model) - initial_state = np.random.randn(nstate) - control = np.random.randn(3, 3, model.nu) + initial_state = randn(nstate) + control = randn(3, 3, model.nu) # Test passing empty lists for data with self.assertRaisesWithLiteralMatch( @@ -852,7 +868,7 @@ class MuJoCoRolloutTest(parameterized.TestCase): ): with rollout.Rollout(nthread=0) as rollout_: rollout_.rollout( - model, [copy.copy(data) for i in range(2)], initial_state, control + model, [copy.copy(data) for _ in range(2)], initial_state, control ) with self.assertRaisesWithLiteralMatch( @@ -872,7 +888,7 @@ class MuJoCoRolloutTest(parameterized.TestCase): ): with rollout.Rollout(nthread=2) as rollout_: rollout_.rollout( - model, [copy.copy(data) for i in range(3)], initial_state, control + model, [copy.copy(data) for _ in range(3)], initial_state, control ) @absltest.skip(reason='Takes a long time to run') @@ -887,7 +903,7 @@ class MuJoCoRolloutTest(parameterized.TestCase): nstep = ((2**31) // (nstate * nbatch)) + 2 assert nstep * nstate * nbatch > 2**31 - initial_state = np.random.randn(nbatch, nstate) + initial_state = randn(nbatch, nstate) rollout.rollout( model, [copy.copy(data) for _ in range(nthread)], @@ -901,7 +917,7 @@ class MuJoCoRolloutTest(parameterized.TestCase): def get_state(model, data): nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS) - state = np.empty(nstate) + state = empty(nstate) mujoco.mj_getState(model, data, state, mujoco.mjtState.mjSTATE_FULLPHYSICS) return state.reshape((1, nstate)) @@ -925,8 +941,8 @@ def one_rollout( ): nstep = control.shape[0] nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS) - state = np.empty((nstep, nstate)) - sensordata = np.empty((nstep, model.nsensordata)) + state = empty(nstep, nstate) + sensordata = empty(nstep, model.nsensordata) mujoco.mj_resetData(model, data) for t in range(nstep): @@ -940,7 +956,7 @@ def ensure_2d(arg): if arg is None: return None else: - return np.ascontiguousarray(np.atleast_2d(arg), dtype=np.float64) + return np.ascontiguousarray(np.atleast_2d(arg), dtype=DTYPE) def ensure_3d(arg): @@ -954,7 +970,7 @@ def ensure_3d(arg): arg = arg[np.newaxis, np.newaxis, ...] elif arg.ndim == 2: arg = arg[np.newaxis, ...] - return np.ascontiguousarray(arg, dtype=np.float64) + return np.ascontiguousarray(arg, dtype=DTYPE) def py_rollout( @@ -974,8 +990,8 @@ def py_rollout( nstate = mujoco.mj_stateSize(model[0], mujoco.mjtState.mjSTATE_FULLPHYSICS) - state = np.empty((nbatch, nstep, nstate)) - sensordata = np.empty((nbatch, nstep, model[0].nsensordata)) + state = empty(nbatch, nstep, nstate) + sensordata = empty(nbatch, nstep, model[0].nsensordata) for r in range(nbatch): state_r, sensordata_r = one_rollout( model[r], data, initial_state[r], control[r], control_spec